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..5ce8b6c84ba 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -8,6 +8,7 @@ legacy_flags=( enterprise-package enterprise-routing mcp-integration + misc proxy-db-auth-checks proxy-db-budgets proxy-db-custom-logging @@ -22,6 +23,7 @@ legacy_flags=( proxy-db-proxy-utils proxy-extras proxy-infra + responses-caching-types ) legacy_paths() { @@ -36,6 +38,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 @@ -48,9 +51,28 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_managed_files_access_check.py echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; 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 +135,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..10ee19f146a 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,13 @@ 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-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..8ed7b917460 100644 --- a/.github/merge-smoke-tests.json +++ b/.github/merge-smoke-tests.json @@ -5,8 +5,8 @@ "CHAT-TOOL-STREAM": "tests/test_litellm/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/scripts/read_rc_version.py b/.github/scripts/read_rc_version.py new file mode 100644 index 00000000000..02b5a067427 --- /dev/null +++ b/.github/scripts/read_rc_version.py @@ -0,0 +1,44 @@ +#!/usr/bin/env python3 +"""Print `version=X.Y.0` from [project].version in pyproject.toml for $GITHUB_OUTPUT. + +Usage +----- + python3 read_rc_version.py [path/to/pyproject.toml] >> "$GITHUB_OUTPUT" + +Exit code 1 with a `::error::` line on stderr when the version is not an X.Y.0 release. +""" + +from __future__ import annotations + +import pathlib +import re +import sys +from typing import Final + +if sys.version_info >= (3, 11): + import tomllib +else: + import tomli as tomllib + +RELEASE_VERSION: Final = re.compile(r"[0-9]+\.[0-9]+\.0") + + +def read_version(pyproject: pathlib.Path) -> str: + with pyproject.open("rb") as f: + return tomllib.load(f)["project"]["version"] + + +def main(argv: list[str]) -> int: + pyproject: Final = pathlib.Path(argv[1]) if len(argv) > 1 else pathlib.Path("pyproject.toml") + version: Final = read_version(pyproject) + if RELEASE_VERSION.fullmatch(version) is None: + print( # noqa: T201 # the ::error:: line to stderr is the workflow's failure signal + f"::error::pyproject.toml version {version} is not an X.Y.0 release version", file=sys.stderr + ) + return 1 + print(f"version={version}") # noqa: T201 # stdout line is appended to $GITHUB_OUTPUT + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv)) 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/create-rc-branch.yml b/.github/workflows/create-rc-branch.yml new file mode 100644 index 00000000000..53760ad553e --- /dev/null +++ b/.github/workflows/create-rc-branch.yml @@ -0,0 +1,66 @@ +name: Create RC Branch + +on: + schedule: + - cron: "0 3 * * 5" + timezone: "America/Los_Angeles" + workflow_dispatch: + +permissions: {} + +jobs: + create-rc-branch: + name: Create RC Branch + if: github.event_name != 'schedule' || github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + permissions: + contents: write + steps: + - name: Require main + env: + REF: ${{ github.ref }} + run: | + if [ "$REF" != "refs/heads/main" ]; then + echo "::error::rc branches are cut from refs/heads/main only, got $REF" + exit 1 + fi + + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + + - name: Read release version + id: version + run: python3 .github/scripts/read_rc_version.py >> "$GITHUB_OUTPUT" + + - name: Create rc branch + env: + VERSION: ${{ steps.version.outputs.version }} + uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1 + with: + script: | + const branchName = `rc/${process.env.VERSION}`; + const ref = `heads/${branchName}`; + + const existing = await github.rest.git.getRef({ + owner: context.repo.owner, + repo: context.repo.repo, + ref, + }).catch((error) => { + if (error.status === 404) { + return null; + } + throw error; + }); + if (existing !== null) { + core.setFailed(`Branch ${branchName} already exists at ${existing.data.object.sha}; leaving it untouched`); + return; + } + + await github.rest.git.createRef({ + owner: context.repo.owner, + repo: context.repo.repo, + ref: `refs/${ref}`, + sha: context.sha, + }); + core.info(`Created branch ${branchName} at ${context.sha}`); diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index 0b58cf9d486..2f5ce4d441a 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -8,9 +8,13 @@ on: paths: - "litellm/_redis.py" - "litellm/_redis_credential_provider.py" - - "tests/test_litellm/test_redis.py" + - "litellm/caching/redis_cache.py" + - "litellm/caching/evicted_client_closer.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" + - "tests/test_litellm/caching/test_evicted_client_closer.py" - ".github/workflows/test-redis-compat.yml" - "pyproject.toml" - "uv.lock" @@ -80,8 +84,10 @@ 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 \ tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \ tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \ --tb=short -vv \ 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..91b54f4ee70 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 @@ -106,26 +105,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 +191,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 +200,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 +209,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 +218,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 +229,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 +237,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/AGENTS.md b/AGENTS.md index 820ea64d4f9..69e034fbdea 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -96,6 +96,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega - No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc. - Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: ` - Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: ` + - Comprehensions take at most one `for` clause and one `if` clause (LIT014); split stacked clauses into a helper generator, a named intermediate, or a plain loop. Suppress with `# comprehension-ok: ` only when unavoidable - Use dependency injection - Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed - Use tagged unions + match diff --git a/Makefile b/Makefile index 28daf589a23..62e6ae53275 100644 --- a/Makefile +++ b/Makefile @@ -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/enterprise/pyproject.toml b/enterprise/pyproject.toml index 8509600ad96..e5e54a3df2c 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.70" +version = "0.1.71" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index e9a4ff90b9e..2835715ef30 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.101" +version = "0.4.102" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index e5ffcd1c57a..70fcc367905 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -8,3 +8,11 @@ - Split a mixed test file along that line instead of widening visibility to move it - A test for another crate's item belongs in that crate, not in a downstream one - Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests//mod.rs` or `tests//support.rs` so it is not picked up as a test crate of its own + +## Error definitions + +- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` +- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each +- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string +- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return +- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0a91f0759c2..e7d911f5fd9 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -73,6 +73,15 @@ version = "1.0.104" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" +[[package]] +name = "arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1" +dependencies = [ + "derive_arbitrary", +] + [[package]] name = "arc-swap" version = "1.9.2" @@ -897,6 +906,12 @@ dependencies = [ "hybrid-array", ] +[[package]] +name = "borrow-or-share" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc0b364ead1874514c8c2855ab558056ebfeb775653e7ae45ff72f28f8f3166c" + [[package]] name = "bstr" version = "1.13.1" @@ -914,6 +929,12 @@ version = "3.20.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" +[[package]] +name = "bytecount" +version = "0.6.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "175812e0be2bccb6abe50bb8d566126198344f707e304f45c648fd8f2cc0365e" + [[package]] name = "byteorder" version = "1.5.0" @@ -1458,6 +1479,17 @@ dependencies = [ "serde_core", ] +[[package]] +name = "derive_arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "derive_builder" version = "0.20.2" @@ -1540,6 +1572,15 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "email_address" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e079f19b08ca6239f47f8ba8509c11cf3ea30095831f7fed61441475edd8c449" +dependencies = [ + "serde", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -1622,6 +1663,16 @@ version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" +[[package]] +name = "filetime" +version = "0.2.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759" +dependencies = [ + "cfg-if", + "libc", +] + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1639,6 +1690,17 @@ dependencies = [ "zlib-rs", ] +[[package]] +name = "fluent-uri" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc74ac4d8359ae70623506d512209619e5cf8f347124910440dbc221714b328e" +dependencies = [ + "borrow-or-share", + "ref-cast", + "serde", +] + [[package]] name = "fnv" version = "1.0.7" @@ -1660,6 +1722,16 @@ dependencies = [ "percent-encoding", ] +[[package]] +name = "fraction" +version = "0.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e246562084dde8ebbcc943b261c406ce4f68e5032ec28029a251a47d6a295500" +dependencies = [ + "num", + "num-bigint 0.4.8", +] + [[package]] name = "fs_extra" version = "1.3.0" @@ -1817,9 +1889,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" dependencies = [ "cfg-if", + "js-sys", "libc", "r-efi 5.3.0", "wasip2", + "wasm-bindgen", ] [[package]] @@ -2660,6 +2734,59 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "jsonschema" +version = "0.55.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b68339c3d874e48151d74ffe256d93a58cffa240983cb0967d3cbaea083a44fe" +dependencies = [ + "ahash", + "bytecount", + "data-encoding", + "email_address", + "fancy-regex 0.19.2", + "fraction", + "getrandom 0.3.4", + "itoa", + "jsonschema-regex", + "jsonschema-value", + "num-cmp", + "num-traits", + "percent-encoding", + "referencing", + "regex", + "serde", + "serde_json", + "strum", + "unicode-general-category", + "uuid-simd", +] + +[[package]] +name = "jsonschema-regex" +version = "0.55.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6307b5b51216ec9b941b52244c74043fa0b1d6b657b56199f57cb1416d3641c5" +dependencies = [ + "regex-syntax", +] + +[[package]] +name = "jsonschema-value" +version = "0.55.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0230ac05e09c6111e96c147b75c390579f5cbd45654b980c68ac60fe17b3f129" +dependencies = [ + "ahash", + "bytecount", + "fraction", + "getrandom 0.3.4", + "num-cmp", + "num-traits", + "serde_json", + "zmij", +] + [[package]] name = "lazy_static" version = "1.5.0" @@ -2995,6 +3122,7 @@ dependencies = [ "tokio-tungstenite", "url", "veil", + "wiremock", ] [[package]] @@ -3013,6 +3141,15 @@ dependencies = [ "url", ] +[[package]] +name = "litellm-coroutine" +version = "0.1.0" +dependencies = [ + "rstest", + "thiserror 2.0.19", + "tokio", +] + [[package]] name = "litellm-cost" version = "0.1.0" @@ -3040,6 +3177,7 @@ name = "litellm-host" version = "0.1.0" dependencies = [ "litellm-auth", + "litellm-coroutine", "rstest", "serde_json", "tokio", @@ -3049,6 +3187,7 @@ dependencies = [ name = "litellm-host-python" version = "0.1.0" dependencies = [ + "bytes", "futures-util", "litellm-host", "pyo3", @@ -3115,14 +3254,14 @@ dependencies = [ name = "litellm-model-catalog" version = "0.1.0" dependencies = [ - "criterion", "indexmap 2.14.0", - "litellm-model-catalog", + "jsonschema", "rstest", "schemars 1.2.2", "serde", "serde_json", "thiserror 2.0.19", + "time", ] [[package]] @@ -3344,6 +3483,27 @@ dependencies = [ "veil", ] +[[package]] +name = "litellm-testkit" +version = "0.1.0" +dependencies = [ + "flate2", + "futures-util", + "reqwest 0.12.28", + "rstest", + "semver", + "serde", + "serde_json", + "sha2 0.10.9", + "tar", + "target-lexicon", + "tempfile", + "thiserror 2.0.19", + "tokio", + "toml", + "zip", +] + [[package]] name = "litellm-token-counter" version = "0.1.0" @@ -3493,6 +3653,12 @@ version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" +[[package]] +name = "micromap" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a86d3146ed3995b5913c414f6664344b9617457320782e64f0bb44afd49d74" + [[package]] name = "mime" version = "0.3.17" @@ -3588,6 +3754,20 @@ dependencies = [ "minimal-lexical", ] +[[package]] +name = "num" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "35bd024e8b2ff75562e5f34e7f4905839deb4b22955ef5e73d2fea1b9813cb23" +dependencies = [ + "num-bigint 0.4.8", + "num-complex", + "num-integer", + "num-iter", + "num-rational", + "num-traits", +] + [[package]] name = "num-bigint" version = "0.4.8" @@ -3608,6 +3788,12 @@ dependencies = [ "num-traits", ] +[[package]] +name = "num-cmp" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63335b2e2c34fae2fb0aa2cecfd9f0832a1e24b3b32ecec612c3426d46dc8aaa" + [[package]] name = "num-complex" version = "0.4.6" @@ -3632,6 +3818,27 @@ dependencies = [ "num-traits", ] +[[package]] +name = "num-iter" +version = "0.1.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b" +dependencies = [ + "num-integer", + "num-traits", +] + +[[package]] +name = "num-rational" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f83d14da390562dca69fc84082e73e548e1ad308d24accdedd2720017cb37824" +dependencies = [ + "num-bigint 0.4.8", + "num-integer", + "num-traits", +] + [[package]] name = "num-traits" version = "0.2.19" @@ -4447,6 +4654,23 @@ dependencies = [ "syn 3.0.0", ] +[[package]] +name = "referencing" +version = "0.55.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a196a5b4a8a12f46b6353174df865a05d41a6055aff212ec30877492788618b6" +dependencies = [ + "ahash", + "fluent-uri", + "getrandom 0.3.4", + "hashbrown 0.17.1", + "itoa", + "micromap", + "parking_lot", + "percent-encoding", + "serde_json", +] + [[package]] name = "regex" version = "1.13.1" @@ -5033,6 +5257,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_spanned" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26" +dependencies = [ + "serde_core", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -5364,6 +5597,17 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" +[[package]] +name = "tar" +version = "0.4.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840" +dependencies = [ + "filetime", + "libc", + "xattr", +] + [[package]] name = "target-lexicon" version = "0.13.5" @@ -5633,6 +5877,30 @@ dependencies = [ "tokio", ] +[[package]] +name = "toml" +version = "0.9.12+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863" +dependencies = [ + "indexmap 2.14.0", + "serde_core", + "serde_spanned", + "toml_datetime 0.7.5+spec-1.1.0", + "toml_parser", + "toml_writer", + "winnow 0.7.15", +] + +[[package]] +name = "toml_datetime" +version = "0.7.5+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347" +dependencies = [ + "serde_core", +] + [[package]] name = "toml_datetime" version = "1.1.1+spec-1.1.0" @@ -5649,9 +5917,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b" dependencies = [ "indexmap 2.14.0", - "toml_datetime", + "toml_datetime 1.1.1+spec-1.1.0", "toml_parser", - "winnow", + "winnow 1.0.4", ] [[package]] @@ -5660,9 +5928,15 @@ version = "1.1.3+spec-1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56" dependencies = [ - "winnow", + "winnow 1.0.4", ] +[[package]] +name = "toml_writer" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2" + [[package]] name = "tonic" version = "0.14.6" @@ -5930,6 +6204,12 @@ version = "2.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" +[[package]] +name = "unicode-general-category" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b993bddc193ae5bd0d623b49ec06ac3e9312875fdae725a975c51db1cc1677f" + [[package]] name = "unicode-ident" version = "1.0.24" @@ -6010,6 +6290,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "uuid-simd" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23b082222b4f6619906941c17eb2297fff4c2fb96cb60164170522942a200bd8" +dependencies = [ + "outref", + "vsimd", +] + [[package]] name = "valuable" version = "0.1.1" @@ -6408,6 +6698,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" + [[package]] name = "winnow" version = "1.0.4" @@ -6470,6 +6766,16 @@ dependencies = [ "time", ] +[[package]] +name = "xattr" +version = "1.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" +dependencies = [ + "libc", + "rustix", +] + [[package]] name = "xmlparser" version = "0.13.6" @@ -6595,6 +6901,23 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "zip" +version = "2.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50" +dependencies = [ + "arbitrary", + "crc32fast", + "crossbeam-utils", + "displaydoc", + "flate2", + "indexmap 2.14.0", + "memchr", + "thiserror 2.0.19", + "zopfli", +] + [[package]] name = "zlib-rs" version = "0.6.7" @@ -6606,3 +6929,15 @@ name = "zmij" version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + +[[package]] +name = "zopfli" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249" +dependencies = [ + "bumpalo", + "crc32fast", + "log", + "simd-adler32", +] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 9d05c8d2b98..022e8f13311 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm" litellm-tracing = { path = "crates/tracing" } tracing = "0.1" litellm-core = { path = "crates/core" } +litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } litellm-framing = { path = "crates/framer" } @@ -80,6 +81,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" +flate2 = "1" +semver = "1" +tar = "0.4" +target-lexicon = "0.13.5" +tempfile = "3" +zip = { version = "2", default-features = false, features = ["deflate"] } moka = { version = "0.12.16", features = ["future"] } strum = { version = "0.28.0", features = ["derive"] } url = "2.5.8" diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index b37790f60a8..9b921070839 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -3,8 +3,8 @@ //! lifetime. No other callback host has that obligation, which is why nothing outside //! this crate holds them. -use litellm_host::{machine::Machine, route::Route}; -use litellm_host_python::{RouteHost, lookup, run_call}; +use litellm_host::{machine::Machine, protocol::Protocol}; +use litellm_host_python::{ProtocolHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -63,25 +63,25 @@ impl PublicCall { } } -/// Runs one native call under the legacy `Logging` contract: the route host projects from +/// Runs one native call under the legacy `Logging` contract: the protocol host projects from /// the keyword view the contract prepares, and the contract observes the call. pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, machine: M, - route: H, + host: H, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine::Response> + 'static, + H: ProtocolHost + 'static, + M: Machine::Response> + 'static, { let arguments = call.kwargs.clone_ref(py); run_call( py, machine, - route, + host, Box::new(LegacyLogging::new(py, surface, call, asynchronous)), arguments, asynchronous, diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index d7096cdd774..12410c187e2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -40,3 +40,4 @@ litellm-auth-gcp.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true rstest_reuse.workspace = true +wiremock = "0.6.5" diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index de926c715d5..2391ab83a60 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -85,3 +85,29 @@ pub(super) async fn outbound_request( other => other, }) } + +#[cfg(test)] +mod tests { + use super::{Error, as_response_error}; + + #[test] + fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { + for original in [ + Error::MissingField("usage"), + Error::Unsupported("non-text response content block"), + Error::InvalidRequest("whatever".to_string()), + Error::Auth(litellm_auth::Error::InvalidHeader), + ] { + let label = format!("{original:?}"); + assert!( + matches!(as_response_error(original), Error::InvalidResponse(_)), + "{label} must not stay retryable once the provider has answered" + ); + } + let upstream = Error::Transport(litellm_http::transport::Error::Http { + status: 500, + body: "boom".to_string(), + }); + assert_eq!(as_response_error(upstream.clone()), upstream); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index afea46221f5..b6425773964 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -736,248 +736,4 @@ mod tests { .unwrap_or_else(|error| panic!("prepare declined {messages}: {error}")); } } - - mod round_trip { - use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - net::{TcpListener, TcpStream}, - }; - - use super::*; - use crate::chat_completions::chat_completions; - - async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") - { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - } - String::from_utf8(request).expect("request is utf8") - } - - fn http_response(status: &str, body: &str) -> String { - format!( - "HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - body.len(), - body - ) - } - - /// Serve one request from a stub upstream and hand back what it received. - async fn serve_once( - status: &'static str, - body: &'static str, - ) -> (String, tokio::task::JoinHandle) { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let port = listener.local_addr().expect("addr").port(); - let handle = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts"); - let received = read_http_request(&mut socket).await; - socket - .write_all(http_response(status, body).as_bytes()) - .await - .expect("writes response"); - socket.flush().await.expect("flushes"); - received - }); - (format!("http://127.0.0.1:{port}/v1/messages"), handle) - } - - fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> { - ChatCompletionsRequest { - model: "anthropic/claude-sonnet-4-5", - messages, - optional_params: match params { - Value::Object(map) => map, - other => panic!("params must be an object, got {other}"), - }, - api_key: Some("sk-test"), - api_base: Some(api_base), - custom_llm_provider: None, - extra_headers: None, - timeout: Some(std::time::Duration::from_secs(10)), - } - } - - const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; - - #[tokio::test] - async fn round_trip_sends_the_translated_body_and_normalizes_the_response() { - let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await; - let response = chat_completions(call( - &api_base, - json!([ - {"role": "system", "content": "be terse"}, - {"role": "user", "content": "hi"} - ]), - json!({"max_tokens": 16}), - )) - .await - .expect("call succeeds"); - - let received = handle.await.expect("server task"); - let sent: Value = serde_json::from_str( - received - .split_once("\r\n\r\n") - .expect("request has a body") - .1, - ) - .expect("body is json"); - assert_eq!( - sent["messages"], - json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) - ); - assert_eq!( - sent["system"], - json!([{"type": "text", "text": "be terse"}]) - ); - assert_eq!(sent["max_tokens"], json!(16)); - assert!(received.to_lowercase().contains("x-api-key: sk-test")); - - assert_eq!( - response.choices[0].message.content.as_deref(), - Some("hello") - ); - assert_eq!(response.usage.total_tokens, 15); - } - - #[tokio::test] - async fn a_response_it_cannot_normalize_is_reported_as_already_sent() { - // The provider was called and billed, so the host must not retry this - // on its own path. `MissingField` here would read as a pre-send - // decline and be retried; `InvalidResponse` cannot. - const NO_USAGE: &str = - r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#; - let (api_base, handle) = serve_once("200 OK", NO_USAGE).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) - .await - .expect_err("response cannot be normalized"); - handle.await.expect("server task"); - assert!( - matches!(err, Error::InvalidResponse(_)), - "expected a post-send error, got {err:?}" - ); - } - - #[tokio::test] - async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() { - const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#; - let (api_base, handle) = serve_once("200 OK", TOOL_USE).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) - .await - .expect_err("response cannot be normalized"); - handle.await.expect("server task"); - assert!( - matches!(err, Error::InvalidResponse(_)), - "expected a post-send error, got {err:?}" - ); - } - - #[tokio::test] - async fn an_upstream_error_status_keeps_its_code() { - let (api_base, handle) = - serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) - .await - .expect_err("upstream rejects"); - handle.await.expect("server task"); - assert!( - matches!( - err, - Error::Transport(litellm_http::transport::Error::Http { status: 429, .. }) - ), - "expected a 429, got {err:?}" - ); - } - - #[tokio::test] - async fn a_connection_that_is_never_established_declines_instead_of_failing() { - // Nothing was sent, so nothing was billed and the host can still serve - // the request. Classing this with the post-send failures would turn a - // recoverable fallback into a user-facing error on exactly the - // deployments whose transport is configured only on the Python client. - let port = { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - listener.local_addr().expect("has an address").port() - // Dropped here, so the port is closed and the connect is refused. - }; - let err = chat_completions(call( - &format!("http://127.0.0.1:{port}/v1/messages"), - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) - .await - .expect_err("nothing is listening"); - assert!( - matches!( - err, - Error::Transport(litellm_http::transport::Error::Connect(_)) - ), - "expected a pre-send connect failure, got {err:?}" - ); - } - - #[test] - fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { - use crate::chat_completions::handler::as_response_error; - - for original in [ - Error::MissingField("usage"), - Error::Unsupported("non-text response content block"), - Error::InvalidRequest("whatever".to_string()), - Error::Auth(litellm_auth::Error::InvalidHeader), - ] { - let label = format!("{original:?}"); - assert!( - matches!(as_response_error(original), Error::InvalidResponse(_)), - "{label} must not stay retryable once the provider has answered" - ); - } - // An upstream status is already unambiguous, so it survives intact. - assert!(matches!( - as_response_error(Error::Transport(litellm_http::transport::Error::Http { - status: 500, - body: "boom".to_string() - })), - Error::Transport(litellm_http::transport::Error::Http { status: 500, .. }) - )); - } - } } diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index 4327754ed05..d27b79bdc04 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -29,151 +29,10 @@ pub(super) fn string_headers( #[cfg(test)] mod tests { - use std::{sync::Arc, time::Duration}; - - use futures_util::future::BoxFuture; - use litellm_secrets::{SecretValue, source::SecretSource}; - use serde_json::{Value, json}; - use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - net::{TcpListener, TcpStream}, - }; + use serde_json::json; use super::{messages_provider_config, string_headers, truncate_error_body}; - use crate::messages::{ - Error, - route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}, - types::MessagesShaping, - }; - - struct RecordingSecrets { - values: Vec<(&'static str, String)>, - requested: std::sync::Mutex>, - } - - impl SecretSource for RecordingSecrets { - fn get_secret_str<'a>( - &'a self, - name: &'a str, - ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { - Box::pin(async move { - self.requested.lock().unwrap().push(name.to_string()); - Ok(self - .values - .iter() - .find(|(key, _)| *key == name) - .map(|(_, value)| SecretValue::new(value.clone()))) - }) - } - } - - fn secrets_call() -> MessagesCall { - let Value::Object(body) = json!({ - "model": "claude-sonnet-4-5", - "max_tokens": 16, - "messages": [{"role": "user", "content": "hi"}] - }) else { - unreachable!("literal object") - }; - MessagesCall { - model: "claude-sonnet-4-5".into(), - body, - api_key: None, - api_base: None, - custom_llm_provider: Some("anthropic".into()), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - } - } - - async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - } - String::from_utf8(request).expect("request is utf8") - } - - #[tokio::test] - async fn route_reads_the_provider_credential_and_base_from_the_secret_source() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#; - let response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - let secrets = Arc::new(RecordingSecrets { - values: vec![ - ("ANTHROPIC_API_KEY", "sk-from-manager".to_string()), - ("ANTHROPIC_BASE_URL", format!("http://{addr}")), - ], - requested: std::sync::Mutex::new(Vec::new()), - }); - - let output = litellm_host::run::run( - messages_machine(secrets.clone()), - &LocalMessagesHost::new(secrets_call()), - ) - .await - .expect("messages request succeeds"); - - assert!(matches!(output, MessagesOutput::Message(_))); - let request = server.await.expect("server task completes"); - assert!( - request - .to_ascii_lowercase() - .contains("x-api-key: sk-from-manager"), - "{request}" - ); - let requested = secrets.requested.lock().unwrap().clone(); - assert_eq!( - requested, - messages_provider_config("anthropic") - .unwrap() - .secret_names() - .iter() - .map(ToString::to_string) - .collect::>() - ); - } + use crate::messages::Error; #[test] fn provider_config_resolves_anthropic_and_azure_ai() { diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index fe7e8bb4b80..de1a5f476ed 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -17,8 +17,10 @@ pub(super) async fn send( body: &Value, timeout: Option, ) -> Result { + let encoded = serde_json::to_vec(body) + .map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?; let builder = headers.iter().fold( - http_client().post(url).json(body), + http_client().post(url).body(encoded), |builder, (key, value)| builder.header(key, value), ); let builder = match timeout { diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 8cd3eaf3aa3..40aff185e81 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,16 +1,16 @@ use std::{ + convert::Infallible, sync::{Arc, Mutex}, time::Duration, }; use bytes::Bytes; use litellm_auth::SecretValue; -use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, - machine::{HostChannel, MachineFault, RouteMachine}, - route::Route, + machine::{CallMachine, HostChannel, MachineFault}, + protocol::Protocol, }; use litellm_secrets::source::SecretSource; use litellm_types::{ @@ -21,22 +21,12 @@ use serde_json::{Map, Value}; use super::{ Error, - common_utils::messages_provider_config, handler::{decode_response, network, provider_error, send}, prepare::{prepare_provider_request, resolve_provider}, types::{MessagesRequest, MessagesShaping}, }; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MessagesOp { - ProjectRequest, -} - -pub enum MessagesOpResult { - Request(Box), -} - /// The caller's request as the host projects it. pub struct MessagesCall { pub model: String, @@ -62,15 +52,20 @@ pub enum MessagesOutput { Streamed, } +/// The upstream response as the caller sees it at stream hand-off, before any chunk. +pub struct MessagesStreamHead { + pub headers: Vec<(String, String)>, +} + pub struct Messages; -impl Route for Messages { +impl Protocol for Messages { type Response = MessagesOutput; type Error = Error; - type Op = MessagesOp; - type OpResult = MessagesOpResult; + type Projection = MessagesCall; + type Op = Infallible; type Chunk = Bytes; - type StreamHead = (); + type StreamHead = MessagesStreamHead; } impl From for Error { @@ -78,26 +73,12 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "messages host driver was abandoned".into(), MachineFault::Protocol(message) => format!("messages {message}"), - MachineFault::Mismatch => "invalid messages host operation result".into(), }) } } pub type MessagesHost = HostChannel; -pub type MessagesMachine = RouteMachine; - -/// Whether this route serves the request, decided before any callback runs so a host -/// can still run its own path. -pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool { - let provider = get_custom_llm_provider(model, custom_llm_provider) - .map(|resolved| resolved.custom_llm_provider) - .or(custom_llm_provider); - match provider { - Some(ANTHROPIC_MESSAGES_PROVIDER) => true, - Some(provider) => !stream && messages_provider_config(provider).is_some(), - None => false, - } -} +pub type MessagesMachine = CallMachine; /// The in-process host for a request already in hand. It answers projection once and /// observes nothing. @@ -114,30 +95,28 @@ impl LocalMessagesHost { } impl Host for LocalMessagesHost { - async fn route(&self, op: MessagesOp) -> Result { - match op { - MessagesOp::ProjectRequest => self - .call - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|call| MessagesOpResult::Request(Box::new(call))) - .ok_or_else(|| { - Error::InvalidRequest("messages request was already projected".into()) - }), - } + async fn project(&self) -> Result { + self.call + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .ok_or_else(|| Error::InvalidRequest("messages request was already projected".into())) + } + + async fn custom_op(&self, op: Infallible) -> Result<(), Error> { + match op {} } } pub fn messages_machine(secrets: Arc) -> MessagesMachine { - RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) + CallMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) } async fn execute( host: MessagesHost, secrets: Arc, ) -> Result { - let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; + let call = host.project().await?; let stream = call.streams(); let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; let secrets = secrets.resolve(resolved.config.secret_names()).await?; @@ -163,8 +142,11 @@ async fn execute( model: request.model.clone(), custom_llm_provider: request.provider.clone(), optional_params: Value::Object( - call.body - .iter() + request + .body + .as_object() + .into_iter() + .flatten() .filter(|(name, _)| !matches!(name.as_str(), "model" | "messages")) .map(|(name, value)| (name.clone(), value.clone())) .collect(), @@ -204,7 +186,14 @@ async fn relay( host: &MessagesHost, mut response: reqwest::Response, ) -> Result { - if host.open(()).await? == Demand::Detached { + let head = MessagesStreamHead { + headers: response + .headers() + .iter() + .filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string()))) + .collect(), + }; + if host.open(head).await? == Demand::Detached { return Ok(MessagesOutput::Streamed); } while let Some(chunk) = response.chunk().await.map_err(network)? { diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index ffa4f045e8e..c33ee053422 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -23,9 +23,6 @@ pub fn prepare_document(input: OcrDocumentInput) -> Result { file_name.as_deref(), mime_type.as_deref(), )?), - OcrDocumentInput::HostReader { .. } => Err(Error::InvalidRequest( - "OCR file reader was not read by the host".into(), - )), } } @@ -207,7 +204,7 @@ mod tests { } #[test] - fn byte_documents_are_encoded_and_host_readers_must_be_read_first() { + fn byte_documents_are_encoded() { assert_eq!( prepare_document(OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), @@ -217,7 +214,6 @@ mod tests { .unwrap(), document("data:application/pdf;base64,YWJj") ); - assert!(prepare_document(OcrDocumentInput::HostReader { mime_type: None }).is_err()); } #[test] @@ -248,160 +244,3 @@ mod tests { } } } - -#[cfg(test)] -mod document_tests { - use litellm_host::event::WireRequest; - use litellm_llms::base_llm::ocr::error::Error; - use rstest::rstest; - use serde_json::{Value, json}; - - use crate::ocr::route::LocalOcrHost; - use crate::ocr::test_support::{ - MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, - request_body, wire_request_with_document, - }; - - #[derive(Clone, Copy, Debug)] - enum Route { - Mistral, - AzureAi, - VertexMistral, - AzureCohereParse, - Cohere, - } - - impl Route { - fn model(self) -> &'static str { - match self { - Self::Mistral => "mistral/model", - Self::AzureAi => "azure_ai/model", - Self::VertexMistral => "vertex_ai/mistral-ocr-maas", - Self::AzureCohereParse => "azure_ai/cohere-parse", - Self::Cohere => "cohere/model", - } - } - - fn document_type(self) -> &'static str { - match self { - Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url", - Self::AzureCohereParse | Self::Cohere => "image_url", - } - } - - fn options(self) -> Value { - match self { - Self::Mistral | Self::AzureAi => json!({"pages": [0]}), - Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}), - Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}), - } - } - } - - /// What the host does to the wire request in `before_send`. - #[derive(Clone, Copy, Debug)] - enum Host { - Detached, - ReplacesDocument, - } - - const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ="; - - impl Host { - fn before_send(self, wire: WireRequest) -> WireRequest { - let Value::Object(fields) = wire.body else { - return wire; - }; - let body = fields - .into_iter() - .map(|(name, value)| match self { - Self::Detached => (name, value), - Self::ReplacesDocument if name == "document" => { - let document_type = value["type"].clone(); - let key = document_type.as_str().unwrap_or_default().to_string(); - (name, json!({"type": document_type, key: REPLACED_DOCUMENT})) - } - Self::ReplacesDocument => (name, value), - }) - .collect(); - WireRequest { - body: Value::Object(body), - ..wire - } - } - } - - struct Sent { - result: Result<(), Error>, - provider_body: Option, - } - - async fn send(route: Route, host: Host, document_base: &str) -> Sent { - let (base, seen, provider) = - mock_server(vec![MockResponse::json(json!({"pages": []}))]).await; - let document_type = route.document_type(); - let document = - json!({"type": document_type, document_type: format!("{document_base}/scan.png")}); - let request = wire_request_with_document(route.model(), &base, document, route.options()); - let local = - LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire))); - let result = perform_ocr_with(local).await.map(|_| ()); - match result { - Ok(()) => provider.await.unwrap(), - Err(_) => provider.abort(), - } - let provider_body = seen - .lock() - .unwrap() - .first() - .map(|request| request_body(request)); - Sent { - result, - provider_body, - } - } - - fn served_document_uri() -> String { - use base64::Engine; - format!( - "data:image/png;base64,{}", - base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT) - ) - } - - #[rstest] - #[case::azure_ai(Route::AzureAi)] - #[case::vertex_mistral(Route::VertexMistral)] - #[case::azure_cohere_parse(Route::AzureCohereParse)] - #[tokio::test] - async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) { - let (document_base, _documents) = document_server().await; - let sent = send(route, Host::Detached, &document_base).await; - sent.result.unwrap(); - assert_eq!( - sent.provider_body.unwrap()["document"][route.document_type()], - json!(served_document_uri()) - ); - } - - #[rstest] - #[tokio::test] - async fn document_replaced_by_the_host_reaches_the_provider( - #[values( - Route::Mistral, - Route::AzureAi, - Route::VertexMistral, - Route::AzureCohereParse, - Route::Cohere - )] - route: Route, - ) { - let (document_base, _documents) = document_server().await; - let sent = send(route, Host::ReplacesDocument, &document_base).await; - sent.result.unwrap(); - assert_eq!( - sent.provider_body.unwrap()["document"][route.document_type()], - json!(REPLACED_DOCUMENT) - ); - } -} diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 270a402c9fa..2a0d20f69c9 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -7,212 +7,3 @@ pub mod provider_config; pub mod route; pub mod types; pub mod wire; - -#[cfg(test)] -pub(crate) mod test_support { - use std::sync::{Arc, Mutex}; - - use futures_util::future::BoxFuture; - use litellm_host::event::WireRequest; - use litellm_llms::base_llm::ocr::{ - error::Error, - handler::{CallHooks, OcrClient}, - transformation::LiteLLMOcrResponse, - }; - use serde_json::{Value, json}; - use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - net::TcpListener, - }; - - use crate::ocr::{ - route::{LocalOcrHost, ocr_machine}, - types::LiteLLMOcrRequest, - wire::{OcrWireRequest, decode_request}, - }; - - /// Stands in for a host with no hooks registered: the wire request goes out unchanged - /// and response events go nowhere. - pub(crate) struct NoHooks; - - impl CallHooks for NoHooks { - fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { - Box::pin(async move { Ok(wire) }) - } - - fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> { - Box::pin(async { Ok(()) }) - } - } - - pub(crate) fn ocr_client() -> OcrClient { - let document_http = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test document client builds"); - OcrClient::for_test(reqwest::Client::new(), document_http) - } - - pub(crate) async fn perform_ocr( - request: LiteLLMOcrRequest, - ) -> Result { - crate::ocr::client::perform(&ocr_client(), request).await - } - - pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result { - litellm_host::run::run(ocr_machine(ocr_client()), &host).await - } - - pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest { - wire_request_with_document( - model, - base, - json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), - options, - ) - } - - pub(crate) fn wire_request_with_document( - model: &str, - base: &str, - document: Value, - options: Value, - ) -> LiteLLMOcrRequest { - decode_request(OcrWireRequest { - model: model.into(), - document, - api_key: Some(litellm_auth::SecretValue::new("test-key")), - api_base: Some(base.into()), - custom_llm_provider: None, - extra_headers: None, - optional_params: options.as_object().unwrap().clone(), - input_sources: Default::default(), - timeout_seconds: Some(2.0), - }) - .unwrap() - } - - pub(crate) fn resolved_request( - request: LiteLLMOcrRequest, - ) -> crate::ocr::types::ResolvedOcrRequest { - request - .map_document(crate::ocr::document::prepare_document) - .unwrap() - } - - pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest { - let request = resolved_request(request); - let document = request.document.clone().with_source(source.into()); - request.with_document(document.into()) - } - - pub(crate) fn request_body(request: &str) -> Value { - serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() - } - - pub(crate) const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document"; - - /// Serves [`SERVED_DOCUMENT`] as `image/png` to every connection until aborted. - pub(crate) async fn document_server() -> (String, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let base = format!("http://{}", listener.local_addr().unwrap()); - let task = tokio::spawn(async move { - loop { - let (mut socket, _) = listener.accept().await.unwrap(); - let mut buffer = [0u8; 4096]; - let _ = socket.read(&mut buffer).await.unwrap(); - let head = format!( - "HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", - SERVED_DOCUMENT.len() - ); - socket.write_all(head.as_bytes()).await.unwrap(); - socket.write_all(SERVED_DOCUMENT).await.unwrap(); - } - }); - (base, task) - } - - pub(crate) struct MockResponse { - pub status: u16, - pub headers: Vec<(&'static str, String)>, - pub body: Value, - } - - impl MockResponse { - pub fn json(body: Value) -> Self { - Self { - status: 200, - headers: vec![], - body, - } - } - } - - pub(crate) async fn mock_server( - responses: Vec, - ) -> (String, Arc>>, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let base = format!("http://{}", listener.local_addr().unwrap()); - let requests = Arc::new(Mutex::new(Vec::new())); - let seen = requests.clone(); - let server_base = base.clone(); - let task = tokio::spawn(async move { - for response in responses { - let (mut socket, _) = listener.accept().await.unwrap(); - let mut bytes = Vec::new(); - let mut buffer = [0u8; 4096]; - let header_end = loop { - let n = socket.read(&mut buffer).await.unwrap(); - assert!(n > 0); - bytes.extend_from_slice(&buffer[..n]); - if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") { - break index + 4; - } - }; - let length = String::from_utf8_lossy(&bytes[..header_end]) - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().unwrap()) - }) - .unwrap_or(0); - while bytes.len() < header_end + length { - let n = socket.read(&mut buffer).await.unwrap(); - assert!(n > 0); - bytes.extend_from_slice(&buffer[..n]); - } - seen.lock() - .unwrap() - .push(String::from_utf8_lossy(&bytes).into_owned()); - let body = serde_json::to_vec(&response.body).unwrap(); - let headers = response - .headers - .into_iter() - .map(|(name, value)| { - format!("{name}: {}\r\n", value.replace("{base}", &server_base)) - }) - .collect::(); - let head = format!( - "HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n", - response.status, - body.len(), - headers - ); - socket.write_all(head.as_bytes()).await.unwrap(); - socket.write_all(&body).await.unwrap(); - } - }); - (base, requests, task) - } - - pub(crate) fn header<'a>(request: &'a str, name: &str) -> Option<&'a str> { - request - .lines() - .take_while(|line| !line.is_empty()) - .find_map(|line| { - let (key, value) = line.split_once(':')?; - key.eq_ignore_ascii_case(name).then(|| value.trim()) - }) - } -} diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 37c0f18f659..18961ec96fa 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -70,20 +70,203 @@ pub(crate) fn prepare_request( } } -#[cfg(test)] -pub(crate) fn prepare_request_for_test(request: ResolvedOcrRequest) -> PreparedOcrRequest { - prepare_request( - request, - true, - &OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()), - std::sync::Arc::new(litellm_core_utils::settings::ProcessEnvironment), - ) -} - #[cfg(test)] mod tests { + use std::time::Duration; + + use futures_util::future::BoxFuture; use litellm_core_utils::call_arguments::{CallArguments, compose_body, parse_options}; - use serde_json::json; + use litellm_host::event::WireRequest; + use litellm_llms::{ + base_llm::ocr::{ + error::Error, + handler::{CallHooks, OcrClient}, + transformation::{BaseOcrConfig, OcrResponseFormat}, + }, + cohere::ocr::transformation::CohereParseConfig, + mistral::ocr::transformation::MistralOcrConfig, + vertex_ai::ocr::transformation::VertexAiOcrConfig, + }; + use serde_json::{Value, json}; + + use super::*; + use crate::ocr::{ + document::prepare_document, + types::LiteLLMOcrRequest, + wire::{OcrWireRequest, decode_request}, + }; + + /// Stands in for a host with no hooks registered. + struct NoHooks; + + impl CallHooks for NoHooks { + fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { + Box::pin(async move { Ok(wire) }) + } + + fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> { + Box::pin(async { Ok(()) }) + } + } + + fn client() -> OcrClient { + OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()) + } + + fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest { + decode_request(OcrWireRequest { + model: model.into(), + document, + api_key: Some(litellm_auth::SecretValue::new("test-key")), + api_base: Some(base.into()), + custom_llm_provider: None, + extra_headers: None, + optional_params: options.as_object().unwrap().clone(), + input_sources: Default::default(), + timeout_seconds: Some(2.0), + }) + .unwrap() + } + + fn prepared(request: LiteLLMOcrRequest) -> PreparedOcrRequest { + prepare_request( + request.map_document(prepare_document).unwrap(), + true, + &client(), + std::sync::Arc::new(litellm_core_utils::settings::ProcessEnvironment), + ) + } + + fn image(url: &str) -> Value { + json!({"type": "image_url", "image_url": url}) + } + + #[tokio::test] + async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() { + let request = request( + "cohere/parse", + "https://example.com", + image("https://example.com/original.png"), + json!({ + "output_format": "markdown", "timeout": 30, + "extra_body": { + "output_format": {"future": true}, + "document": {"type": "image_url", "image_url": "https://example.com/a.png", + "provider_options": {"nested": [false, 0, null]}} + } + }), + ); + + let http = CohereParseConfig + .prepare_request(&prepared(request), &client(), &NoHooks) + .await + .unwrap(); + + let body: Value = serde_json::from_slice(http.body()).unwrap(); + assert_eq!( + body, + json!({ + "model": "parse", "output_format": {"future": true}, + "document": {"type": "image_url", "image_url": "https://example.com/a.png", + "provider_options": {"nested": [false, 0, null]}} + }) + ); + } + + #[tokio::test] + async fn explicit_null_options_use_defaults_before_http() { + let request = request( + "cohere/parse", + "https://example.com", + image("https://example.com/a.png"), + json!({"output_format": null, "req_format": null}), + ); + assert_eq!( + request.response_format().unwrap(), + OcrResponseFormat::Litellm + ); + + let http = CohereParseConfig + .prepare_request(&prepared(request), &client(), &NoHooks) + .await + .unwrap(); + + let body: Value = serde_json::from_slice(http.body()).unwrap(); + assert_eq!(body["output_format"], "markdown"); + assert!(body.get("req_format").is_none()); + } + + #[tokio::test] + async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() { + let options = json!({ + "pages": [0, 2], + "include_image_base64": true, + "vertex_project": "project-1", + "vertex_location": "us-central1", + "unknown": "preserved" + }); + let document = + json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}); + let direct = prepared(request( + "mistral/mistral-ocr-maas", + "https://mistral.test", + document.clone(), + options.clone(), + )); + let vertex = prepared(request( + "vertex_ai/mistral-ocr-maas", + "https://vertex.test", + document, + options, + )); + + let direct_http = MistralOcrConfig + .prepare_request(&direct, &client(), &NoHooks) + .await + .unwrap(); + let vertex_http = VertexAiOcrConfig + .prepare_request(&vertex, &client(), &NoHooks) + .await + .unwrap(); + + assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr"); + assert_eq!( + vertex_http.url(), + "https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict" + ); + for http in [&direct_http, &vertex_http] { + assert_eq!(http.header("authorization").unwrap(), "Bearer test-key"); + assert_eq!(http.header("content-type").unwrap(), "application/json"); + assert_eq!(http.timeout(), Some(Duration::from_secs(2))); + let body: Value = serde_json::from_slice(http.body()).unwrap(); + assert_eq!( + body, + json!({ + "model": "mistral-ocr-maas", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "pages": [0, 2], + "include_image_base64": true, + "unknown": "preserved" + }) + ); + } + let payload = serde_json::to_vec( + &json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}), + ) + .unwrap(); + let direct_response = MistralOcrConfig + .transform_ocr_response(&direct.model, &payload, OcrResponseFormat::Litellm) + .unwrap() + .into_json(); + let vertex_response = VertexAiOcrConfig + .transform_ocr_response(&vertex.model, &payload, OcrResponseFormat::Litellm) + .unwrap() + .into_json(); + assert_eq!(direct_response, vertex_response); + assert_eq!(direct_response["model"], "mistral-ocr-maas"); + assert_eq!(direct_response["object"], "ocr"); + assert_eq!(direct_response["extra"], "preserved"); + } #[derive(serde::Deserialize)] struct KnownParams { diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index 7f83291bdab..adb704a15d1 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -3,102 +3,73 @@ use std::sync::{Arc, Mutex}; use litellm_auth::ResolvedCredential; use litellm_host::{ event::{CallEvent, RequestContext, WireRequest}, - machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute}, - route::Route, + host::Reply, + machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol}, + protocol::Protocol, }; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, }; use super::handler::perform_ocr_request; -use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest}; +use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest}; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum OcrOp { - ProjectRequest, - ReadDocument, - AcquireAzureAdToken, + AcquireAzureAdToken(Reply), } -pub enum OcrOpResult { - Request { - request: Box>, - caller_token: bool, - }, - Document(OcrFileContent), - AzureAdToken(ResolvedCredential), +/// The caller's request as the host projects it. +pub struct OcrProjection { + pub request: LiteLLMOcrRequest, + /// The caller passed its own Azure AD token provider, which the host keeps. + pub caller_token: bool, } pub struct Ocr; -impl Route for Ocr { +impl Protocol for Ocr { type Response = LiteLLMOcrResponse; type Error = Error; + type Projection = OcrProjection; type Op = OcrOp; - type OpResult = OcrOpResult; type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } -impl TokenRoute for Ocr { - fn acquire_token_op() -> OcrOp { - OcrOp::AcquireAzureAdToken - } - - fn token_credential(result: OcrOpResult) -> Option { - match result { - OcrOpResult::AzureAdToken(credential) => Some(credential), - _ => None, - } +impl TokenProtocol for Ocr { + fn acquire_token_op(reply: Reply) -> OcrOp { + OcrOp::AcquireAzureAdToken(reply) } } pub type OcrHost = HostChannel; -pub type OcrMachine = RouteMachine; +pub type OcrMachine = CallMachine; -/// The OCR call as a machine: projection, document reading and token acquisition are -/// host operations; everything else runs in Rust. +/// The OCR call as a machine: projection and token acquisition are host operations; +/// everything else runs in Rust. pub fn ocr_machine(client: OcrClient) -> OcrMachine { - RouteMachine::new(move |host| Box::pin(execute(client, host))) + CallMachine::new(move |host| Box::pin(execute(client, host))) } async fn execute(client: OcrClient, host: OcrHost) -> Result { - let OcrOpResult::Request { + let OcrProjection { request, caller_token, - } = host.route(OcrOp::ProjectRequest).await? - else { - return Err(MachineFault::Mismatch.into()); - }; + } = host.project().await?; let request = LiteLLMOcrRequest { azure_ad_token_provider: caller_token .then(|| HostTokenProvider::handle(host.clone())) .or(request.azure_ad_token_provider), - ..*request + ..request }; let caller_document = matches!(request.document, OcrDocumentInput::Document(_)); - let request = prepare_request_document(request, &host).await?; + let request = prepare_request_document(request).await?; perform_ocr_request(&client, request, &host, caller_document).await } async fn prepare_request_document( request: LiteLLMOcrRequest, - host: &OcrHost, ) -> Result { - let request = match &request.document { - OcrDocumentInput::HostReader { mime_type } => { - let mime_type = mime_type.clone(); - let OcrOpResult::Document(content) = host.route(OcrOp::ReadDocument).await? else { - return Err(MachineFault::Mismatch.into()); - }; - request.with_document(OcrDocumentInput::Bytes { - bytes: content.bytes, - file_name: content.file_name, - mime_type, - }) - } - _ => request, - }; if let OcrDocumentInput::Document(_) = &request.document { return request.map_document(super::document::prepare_document); } @@ -107,7 +78,6 @@ async fn prepare_request_document( .map_err(|error| Error::DocumentTask(Arc::new(error)))? } -type Reader = Box Result + Send + Sync>; type BeforeSend = Box Result + Send + Sync>; type Observer = Box; @@ -116,7 +86,6 @@ type Observer = Box; /// projection, and the optional observer sees and may rewrite the wire request. pub struct LocalOcrHost { request: Mutex>>, - reader: Option, before_send: Option, observer: Option, } @@ -125,22 +94,11 @@ impl LocalOcrHost { pub fn new(request: LiteLLMOcrRequest) -> Self { Self { request: Mutex::new(Some(request)), - reader: None, before_send: None, observer: None, } } - pub fn with_reader( - self, - reader: impl Fn() -> Result + Send + Sync + 'static, - ) -> Self { - Self { - reader: Some(Box::new(reader)), - ..self - } - } - pub fn with_before_send( self, before_send: impl Fn(WireRequest, &RequestContext) -> Result @@ -163,25 +121,21 @@ impl LocalOcrHost { } impl litellm_host::host::Host for LocalOcrHost { - async fn route(&self, op: OcrOp) -> Result { + async fn project(&self) -> Result { + self.request + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|request| OcrProjection { + request, + caller_token: false, + }) + .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { match op { - OcrOp::ProjectRequest => self - .request - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|request| OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }) - .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())), - OcrOp::ReadDocument => self - .reader - .as_ref() - .ok_or_else(|| Error::InvalidRequest("OCR host has no document reader".into())) - .and_then(|reader| reader()) - .map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => { + OcrOp::AcquireAzureAdToken(_) => { Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition( "OCR host has no Azure AD token provider".into(), ))) @@ -207,3589 +161,3 @@ impl litellm_host::host::Host for LocalOcrHost { Ok(()) } } - -#[cfg(test)] -mod aws_textract_tests { - use std::{collections::BTreeMap, time::SystemTime}; - - use litellm_auth_aws::{Credentials, aws_signature_headers, sign_post}; - use litellm_llms::base_llm::ocr::error::Error; - use serde_json::{Value, json}; - use time::{PrimitiveDateTime, format_description}; - - use crate::ocr::{ - route::LocalOcrHost, - test_support::{ - MockResponse, header, mock_server, perform_ocr_with, request_body, - wire_request_with_document, - }, - types::LiteLLMOcrRequest, - }; - - const ACCESS_KEY_ID: &str = "AKIDEXAMPLE"; - const SECRET_ACCESS_KEY: &str = "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"; - - fn textract_request(base: &str) -> LiteLLMOcrRequest { - textract_request_for("aws_textract/detect-document-text", base) - } - - fn textract_request_for(model: &str, base: &str) -> LiteLLMOcrRequest { - wire_request_with_document( - model, - &format!("{base}/"), - json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}), - json!({ - "aws_access_key_id": ACCESS_KEY_ID, - "aws_secret_access_key": SECRET_ACCESS_KEY, - "aws_region_name": "eu-west-1" - }), - ) - } - - fn textract_response() -> MockResponse { - MockResponse::json(json!({ - "DocumentMetadata": {"Pages": 1}, - "Blocks": [{"BlockType": "PAGE"}, {"BlockType": "LINE", "Text": "Invoice 12345"}] - })) - } - - /// Recomputes SigV4 over the bytes the server received, at the time the client claimed. - fn expected_authorization(url: &str, raw_request: &str) -> String { - let format = - format_description::parse_borrowed::<2>("[year][month][day]T[hour][minute][second]Z") - .unwrap(); - let signed_at: SystemTime = - PrimitiveDateTime::parse(header(raw_request, "x-amz-date").unwrap(), &format) - .unwrap() - .assume_utc() - .into(); - let headers: BTreeMap = ["content-type", "x-amz-target"] - .into_iter() - .map(|name| { - ( - name.to_string(), - header(raw_request, name).unwrap().to_string(), - ) - }) - .collect(); - let body = raw_request.split_once("\r\n\r\n").unwrap().1; - sign_post( - url, - body.as_bytes(), - &aws_signature_headers(&headers), - "eu-west-1", - "textract", - &Credentials::new(ACCESS_KEY_ID, SECRET_ACCESS_KEY, None, None, "test"), - signed_at, - ) - .unwrap()["Authorization"] - .clone() - } - - #[tokio::test] - async fn the_request_is_signed_for_textract_and_lines_become_the_page() { - let (base, seen, server) = mock_server(vec![textract_response()]).await; - - let response = perform_ocr_with(LocalOcrHost::new(textract_request(&base))) - .await - .unwrap(); - server.await.unwrap(); - - let raw = seen.lock().unwrap()[0].clone(); - assert_eq!( - header(&raw, "x-amz-target"), - Some("Textract.DetectDocumentText") - ); - assert_eq!( - header(&raw, "content-type"), - Some("application/x-amz-json-1.1") - ); - assert_eq!( - request_body(&raw), - json!({"Document": {"Bytes": "b3JpZ2luYWw="}}) - ); - assert_eq!( - header(&raw, "authorization"), - Some(expected_authorization(&format!("{base}/"), &raw).as_str()) - ); - assert_eq!(response.pages[0].markdown, "Invoice 12345"); - assert_eq!(response.usage_info.unwrap().pages_processed, Some(1)); - } - - #[tokio::test] - async fn a_body_rewritten_by_before_send_is_what_gets_signed_and_sent() { - let (base, seen, server) = mock_server(vec![textract_response()]).await; - let host = LocalOcrHost::new(textract_request(&base)).with_before_send(|mut wire, _| { - assert!( - !wire - .headers - .iter() - .any(|(name, _)| name.eq_ignore_ascii_case("authorization")), - "the hook ran after signing" - ); - wire.body["Document"]["Bytes"] = Value::from("cmVkYWN0ZWQ="); - Ok(wire) - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - - let raw = seen.lock().unwrap()[0].clone(); - assert_eq!( - request_body(&raw), - json!({"Document": {"Bytes": "cmVkYWN0ZWQ="}}) - ); - assert_eq!( - header(&raw, "authorization"), - Some(expected_authorization(&format!("{base}/"), &raw).as_str()) - ); - } - - #[tokio::test] - async fn a_multi_page_rejection_reaches_the_caller_with_the_single_page_limit() { - let (base, _, server) = mock_server(vec![MockResponse { - status: 400, - headers: vec![], - body: json!({ - "__type": "UnsupportedDocumentException", - "Message": "Request has unsupported document format" - }), - }]) - .await; - - let error = perform_ocr_with(LocalOcrHost::new(textract_request(&base))) - .await - .unwrap_err(); - server.await.unwrap(); - - let Error::Provider { status, body, .. } = error else { - panic!("expected a provider error, got {error:?}"); - }; - assert_eq!(status, 400); - assert!( - body.contains("multi-page documents are not supported"), - "{body}" - ); - } - - #[tokio::test] - async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "DocumentMetadata": {"Pages": 1}, - "Blocks": [ - {"Id": "l1", "BlockType": "LINE", "Text": "Quarterly Report"}, - {"Id": "t", "BlockType": "LAYOUT_TITLE", - "Relationships": [{"Type": "CHILD", "Ids": ["l1"]}]} - ] - }))]) - .await; - let request = textract_request_for("aws_textract/analyze-document", &base); - - let response = perform_ocr_with(LocalOcrHost::new(request)).await.unwrap(); - server.await.unwrap(); - - let raw = seen.lock().unwrap()[0].clone(); - assert_eq!( - header(&raw, "x-amz-target"), - Some("Textract.AnalyzeDocument") - ); - assert_eq!( - request_body(&raw)["FeatureTypes"], - json!(["LAYOUT", "TABLES"]) - ); - assert_eq!( - header(&raw, "authorization"), - Some(expected_authorization(&format!("{base}/"), &raw).as_str()) - ); - assert_eq!(response.pages[0].markdown, "# Quarterly Report"); - } -} - -#[cfg(test)] -mod azure_ai_tests { - use litellm_llms::base_llm::ocr::error::Error; - use serde_json::{Value, json}; - - use crate::ocr::route::LocalOcrHost; - use crate::ocr::test_support::{ - MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request, - }; - - #[tokio::test] - async fn facade_executes_azure_mistral_with_prepared_auth() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"hello"}], - "usage_info":{"pages_processed":1} - }))]) - .await; - let mut request = wire_request( - "azure_ai/model", - &base, - json!({"include_image_base64":true}), - ); - request.credentials.api_key = None; - request.transport.extra_headers = vec![( - "Authorization".into(), - "Bearer python-prepared-token".into(), - )]; - - let result = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(result.pages[0].markdown, "hello"); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr ")); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer python-prepared-token\r\n") - ); - let body: Value = - serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!( - body, - json!({ - "model":"model", - "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, - "include_image_base64":true - }) - ); - } - - #[tokio::test] - async fn facade_acquires_supplied_entra_token_for_final_request() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let mut request = wire_request( - "azure_ai/model", - &base, - json!({"azure_ad_token":"rust-owned-token"}), - ); - request.credentials.api_key = None; - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer rust-owned-token\r\n") - ); - } - - #[tokio::test] - async fn rejects_non_inline_body_after_guardrails() { - let request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({})); - let host = LocalOcrHost::new(request).with_before_send(|mut wire, _| { - wire.body["document"] = json!({ - "type":"document_url", - "document_url":"https://example.com/not-inline.pdf" - }); - Ok(wire) - }); - let error = perform_ocr_with(host).await.unwrap_err(); - assert!(error.to_string().contains("data URI")); - } - - mod transformation { - use std::sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, - }; - - use litellm_auth::{ - ResolvedCredential, SecretValue, TokenFuture, TokenProvider, TokenProviderHandle, - }; - use rstest::rstest; - use serde_json::json; - - use super::*; - use crate::ocr::{ - test_support::{MockResponse, header, mock_server, perform_ocr}, - types::LiteLLMOcrRequest, - wire::decode_request, - }; - - #[derive(Debug)] - struct CountingToken { - token: fn(usize) -> String, - calls: AtomicUsize, - } - - impl CountingToken { - fn new(token: fn(usize) -> String) -> Arc { - Arc::new(Self { - token, - calls: AtomicUsize::new(0), - }) - } - - fn calls(&self) -> usize { - self.calls.load(Ordering::SeqCst) - } - } - - impl TokenProvider for CountingToken { - fn acquire(&self) -> TokenFuture<'_> { - let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1; - let token = SecretValue::new((self.token)(call)); - Box::pin(async move { - Ok(ResolvedCredential::AccessToken { - token, - expires_on: None, - }) - }) - } - } - - fn numbered_token(call: usize) -> String { - format!("callback-{call}") - } - - fn azure_request( - provider: &Arc, - api_base: Option<&str>, - api_key: Option<&str>, - extra_headers: Value, - optional_params: Value, - ) -> LiteLLMOcrRequest { - let wire = serde_json::from_value(json!({ - "model": "azure_ai/mistral-ocr-latest", - "document": {"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": null, - "extra_headers": extra_headers, - "optional_params": optional_params, - "timeout_seconds": 2.0 - })) - .unwrap(); - LiteLLMOcrRequest { - azure_ad_token_provider: Some(TokenProviderHandle::new(provider.clone())), - ..decode_request(wire).unwrap() - } - } - - fn ocr_page() -> MockResponse { - MockResponse::json(json!({"pages":[{"index":0,"markdown":"hello"}]})) - } - - #[tokio::test] - async fn token_provider_result_is_the_bearer_and_is_acquired_for_each_request() { - let provider = CountingToken::new(numbered_token); - let (base, seen, server) = mock_server(vec![ocr_page(), ocr_page()]).await; - - for _ in 0..2 { - perform_ocr(azure_request( - &provider, - Some(&base), - None, - Value::Null, - json!({}), - )) - .await - .unwrap(); - } - server.await.unwrap(); - - assert_eq!(provider.calls(), 2); - let requests = seen.lock().unwrap(); - assert_eq!( - requests - .iter() - .map(|request| header(request, "authorization")) - .collect::>(), - [Some("Bearer callback-1"), Some("Bearer callback-2")] - ); - } - - #[rstest] - #[case::api_key_skips_provider(Some("resource-key"), Value::Null, json!({}), "Bearer resource-key", 0)] - #[case::provider_beats_static_token( - None, - Value::Null, - json!({"azure_ad_token":"static-token"}), - "Bearer callback-1", - 1 - )] - #[case::header_wins_on_the_wire_but_provider_still_runs( - None, - json!({"Authorization":"Bearer override"}), - json!({}), - "Bearer override", - 1 - )] - #[tokio::test] - async fn credential_precedence( - #[case] api_key: Option<&str>, - #[case] extra_headers: Value, - #[case] optional_params: Value, - #[case] expected_authorization: &str, - #[case] expected_calls: usize, - ) { - let provider = CountingToken::new(numbered_token); - let (base, seen, server) = mock_server(vec![ocr_page()]).await; - - perform_ocr(azure_request( - &provider, - Some(&base), - api_key, - extra_headers, - optional_params, - )) - .await - .unwrap(); - server.await.unwrap(); - - assert_eq!(provider.calls(), expected_calls); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert_eq!( - header(&requests[0], "authorization"), - Some(expected_authorization) - ); - } - - #[rstest] - #[case::missing_api_base( - false, - json!({}), - numbered_token, - |error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase { - provider: "Azure AI", - environment_variable: "AZURE_AI_API_BASE", - })), - 0 - )] - #[case::unsupported_oidc_reference( - true, - json!({"azure_ad_token":"oidc/assertion","client_id":"client","tenant_id":"tenant"}), - numbered_token, - |error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)), - 0 - )] - #[case::empty_provider_token_ignores_static_token( - true, - json!({"azure_ad_token":"static-token"}), - |_| String::new(), - |error: &Error| matches!(error, Error::MissingAzureAiCredentials), - 1 - )] - #[tokio::test] - async fn credential_failures_send_no_provider_request( - #[case] with_api_base: bool, - #[case] optional_params: Value, - #[case] token: fn(usize) -> String, - #[case] expected: fn(&Error) -> bool, - #[case] expected_calls: usize, - ) { - let provider = CountingToken::new(token); - let (base, seen, server) = mock_server(vec![ocr_page()]).await; - - let error = perform_ocr(azure_request( - &provider, - with_api_base.then_some(base.as_str()), - None, - Value::Null, - optional_params, - )) - .await - .unwrap_err(); - server.abort(); - - assert!(expected(&error), "unexpected error: {error:?}"); - assert_eq!(provider.calls(), expected_calls); - assert!(seen.lock().unwrap().is_empty()); - } - } -} - -#[cfg(test)] -mod azure_document_intelligence_tests { - use litellm_host::event::{CallEvent, MachineEvent}; - use litellm_llms::base_llm::ocr::{error::Error, settings::OcrSettings}; - use rstest::rstest; - use serde_json::{Value, json}; - - use crate::ocr::route::LocalOcrHost; - use crate::ocr::{ - test_support::{ - MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request, - }, - wire::{OcrWireRequest, decode_request}, - }; - - fn query_value(url: &str, key: &str) -> Option { - url::Url::parse(url) - .unwrap() - .query_pairs() - .find_map(|(name, value)| (name == key).then(|| value.into_owned())) - } - - #[tokio::test] - async fn facade_maps_pages_features_and_url_document() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded", - "analyzeResult":{"pages":[]} - }))]) - .await; - let mut request = wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}), - ); - request.document = serde_json::from_value::< - litellm_llms::base_llm::ocr::transformation::OcrDocument, - >(json!({ - "type":"document_url", - "document_url":"https://example.com/document.pdf" - })) - .unwrap() - .into(); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let request = &seen.lock().unwrap()[0]; - let target = request.split_whitespace().nth(1).unwrap(); - let url = format!("{base}{target}"); - assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3")); - assert_eq!( - query_value(&url, "features").as_deref(), - Some("keyValuePairs,languages") - ); - let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!( - body, - json!({"urlSource":"https://example.com/document.pdf"}) - ); - } - - #[rstest] - #[case(json!({"pages":[true]}), Error::Pages("expected only integers or only strings".into()))] - #[case(json!({"pages":[1,"2"]}), Error::Pages("expected only integers or only strings".into()))] - #[case(json!({"pages":[-1]}), Error::Pages("negative page index".into()))] - #[case(json!({"pages":"1&&features=bad"}), Error::Pages("invalid native page range".into()))] - #[case(json!({"features":"languages&pages=1"}), Error::Features)] - #[case(json!({"req_format":"azure"}), Error::RequestFormat)] - #[tokio::test] - async fn rejects_invalid_pages_features_and_format( - #[case] options: Value, - #[case] expected: Error, - ) { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await; - let result = decode_request(OcrWireRequest { - model: "azure_ai/doc-intelligence/prebuilt-read".into(), - document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), - api_key: Some(litellm_auth::SecretValue::new("key")), - api_base: Some(base), - custom_llm_provider: None, - extra_headers: None, - optional_params: options.as_object().unwrap().clone(), - input_sources: Default::default(), - timeout_seconds: Some(2.0), - }); - let result = match result { - Ok(request) => perform_ocr(request).await, - Err(error) => Err(error), - }; - server.abort(); - let _ = server.await; - assert!( - seen.lock().unwrap().is_empty(), - "sent invalid options: {options}" - ); - let error = result.unwrap_err(); - assert_eq!( - std::mem::discriminant(&error), - std::mem::discriminant(&expected) - ); - assert_eq!(error.http_status_code(), Some(400)); - assert_eq!(error.to_string(), expected.to_string()); - } - - #[rstest] - #[case(json!({}))] - #[case(json!({"req_format":"litellm"}))] - #[tokio::test] - async fn missing_native_fields_keep_page_text_without_retaining_raw_response( - #[case] options: Value, - ) { - let operation = json!({ - "status":"succeeded", - "analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"hello"}]}]} - }); - let (base, seen, server) = mock_server(vec![MockResponse::json(operation)]).await; - let response = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - options, - )) - .await - .unwrap(); - server.await.unwrap(); - - assert_eq!(response.pages.len(), 1); - assert_eq!(response.pages[0].index, 0); - assert_eq!(response.pages[0].markdown, "hello"); - assert_eq!(response.provider_native_response, None); - let serialized = response.into_json(); - assert_eq!(serialized.get("content"), Some(&Value::Null)); - assert_eq!(serialized.get("tables"), Some(&Value::Null)); - assert_eq!(serialized.get("keyValuePairs"), Some(&Value::Null)); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - let target = requests[0].split_whitespace().nth(1).unwrap(); - let url = format!("{base}{target}"); - for field in ["pages", "features", "req_format"] { - assert_eq!(query_value(&url, field), None); - } - let body: Value = - serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!(body, json!({"base64Source":"YWJj"})); - } - - #[tokio::test] - async fn inline_document_decodes_to_base64_source() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded" - }))]) - .await; - let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let request = &seen.lock().unwrap()[0]; - let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!(body, json!({"base64Source":"YWJj"})); - } - - #[tokio::test] - async fn immediate_response_normalizes_pages_and_preserves_native() { - let operation = json!({ - "status":"succeeded", - "operationExtension":42, - "analyzeResult":{ - "content":"A\n\nB", - "tables":[{"cells":[]}], - "keyValuePairs":[{"key":{"content":"A"}}], - "pages":[{ - "pageNumber":"2", - "width":"8.5", - "height":11, - "unit":"inch", - "lines":[{"content":"A"},{"content":null},{"content":"B"}] - }] - } - }); - let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await; - let result = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"req_format":"native"}), - )) - .await - .unwrap(); - server.await.unwrap(); - - assert_eq!(result.pages[0].index, 1); - assert_eq!(result.pages[0].markdown, "A\n\nB"); - assert_eq!( - serde_json::to_value(&result.pages[0].dimensions).unwrap(), - json!({"width":816,"height":1056,"dpi":96}) - ); - assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1)); - let serialized = result.clone().into_json(); - assert_eq!(serialized["content"], "A\n\nB"); - assert_eq!(serialized["tables"], json!([{"cells":[]}])); - assert_eq!( - serialized["keyValuePairs"], - json!([{"key":{"content":"A"}}]) - ); - assert!(serialized.get("key_value_pairs").is_none()); - assert_eq!( - result.provider_native_response.map(Value::Object), - Some(operation) - ); - } - - #[tokio::test] - async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded", - "analyzeResult":{"pages":[{"pageNumber":1,"width":8.5,"height":11,"unit":"inch"}]} - }))]) - .await; - let client = ocr_client().with_settings(OcrSettings { - document_intelligence_api_version: "2099-01-01".into(), - document_intelligence_dpi: 72, - ..OcrSettings::default() - }); - - let result = crate::ocr::client::perform( - &client, - wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})), - ) - .await - .unwrap(); - server.await.unwrap(); - - let target = seen.lock().unwrap()[0] - .split_whitespace() - .nth(1) - .unwrap() - .to_string(); - assert_eq!( - query_value(&format!("{base}{target}"), "api-version").as_deref(), - Some("2099-01-01") - ); - assert_eq!( - serde_json::to_value(&result.pages[0].dimensions).unwrap(), - json!({"width":612,"height":792,"dpi":72}) - ); - } - - #[tokio::test] - async fn accepted_response_polls_to_success_with_only_credentials() { - let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}}); - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse { - status: 200, - headers: vec![("Retry-After", "0".into())], - body: json!({"status":"running"}), - }, - MockResponse::json(operation.clone()), - ]) - .await; - let mut request = wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"req_format":"native"}), - ); - request - .transport - .extra_headers - .push(("X-Trace".into(), "initial-only".into())); - - let result = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!( - result.provider_native_response.map(Value::Object), - Some(operation) - ); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 3); - assert!(requests[0].to_ascii_lowercase().contains("x-trace:")); - for poll in &requests[1..] { - assert!(!poll.to_ascii_lowercase().contains("x-trace:")); - assert!( - poll.to_ascii_lowercase() - .contains("ocp-apim-subscription-key: test-key") - ); - } - } - - #[tokio::test] - async fn accepted_response_emits_response_received_before_polling() { - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({"submitted": true}), - }, - MockResponse::json(json!({"status":"succeeded"})), - ]) - .await; - let request_count = seen.clone(); - let host = LocalOcrHost::new(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .with_observer(move |event| { - let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event else { - return; - }; - match request_count.lock().unwrap().len() { - 1 => assert_eq!(raw.body, r#"{"submitted":true}"#), - 2 => assert!(raw.body.contains("succeeded")), - count => panic!("unexpected callback after {count} requests"), - } - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 2); - } - - #[tokio::test] - async fn polling_forwards_bearer_credentials() { - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse::json(json!({"status":"succeeded"})), - ]) - .await; - let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); - request.credentials.api_key = None; - request.transport.extra_headers = vec![("Authorization".into(), "Bearer token".into())]; - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let requests = seen.lock().unwrap(); - assert!( - requests[1] - .to_ascii_lowercase() - .contains("authorization: bearer token") - ); - } - - #[tokio::test] - async fn polling_does_not_follow_redirects() { - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse { - status: 302, - headers: vec![("Location", "{base}/redirected".into())], - body: json!({}), - }, - MockResponse::json(json!({"status":"succeeded"})), - ]) - .await; - - let error = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .await - .unwrap_err(); - - assert!(error.to_string().contains("status 302"), "{error}"); - assert_eq!(seen.lock().unwrap().len(), 2); - server.abort(); - } - - #[tokio::test] - async fn polling_rejects_terminal_failure() { - let (base, _, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse::json(json!({"status":"failed"})), - ]) - .await; - - let error = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .await - .unwrap_err(); - server.await.unwrap(); - assert!(error.to_string().contains("status failed")); - } - - #[tokio::test] - async fn malformed_provider_pages_report_response_paths() { - for (analysis, path) in [ - (json!({"pages":null}), "pages"), - (json!({"pages":[null]}), "pages[0]"), - (json!({"pages":[{"lines":null}]}), "lines"), - (json!({"pages":[{"width":"bad"}]}), "width"), - ] { - let (base, _, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded", - "analyzeResult":analysis - }))]) - .await; - let error = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .await - .unwrap_err(); - server.await.unwrap(); - assert!(error.to_string().contains(path), "{error}"); - } - } - - #[tokio::test] - async fn rejects_missing_invalid_and_cross_origin_operation_locations() { - for headers in [ - Vec::new(), - vec![("Operation-Location", "/relative".into())], - vec![("Operation-Location", "http://example.com/operation".into())], - vec![( - "Operation-Location", - "http://user:password@127.0.0.1/operation".into(), - )], - ] { - let (base, _, server) = mock_server(vec![MockResponse { - status: 202, - headers, - body: json!({}), - }]) - .await; - let error = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .await - .unwrap_err(); - server.await.unwrap(); - assert!(error.to_string().contains("operation-location")); - } - } - - #[tokio::test] - async fn polling_deadline_bounds_retry_delay() { - let (base, _, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse { - status: 200, - headers: vec![("Retry-After", "9999".into())], - body: json!({"status":"notStarted"}), - }, - ]) - .await; - let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); - let client = ocr_client().with_settings(OcrSettings { - poll_timeout: std::time::Duration::from_millis(100), - ..OcrSettings::default() - }); - - let error = tokio::time::timeout( - std::time::Duration::from_secs(1), - crate::ocr::client::perform(&client, request), - ) - .await - .unwrap() - .unwrap_err(); - server.await.unwrap(); - assert!(error.to_string().contains("timed out")); - } - - #[tokio::test] - async fn model_id_is_encoded_and_dot_segments_are_rejected() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded" - }))]) - .await; - perform_ocr(wire_request( - "azure_ai/doc-intelligence/a ?#é", - &base, - json!({}), - )) - .await - .unwrap(); - server.await.unwrap(); - assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze")); - - for model in [ - "azure_ai/doc-intelligence/.", - "azure_ai/doc-intelligence/..", - ] { - let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({}))) - .await - .unwrap_err(); - assert!(error.to_string().contains("dot segment")); - } - } - - mod transformation { - use std::sync::{Arc, Mutex}; - - use litellm_host::event::{CallEvent, MachineEvent}; - use litellm_llms::base_llm::ocr::transformation::OcrDocument; - use serde_json::{Value, json}; - - use super::*; - use crate::ocr::{ - route::LocalOcrHost, - test_support::{ - MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request, - }, - }; - - #[tokio::test] - async fn facade_maps_pages_features_and_url_document() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "status":"succeeded", - "analyzeResult":{"pages":[]} - }))]) - .await; - let mut request = wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}), - ); - request.document = serde_json::from_value::(json!({ - "type":"document_url", - "document_url":"https://example.com/document.pdf" - })) - .unwrap() - .into(); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let request = &seen.lock().unwrap()[0]; - let target = request.split_whitespace().nth(1).unwrap(); - let url = format!("{base}{target}"); - assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3")); - assert_eq!( - query_value(&url, "features").as_deref(), - Some("keyValuePairs,languages") - ); - let body: Value = - serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!( - body, - json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false}) - ); - } - - #[tokio::test] - async fn rejects_invalid_pages_features_and_format() { - for options in [ - json!({"pages":[true]}), - json!({"pages":[1,"2"]}), - json!({"pages":[-1]}), - json!({"pages":"1&&features=bad"}), - json!({"features":"languages&pages=1"}), - json!({"req_format":"azure"}), - ] { - let request = wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - "http://127.0.0.1:1", - options.clone(), - ); - let rejected = perform_ocr(request).await.is_err(); - assert!(rejected, "accepted {options}"); - } - } - - #[tokio::test] - async fn immediate_response_normalizes_pages_and_preserves_native() { - let operation = json!({ - "status":"succeeded", - "operationExtension":42, - "analyzeResult":{ - "content":"A\n\nB", - "tables":[{"cells":[]}], - "keyValuePairs":[{"key":{"content":"A"}}], - "pages":[{ - "pageNumber":"2", - "width":"8.5", - "height":11, - "unit":"inch", - "lines":[{"content":"A"},{"content":null},{"content":"B"}] - }] - } - }); - let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await; - let result = perform_ocr(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"req_format":"native"}), - )) - .await - .unwrap(); - server.await.unwrap(); - - assert_eq!(result.pages[0].index, 1); - assert_eq!(result.pages[0].markdown, "A\n\nB"); - assert_eq!( - serde_json::to_value(&result.pages[0].dimensions).unwrap(), - json!({"width":816,"height":1056,"dpi":96}) - ); - assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1)); - let serialized = result.clone().into_json(); - assert_eq!(serialized["content"], "A\n\nB"); - assert_eq!(serialized["tables"], json!([{"cells":[]}])); - assert_eq!( - serialized["keyValuePairs"], - json!([{"key":{"content":"A"}}]) - ); - assert!(serialized.get("key_value_pairs").is_none()); - assert_eq!( - result.provider_native_response.as_ref(), - operation.as_object() - ); - } - - #[tokio::test] - async fn accepted_response_polls_to_success_with_only_credentials() { - let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}}); - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({}), - }, - MockResponse { - status: 200, - headers: vec![("Retry-After", "0".into())], - body: json!({"status":"running"}), - }, - MockResponse::json(operation.clone()), - ]) - .await; - let mut request = wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({"req_format":"native"}), - ); - request - .transport - .extra_headers - .push(("X-Trace".into(), "initial-only".into())); - - let result = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!( - result.provider_native_response.as_ref(), - operation.as_object() - ); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 3); - assert!(requests[0].to_ascii_lowercase().contains("x-trace:")); - for poll in &requests[1..] { - assert!(!poll.to_ascii_lowercase().contains("x-trace:")); - assert!( - poll.to_ascii_lowercase() - .contains("ocp-apim-subscription-key: test-key") - ); - } - } - - #[tokio::test] - async fn accepted_response_emits_response_received_for_submission_and_completed_poll() { - let (base, seen, server) = mock_server(vec![ - MockResponse { - status: 202, - headers: vec![("Operation-Location", "{base}/operation".into())], - body: json!({"submitted": true}), - }, - MockResponse::json(json!({"status":"succeeded"})), - ]) - .await; - let responses_received = Arc::new(Mutex::new(Vec::new())); - let request_count = seen.clone(); - let observed = responses_received.clone(); - let host = LocalOcrHost::new(wire_request( - "azure_ai/doc-intelligence/prebuilt-read", - &base, - json!({}), - )) - .with_observer(move |event| { - if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { - observed - .lock() - .unwrap() - .push((request_count.lock().unwrap().len(), raw.body.clone())); - } - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 2); - assert_eq!( - *responses_received.lock().unwrap(), - [ - (1, r#"{"submitted":true}"#.to_string()), - (2, r#"{"status":"succeeded"}"#.to_string()), - ] - ); - } - } -} - -#[cfg(test)] -mod cohere_tests { - mod transformation { - use litellm_llms::{ - base_llm::ocr::{ - error::Error, - transformation::{BaseOcrConfig, OcrDocument, OcrResponseFormat}, - }, - cohere::ocr::transformation::*, - }; - use rstest::rstest; - use serde_json::{Value, json}; - - #[tokio::test] - async fn composed_body_preserves_native_document_fields_and_untyped_overrides() { - let request = crate::ocr::test_support::wire_request( - "cohere/parse", - "https://example.com", - json!({ - "output_format":"markdown", "timeout":30, - "extra_body":{ - "output_format": {"future":true}, - "document":{"type":"image_url","image_url":"https://example.com/a.png", - "provider_options":{"nested":[false,0,null]}} - } - }), - ); - let request = request.with_document( - serde_json::from_value(json!({ - "type":"image_url","image_url":"https://example.com/original.png" - })) - .unwrap(), - ); - let request = crate::ocr::prepare::prepare_request_for_test(request); - let http = CohereParseConfig - .prepare_request( - &request, - &crate::ocr::test_support::ocr_client(), - &crate::ocr::test_support::NoHooks, - ) - .await - .unwrap(); - let body: Value = serde_json::from_slice(http.body()).unwrap(); - assert_eq!( - body, - json!({ - "model":"parse", "output_format":{"future":true}, - "document":{"type":"image_url","image_url":"https://example.com/a.png", - "provider_options":{"nested":[false,0,null]}} - }) - ); - } - - #[tokio::test] - async fn explicit_null_options_use_defaults_before_http() { - let request = crate::ocr::test_support::wire_request( - "cohere/parse", - "https://example.com", - json!({"output_format":null,"req_format":null}), - ); - let request = request.with_document( - serde_json::from_value( - json!({"type":"image_url","image_url":"https://example.com/a.png"}), - ) - .unwrap(), - ); - assert_eq!( - request.response_format().unwrap(), - OcrResponseFormat::Litellm - ); - let request = crate::ocr::prepare::prepare_request_for_test(request); - let http = CohereParseConfig - .prepare_request( - &request, - &crate::ocr::test_support::ocr_client(), - &crate::ocr::test_support::NoHooks, - ) - .await - .unwrap(); - let body: Value = serde_json::from_slice(http.body()).unwrap(); - assert_eq!(body["output_format"], "markdown"); - assert!(body.get("req_format").is_none()); - } - - #[rstest] - #[case::cohere("cohere/parse-v5.0", "POST /v2/parse ")] - #[case::azure_ai("azure_ai/Cohere-parse-v5.0", "POST /providers/cohere/v2/parse ")] - #[tokio::test] - async fn route_sends_image_to_its_parse_endpoint_with_the_bearer_key( - #[case] model: &str, - #[case] request_line: &str, - ) { - use crate::ocr::test_support::{MockResponse, header, mock_server, perform_ocr}; - - let (base, seen, server) = - mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let request = crate::ocr::test_support::wire_request(model, &base, json!({})) - .with_document( - serde_json::from_value::( - json!({"type":"image_url","image_url":"data:image/png;base64,YWJj"}), - ) - .unwrap() - .into(), - ); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with(request_line), "{}", requests[0]); - assert_eq!( - header(&requests[0], "authorization"), - Some("Bearer test-key") - ); - } - - #[rstest] - #[tokio::test] - async fn route_rejects_non_image_document_without_a_request( - #[values("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")] model: &str, - ) { - use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr}; - - let (base, seen, server) = - mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - - let error = perform_ocr(crate::ocr::test_support::wire_request( - model, - &base, - json!({}), - )) - .await - .unwrap_err(); - server.abort(); - - assert!(matches!(error, Error::CohereImageOnly), "{error:?}"); - assert!(seen.lock().unwrap().is_empty()); - } - } -} - -#[cfg(test)] -mod deepseek_tests { - use litellm_llms::{ - base_llm::ocr::transformation::{BaseOcrConfig, OcrDocument}, - vertex_ai::ocr::deepseek_transformation::{ - DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, - normalize_response as transform_ocr_response, - }, - }; - use rstest::rstest; - use serde_json::{Value, json}; - - fn document() -> OcrDocument { - serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap() - } - - #[rstest] - #[case("stream", json!(true))] - #[case("temperature", json!(0.1))] - #[case("max_tokens", json!(1024))] - #[case("top_p", json!(0.9))] - #[case("n", json!(2))] - #[case("stop", json!("done"))] - #[case("stop", json!(["done", "stop"]))] - fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) { - let params: DeepSeekOcrParams = - serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap(); - let result = serde_json::to_value( - VertexAIDeepSeekOCRConfig - .transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), ¶ms, &[]) - .unwrap(), - ) - .unwrap(); - assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas"); - assert_eq!( - result["messages"][0]["content"][0], - json!({"type":"image_url","image_url":"gs://bucket/a.png"}) - ); - assert_eq!(result[name], value); - assert!(result.get("ignored").is_none()); - } - - #[rstest] - #[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))] - #[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))] - fn request_maps_both_document_types_to_image_content(#[case] document: Value) { - let source = document - .get("image_url") - .or_else(|| document.get("document_url")) - .unwrap() - .clone(); - let request = VertexAIDeepSeekOCRConfig - .transform_ocr_request( - "deepseek-ai/deepseek-ocr-maas", - serde_json::from_value(document).unwrap(), - &DeepSeekOcrParams::default(), - &[], - ) - .unwrap(); - let result = serde_json::to_value(request).unwrap(); - assert_eq!( - result["messages"][0]["content"][0], - json!({"type":"image_url","image_url":source}) - ); - } - - #[rstest] - #[case(json!("# hello"), "# hello")] - #[case(json!("{broken"), "{broken")] - #[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")] - #[case(json!({"pages":[]}), "")] - #[case(json!("[]"), "[]")] - #[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")] - #[case(json!({"pages":[{"markdown":"object"}]}), "object")] - fn response_codec_handles_text_json_and_objects( - #[case] content: Value, - #[case] expected: &str, - ) { - let structured = content - .as_object() - .is_some_and(|object| object.contains_key("pages")) - || content - .as_str() - .is_some_and(|text| text.contains("\"pages\"")); - let response: DeepSeekOcrResponse = serde_json::from_value( - json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}), - ) - .unwrap(); - let result = transform_ocr_response("model", response) - .unwrap() - .into_json(); - assert_eq!(result["pages"][0]["markdown"], expected); - assert_eq!(result["pages"][0]["index"], 0); - if structured { - assert!(result["usage_info"].is_null()); - } else { - assert_eq!(result["usage_info"]["prompt_tokens"], 1); - } - } - - #[test] - fn structured_result_maps_pages_usage_model_and_annotation() { - let response: DeepSeekOcrResponse = serde_json::from_value(json!({ - "choices":[{"message":{"content":{ - "pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}], - "model":"provider-model", - "usage_info":{"pages_processed":1}, - "document_annotation":{"language":"en"}, - "future":"kept" - }}}] - })) - .unwrap(); - let result = transform_ocr_response("requested", response) - .unwrap() - .into_json(); - assert_eq!(result["pages"][0]["index"], 2); - assert_eq!(result["pages"][0]["images"][0]["id"], "one"); - assert_eq!(result["model"], "provider-model"); - assert_eq!(result["usage_info"]["pages_processed"], 1); - assert_eq!(result["document_annotation"]["language"], "en"); - assert_eq!(result["future"], "kept"); - } - - #[test] - fn response_codec_rejects_missing_empty_and_malformed_content() { - for value in [ - json!({"choices":[{"message":{"content":{}}}]}), - json!({"choices":[]}), - json!({"choices":[{"message":{"content":""}}]}), - json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}), - json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}), - ] { - let result = serde_json::from_value::(value) - .map_err(|_| ()) - .and_then(|response| transform_ocr_response("model", response).map_err(|_| ())); - assert!(result.is_err()); - } - } -} - -#[cfg(test)] -mod reducto_tests { - use litellm_host::event::{CallEvent, MachineEvent, WireRequest}; - use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument}; - use rstest::rstest; - use serde_json::{Value, json}; - - use crate::ocr::route::LocalOcrHost; - use crate::ocr::test_support::{ - MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request, - }; - - fn request_body(request: &str) -> Value { - serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() - } - - #[rstest] - #[case( - "reducto/parse-v3", - json!({ - "formatting":{"table_output_format":"html"}, - "retrieval":{"chunk_mode":"section"}, - "settings":{"ocr_system":"standard"}, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - "reducto://already.pdf", - json!({ - "input":"reducto://already.pdf", - "formatting":{"table_output_format":"html"}, - "retrieval":{"chunk_mode":"section"}, - "settings":{"ocr_system":"standard"}, - "future_ocr_option":true, - "provider_option":"value" - }) - )] - #[case( - "reducto/parse-legacy", - json!({ - "enhance":{"agentic":[{"type":"table"}]}, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - "reducto://legacy.pdf", - json!({ - "document_url":"reducto://legacy.pdf", - "options":{"enhance":{"agentic":[{"type":"table"}]}}, - "future_ocr_option":true, - "provider_option":"value" - }) - )] - #[tokio::test] - async fn request_mapping_matches_python( - #[case] model: &str, - #[case] options: Value, - #[case] source: &str, - #[case] expected: Value, - ) { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "result":{"chunks":[]} - }))]) - .await; - let request = - crate::ocr::test_support::with_source(wire_request(model, &base, options), source); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with("POST /parse ")); - assert_eq!(request_body(&requests[0]), expected); - } - - #[rstest] - #[case("parse-v3")] - #[case("parse-legacy")] - #[tokio::test] - async fn data_uri_upload_preserves_multipart_headers( - #[case] model: &str, - #[values("application/pdf", "image/png")] mime_type: &str, - ) { - let (base, seen, server) = mock_server(vec![ - MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), - MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})), - ]) - .await; - let document = if mime_type.starts_with("image/") { - json!({"type":"image_url","image_url":format!("data:{mime_type};base64,YWJj")}) - } else { - json!({"type":"document_url","document_url":format!("data:{mime_type};base64,YWJj")}) - }; - let mut request = crate::ocr::types::LiteLLMOcrRequest { - document: serde_json::from_value::(document) - .unwrap() - .into(), - ..wire_request(&format!("reducto/{model}"), &base, json!({})) - }; - request.transport.extra_headers = vec![ - ("Content-Type".into(), "application/json".into()), - ("X-Trace".into(), "upload-test".into()), - ]; - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.pages[0].markdown, "hello"); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 2); - assert!(requests[0].starts_with("POST /upload ")); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("content-type: multipart/form-data; boundary=") - ); - assert!(requests[0].contains("x-trace: upload-test")); - let multipart = requests[0].split_once("\r\n\r\n").unwrap().1; - assert!(multipart.contains(&format!("Content-Type: {mime_type}\r\n"))); - assert!(multipart.contains("\r\n\r\nabc\r\n--")); - assert!(requests[1].starts_with("POST /parse ")); - let source_field = if model == "parse-legacy" { - "document_url" - } else { - "input" - }; - assert_eq!( - request_body(&requests[1]), - json!({source_field:"reducto://uploaded.pdf"}) - ); - for request in requests.iter() { - assert!( - request - .to_ascii_lowercase() - .contains("authorization: bearer test-key\r\n") - ); - } - } - - #[tokio::test] - async fn response_received_stays_after_reducto_upload_and_parse() { - let (base, seen, server) = mock_server(vec![ - MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), - MockResponse::json(json!({"result":{"chunks":[]}})), - ]) - .await; - let request_count = seen.clone(); - let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))) - .with_observer(move |event| { - if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { - assert_eq!(request_count.lock().unwrap().len(), 2); - assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#); - } - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 2); - } - - #[rstest] - #[case(json!({"file_id":""}))] - #[case(json!({}))] - #[case(json!({"file_id":null}))] - #[tokio::test] - async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) { - let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await; - let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({}))) - .await - .unwrap_err(); - server.await.unwrap(); - assert!(error.to_string().contains("file_id")); - assert_eq!(seen.lock().unwrap().len(), 1); - } - - #[tokio::test] - async fn upload_failure_stops_before_parse() { - let (base, seen, server) = mock_server(vec![MockResponse { - status: 503, - headers: vec![], - body: json!({"error":"unavailable"}), - }]) - .await; - assert!( - perform_ocr(wire_request("reducto/parse-v3", &base, json!({}))) - .await - .is_err() - ); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 1); - } - - #[rstest] - #[case("https://example.com/a.pdf", Error::ReductoSource)] - #[case("reducto://", Error::RequestField { path: "document file id".into() })] - #[case("data:application/pdf;base64", Error::InvalidDataUri)] - #[case("data:application/pdf;base64,INVALID!", Error::InvalidDataUri)] - #[tokio::test] - async fn rejects_invalid_document_sources_before_network( - #[case] source: &str, - #[case] expected: Error, - ) { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await; - let request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({})), - source, - ); - let result = perform_ocr(request).await; - server.abort(); - let _ = server.await; - assert!( - seen.lock().unwrap().is_empty(), - "sent invalid source: {source}" - ); - let error = result.unwrap_err(); - assert_eq!( - std::mem::discriminant(&error), - std::mem::discriminant(&expected) - ); - assert_eq!(error.http_status_code(), Some(400)); - assert_eq!(error.to_string(), expected.to_string()); - } - - #[test] - fn response_normalization_groups_blocks_and_distinguishes_null_result() { - use litellm_llms::reducto::ocr::transformation::{ - ReductoResponse, normalize_response as transform_ocr_response, - }; - - let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[ - {"blocks":[{ - "type":"Table", - "content":"B", - "bbox":{"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4}, - "confidence":"high", - "granular_confidence":{"parse_confidence":0.95,"extract_confidence":null}, - "image_url":null - }]}, - {"blocks":[{"content":"A","bbox":{"page":1},"type":"Text"},{"content":"C","bbox":{"page":1}}]} - ]}}); - let response: ReductoResponse = serde_json::from_value(raw).unwrap(); - let normalized = transform_ocr_response("parse-v3", response) - .unwrap() - .into_json(); - assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC"); - assert_eq!(normalized["pages"][1]["markdown"], "B"); - assert_eq!(normalized["pages"][1]["blocks"][0]["type"], "Table"); - assert_eq!( - normalized["pages"][1]["blocks"][0]["bbox"], - json!({"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4}) - ); - assert_eq!(normalized["pages"][1]["blocks"][0]["confidence"], "high"); - assert_eq!( - normalized["pages"][1]["blocks"][0]["granular_confidence"]["parse_confidence"], - 0.95 - ); - assert!(normalized["pages"][1]["blocks"][0]["image_url"].is_null()); - assert_eq!(normalized["usage_info"]["pages_processed"], 2); - assert_eq!(normalized["usage_info"]["credits"], 3.0); - - let missing: ReductoResponse = - serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap(); - let missing = transform_ocr_response("parse-v3", missing).unwrap(); - assert_eq!(missing.pages[0].markdown, "text"); - let null: ReductoResponse = serde_json::from_value( - json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}), - ) - .unwrap(); - let null = transform_ocr_response("parse-v3", null).unwrap(); - assert!(null.pages.is_empty()); - } - - #[tokio::test] - async fn facade_omits_native_response_by_default_and_preserves_auth_priority() { - let raw = json!({"job_id":"job-1","result":{"chunks":[]}}); - let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; - let mut request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({})), - "reducto://ready.pdf", - ); - request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.provider_native_response, None); - assert!( - seen.lock().unwrap()[0] - .to_ascii_lowercase() - .contains("authorization: bearer existing") - ); - } - - #[tokio::test] - async fn native_format_retains_the_provider_response() { - let raw = json!({ - "result":{"chunks":[{"content":"native OCR response"}]}, - "usage":{"num_pages":1} - }); - let (base, _, server) = mock_server(vec![MockResponse::json(raw.clone())]).await; - let request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({"req_format":"native"})), - "reducto://ready.pdf", - ); - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - - assert_eq!(response.pages[0].markdown, "native OCR response"); - assert_eq!(response.provider_native_response.as_ref(), raw.as_object()); - } - - #[tokio::test] - async fn unknown_model_reaches_parse_and_keeps_its_name() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "result":{"chunks":[{"content":"future model response"}]} - }))]) - .await; - let request = crate::ocr::test_support::with_source( - wire_request("reducto/future-parse-model", &base, json!({})), - "reducto://ready.pdf", - ); - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - - assert_eq!(response.model, "future-parse-model"); - assert_eq!(response.pages[0].markdown, "future model response"); - let requests = seen.lock().unwrap(); - assert!(requests[0].starts_with("POST /parse ")); - assert_eq!( - request_body(&requests[0]), - json!({"input":"reducto://ready.pdf"}) - ); - } - - #[tokio::test] - async fn guardrail_rewrites_document_before_upload() { - let (base, seen, server) = - mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await; - let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))) - .with_before_send(|wire, _| { - assert_eq!( - wire.body["document_url"], - "data:application/pdf;base64,YWJj" - ); - Ok(WireRequest { - body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}), - ..wire - }) - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with("POST /parse ")); - assert!(requests[0].contains("reducto://guarded.pdf")); - } - - mod transformation { - use litellm_host::event::{CallEvent, MachineEvent, WireRequest}; - use litellm_llms::{ - base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext}, - reducto::ocr::transformation::*, - }; - use rstest::rstest; - - use super::*; - use crate::ocr::{ - route::LocalOcrHost, - test_support::{ - MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request, - }, - }; - - #[tokio::test] - async fn v3_options_preserve_explicit_null() { - let overrides = - serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true})) - .unwrap(); - let params = ReductoParseV3Config - .map_ocr_params(&overrides, "parse-v3") - .unwrap(); - let client = crate::ocr::test_support::ocr_client(); - let connection = OcrConnection::default(); - let document = serde_json::from_value( - json!({"type":"document_url","document_url":"reducto://ready.pdf"}), - ) - .unwrap(); - let body = ReductoParseV3Config - .async_transform_ocr_request( - "parse-v3", - document, - ¶ms, - &[], - OcrRequestContext { - client: &client, - connection: &connection, - }, - ) - .await - .unwrap(); - assert_eq!( - serde_json::to_value(body).unwrap(), - json!({ - "input":"reducto://ready.pdf", "formatting":null, "settings":{} - }) - ); - let absent = ReductoParseV3Config - .map_ocr_params( - &litellm_core_utils::call_arguments::CallArguments::default(), - "parse-v3", - ) - .unwrap(); - assert_eq!(serde_json::to_value(absent).unwrap(), json!({})); - } - - #[rstest] - #[case( - "reducto/parse-v3", - json!({ - "formatting":{"table_output_format":"html"}, - "retrieval":{"chunk_mode":"section"}, - "settings":{"ocr_system":"standard"}, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - "reducto://already.pdf", - json!({ - "input":"reducto://already.pdf", - "formatting":{"table_output_format":"html"}, - "retrieval":{"chunk_mode":"section"}, - "settings":{"ocr_system":"standard"}, - "future_ocr_option":true, - "provider_option":"value" - }) - )] - #[case( - "reducto/parse-legacy", - json!({ - "enhance":{"agentic":[{"type":"table"}]}, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - "reducto://legacy.pdf", - json!({ - "document_url":"reducto://legacy.pdf", - "options":{"enhance":{"agentic":[{"type":"table"}]}}, - "future_ocr_option":true, - "provider_option":"value" - }) - )] - #[tokio::test] - async fn request_mapping_matches_python( - #[case] model: &str, - #[case] options: Value, - #[case] source: &str, - #[case] expected: Value, - ) { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "result":{"chunks":[]} - }))]) - .await; - let request = - crate::ocr::test_support::with_source(wire_request(model, &base, options), source); - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with("POST /parse ")); - assert_eq!(request_body(&requests[0]), expected); - } - - #[rstest] - #[case("parse-v3")] - #[case("parse-legacy")] - #[tokio::test] - async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) { - let (base, seen, server) = mock_server(vec![ - MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), - MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})), - ]) - .await; - let mut request = wire_request(&format!("reducto/{model}"), &base, json!({})); - request.transport.extra_headers = vec![ - ("Content-Type".into(), "application/json".into()), - ("X-Trace".into(), "upload-test".into()), - ]; - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.pages[0].markdown, "hello"); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 2); - assert!(requests[0].starts_with("POST /upload ")); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("content-type: multipart/form-data; boundary=") - ); - assert!(requests[0].contains("x-trace: upload-test")); - assert!(requests[0].contains("application/pdf")); - assert!(requests[0].contains("abc")); - assert!(requests[1].starts_with("POST /parse ")); - } - - #[tokio::test] - async fn response_received_stays_after_reducto_upload_and_parse() { - let (base, seen, server) = mock_server(vec![ - MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), - MockResponse::json(json!({"result":{"chunks":[]}})), - ]) - .await; - let request_count = seen.clone(); - let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))) - .with_observer(move |event| { - if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { - assert_eq!(request_count.lock().unwrap().len(), 2); - assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#); - } - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 2); - } - - #[rstest] - #[case("https://example.com/a.pdf")] - #[case("reducto://")] - #[case("data:application/pdf;base64")] - #[case("data:application/pdf;base64,INVALID!")] - #[tokio::test] - async fn rejects_invalid_document_sources_before_network(#[case] source: &str) { - let request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), - source, - ); - assert!(perform_ocr(request).await.is_err()); - } - - #[tokio::test] - async fn facade_omits_native_response_by_default_and_preserves_auth_priority() { - let raw = json!({"job_id":"job-1","result":{"chunks":[]}}); - let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; - let mut request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({})), - "reducto://ready.pdf", - ); - request.transport.extra_headers = - vec![("authorization".into(), "Bearer existing".into())]; - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.provider_native_response, None); - assert!( - seen.lock().unwrap()[0] - .to_ascii_lowercase() - .contains("authorization: bearer existing") - ); - } - - #[rstest] - #[case("reducto/parse-v3")] - #[case("reducto/parse-legacy")] - #[tokio::test] - async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) { - let (base, seen, server) = mock_server(vec![ - MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), - MockResponse::json(json!({"result":{"chunks":[]}})), - ]) - .await; - let mut request = wire_request(model, &base, json!({})); - request.transport.extra_headers = - vec![("authorization".into(), "Bearer original".into())]; - let host = LocalOcrHost::new(request).with_before_send(|wire, _| { - Ok(WireRequest { - headers: vec![("authorization".into(), "Bearer guarded".into())], - ..wire - }) - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 2); - assert!(requests[0].starts_with("POST /upload ")); - assert!(requests[1].starts_with("POST /parse ")); - for request in requests.iter() { - assert!(request.contains("authorization: Bearer guarded")); - assert!(!request.contains("Bearer original")); - } - } - } -} - -#[cfg(test)] -mod vertex_ai_tests { - use litellm_auth::InputSource; - use litellm_llms::base_llm::ocr::{settings::OcrSettings, transformation::OcrResponseFormat}; - use serde_json::{Value, json}; - - use crate::ocr::test_support::{ - MockResponse, mock_server, ocr_client, perform_ocr, wire_request, - }; - - fn request_body(request: &str) -> Value { - serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() - } - - #[tokio::test] - async fn facade_executes_vertex_mistral_with_resolved_project_and_location() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"hello"}], - "usage_info":{"pages_processed":1} - }))]) - .await; - let request = wire_request( - "vertex_ai/mistral-ocr-maas", - &base, - json!({ - "vertex_project":"project-1", - "vertex_location":"europe-west4", - "extract_footer":true - }), - ); - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.pages[0].markdown, "hello"); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with( - "POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict " - )); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer test-key") - ); - assert_eq!( - request_body(&requests[0]), - json!({ - "model":"mistral-ocr-maas", - "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, - "extract_footer":true - }) - ); - } - - #[tokio::test] - async fn configured_project_and_location_apply_when_the_call_sets_neither() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let client = ocr_client().with_settings(OcrSettings { - vertex_project: Some("configured-project".into()), - vertex_location: Some("europe-west4".into()), - ..OcrSettings::default() - }); - - crate::ocr::client::perform( - &client, - wire_request("vertex_ai/mistral-ocr-maas", &base, json!({})), - ) - .await - .unwrap(); - server.await.unwrap(); - assert!(seen.lock().unwrap()[0].starts_with( - "POST /v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict " - )); - } - - #[tokio::test] - async fn supplied_authorization_is_forwarded_without_a_static_token() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let mut request = wire_request( - "vertex_ai/model", - &base, - json!({"vertex_project":"project-1"}), - ); - request.credentials.api_key = None; - request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())]; - - perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert!( - seen.lock().unwrap()[0] - .to_ascii_lowercase() - .contains("authorization: bearer supplied") - ); - } - - #[tokio::test] - async fn invalid_credentials_fail_before_provider_http() { - let request = wire_request( - "vertex_ai/model", - "http://127.0.0.1:1", - json!({"vertex_credentials": true}), - ); - let error = perform_ocr(request).await.unwrap_err(); - assert!(error.to_string().contains("vertex_credentials")); - } - - #[tokio::test] - async fn request_controlled_api_base_is_rejected_before_vertex_auth() { - let mut request = wire_request( - "vertex_ai/mistral-ocr-maas", - "https://caller.example", - json!({"vertex_project":"project-1"}), - ); - request.credentials.api_base = Some(litellm_auth::Sourced::new( - "https://caller.example".into(), - InputSource::Request, - )); - - let error = perform_ocr(request).await.unwrap_err(); - assert!( - error - .to_string() - .contains("request-controlled Vertex AI endpoint") - ); - } - - #[tokio::test] - async fn adapters_build_complete_requests_and_share_mistral_normalization() { - use std::time::Duration; - - use litellm_llms::{ - base_llm::ocr::transformation::BaseOcrConfig, - mistral::ocr::transformation::MistralOcrConfig, - vertex_ai::ocr::transformation::VertexAiOcrConfig, - }; - - use crate::ocr::test_support::ocr_client; - - let client = ocr_client(); - let options = json!({ - "pages": [0, 2], - "include_image_base64": true, - "vertex_project": "project-1", - "vertex_location": "us-central1", - "unknown": "ignored" - }); - let direct = wire_request( - "mistral/mistral-ocr-maas", - "https://mistral.test", - options.clone(), - ); - let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); - let direct = crate::ocr::prepare::prepare_request_for_test( - crate::ocr::test_support::resolved_request(direct), - ); - let vertex = crate::ocr::prepare::prepare_request_for_test( - crate::ocr::test_support::resolved_request(vertex), - ); - let direct_http = MistralOcrConfig - .prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks) - .await - .unwrap(); - let vertex_http = VertexAiOcrConfig - .prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks) - .await - .unwrap(); - assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr"); - assert_eq!( - vertex_http.url(), - "https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict" - ); - for http in [&direct_http, &vertex_http] { - assert_eq!(http.header("authorization").unwrap(), "Bearer test-key"); - assert_eq!(http.header("content-type").unwrap(), "application/json"); - assert_eq!(http.timeout(), Some(Duration::from_secs(2))); - let body: Value = serde_json::from_slice(http.body()).unwrap(); - assert_eq!( - body, - json!({ - "model": "mistral-ocr-maas", - "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "pages": [0, 2], - "include_image_base64": true, - "unknown": "ignored" - }) - ); - } - let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}); - let raw = serde_json::to_vec(&payload).unwrap(); - let direct_response = MistralOcrConfig - .transform_ocr_response(&direct.model, &raw, OcrResponseFormat::Litellm) - .unwrap() - .into_json(); - let vertex_response = VertexAiOcrConfig - .transform_ocr_response(&vertex.model, &raw, OcrResponseFormat::Litellm) - .unwrap() - .into_json(); - assert_eq!(direct_response, vertex_response); - assert_eq!(direct_response["model"], "mistral-ocr-maas"); - assert_eq!(direct_response["object"], "ocr"); - assert_eq!(direct_response["extra"], "preserved"); - } - - mod transformation { - - use rstest::rstest; - use serde_json::{Value, json}; - - use crate::ocr::test_support::wire_request; - - #[rstest] - #[case::mistral(false)] - #[case::vertex(true)] - #[tokio::test] - async fn configs_build_complete_requests_and_share_mistral_normalization( - #[case] use_vertex: bool, - ) { - use std::time::Duration; - - use litellm_llms::{ - base_llm::ocr::transformation::BaseOcrConfig, - mistral::ocr::transformation::MistralOcrConfig, - vertex_ai::ocr::transformation::VertexAiOcrConfig, - }; - - use crate::ocr::test_support::ocr_client; - - let client = ocr_client(); - let options = json!({ - "pages": [0, 2], - "include_image_base64": true, - "vertex_project": "project-1", - "vertex_location": "us-central1", - "unknown": "preserved" - }); - let direct = wire_request( - "mistral/mistral-ocr-maas", - "https://mistral.test", - options.clone(), - ); - let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); - let direct = crate::ocr::prepare::prepare_request_for_test( - crate::ocr::test_support::resolved_request(direct), - ); - let vertex = crate::ocr::prepare::prepare_request_for_test( - crate::ocr::test_support::resolved_request(vertex), - ); - let direct_http = MistralOcrConfig - .prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks) - .await - .unwrap(); - let vertex_http = VertexAiOcrConfig - .prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks) - .await - .unwrap(); - assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr"); - assert_eq!( - vertex_http.url(), - "https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict" - ); - let http = if use_vertex { - &vertex_http - } else { - &direct_http - }; - assert_eq!(http.header("authorization").unwrap(), "Bearer test-key"); - assert_eq!(http.header("content-type").unwrap(), "application/json"); - assert_eq!(http.timeout(), Some(Duration::from_secs(2))); - let body: Value = serde_json::from_slice(http.body()).unwrap(); - assert_eq!( - body, - json!({ - "model": "mistral-ocr-maas", - "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "pages": [0, 2], - "include_image_base64": true, - "unknown": "preserved" - }) - ); - let payload = serde_json::to_vec( - &json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}), - ) - .unwrap(); - let direct_response = MistralOcrConfig - .transform_ocr_response(&direct.model, &payload, Default::default()) - .unwrap() - .into_json(); - let vertex_response = VertexAiOcrConfig - .transform_ocr_response(&vertex.model, &payload, Default::default()) - .unwrap() - .into_json(); - assert_eq!(direct_response, vertex_response); - assert_eq!(direct_response["model"], "mistral-ocr-maas"); - assert_eq!(direct_response["object"], "ocr"); - assert_eq!(direct_response["extra"], "preserved"); - } - } -} - -#[cfg(test)] -mod vertex_ai_deepseek_tests { - use litellm_auth::InputSource; - use serde_json::{Value, json}; - - use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; - - fn request_body(request: &str) -> Value { - serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() - } - - #[tokio::test] - async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "choices":[{"message":{"content":"recognized"}}], - "usage":{"prompt_tokens":1} - }))]) - .await; - let request = wire_request( - "vertex_ai/deepseek-ocr-maas", - &base, - json!({ - "vertex_project":"project-1", - "vertex_location":"europe-west4", - "temperature":0.1, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - ); - let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf"); - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.pages[0].markdown, "recognized"); - assert_eq!( - response.usage_info.unwrap().extra_fields["prompt_tokens"], - 1 - ); - let requests = seen.lock().unwrap(); - assert!(requests[0].starts_with( - "POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions " - )); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer test-key") - ); - let body = request_body(&requests[0]); - assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); - assert_eq!(body["temperature"], 0.1); - assert_eq!(body["future_ocr_option"], true); - assert!(body.get("extra_body").is_none()); - assert_eq!( - body["messages"][0]["content"][0], - json!({"type":"image_url","image_url":"gs://bucket/document.pdf"}) - ); - } - - #[test] - fn host_registration_selects_deepseek_without_affecting_mistral() { - assert!(crate::ocr::arguments::is_supported_request( - "deepseek-ocr-maas", - Some("vertex_ai") - )); - assert!(crate::ocr::arguments::is_supported_request( - "mistral-ocr-maas", - Some("vertex_ai") - )); - } - - #[tokio::test] - async fn request_controlled_api_base_is_rejected_before_vertex_auth() { - let mut request = wire_request( - "vertex_ai/deepseek-ocr-maas", - "https://caller.example", - json!({"vertex_project":"project-1"}), - ); - request.credentials.api_base = Some(litellm_auth::Sourced::new( - "https://caller.example".into(), - InputSource::Request, - )); - - let error = perform_ocr(request).await.unwrap_err(); - assert!( - error - .to_string() - .contains("request-controlled Vertex AI endpoint") - ); - } - - mod deepseek_transformation { - use serde_json::json; - - use super::*; - use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; - - #[tokio::test] - async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "choices":[{"message":{"content":"recognized"}}], - "usage":{"prompt_tokens":1} - }))]) - .await; - let request = wire_request( - "vertex_ai/deepseek-ocr-maas", - &base, - json!({ - "vertex_project":"project-1", - "vertex_location":"europe-west4", - "temperature":0.1, - "future_ocr_option":true, - "extra_body":{"provider_option":"value"} - }), - ); - let request = - crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf"); - - let response = perform_ocr(request).await.unwrap(); - server.await.unwrap(); - assert_eq!(response.pages[0].markdown, "recognized"); - assert_eq!( - response.usage_info.unwrap().extra_fields["prompt_tokens"], - 1 - ); - let requests = seen.lock().unwrap(); - assert!(requests[0].starts_with( - "POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions " - )); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer test-key") - ); - let body = request_body(&requests[0]); - assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); - assert_eq!(body["temperature"], 0.1); - assert_eq!(body["future_ocr_option"], true); - assert_eq!(body["provider_option"], "value"); - assert!(body.get("vertex_project").is_none()); - assert!(body.get("extra_body").is_none()); - assert_eq!( - body["messages"][0]["content"][0], - json!({"type":"image_url","image_url":"gs://bucket/document.pdf"}) - ); - } - } -} - -#[cfg(test)] -pub(crate) mod tests { - use std::sync::{Arc, Mutex}; - - use futures_util::future::BoxFuture; - use litellm_auth_gcp::VertexAuth; - use litellm_host::{ - event::{CallEvent, MachineEvent, WireRequest}, - host::{Host, HostOp, HostResult}, - machine::{HostFailure, Machine, MachineStep}, - }; - use litellm_http::{ - HttpClientPool, HttpSettings, Resolution, - media::{PublicDnsResolver, UrlPolicy}, - }; - use litellm_llms::base_llm::ocr::{ - error::Error as OcrError, - handler::OcrClient, - settings::OcrSettings, - transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OCR_RESPONSE_MAX_BYTES, OcrTransportConfig, - }, - }; - use litellm_secrets::source::SecretSource; - use rstest::rstest; - use serde_json::{Value, json}; - - use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine}; - use crate::ocr::{ - test_support::{ - MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request, - }, - wire::{OcrWireRequest, decode_request}, - }; - - struct RecordingSecretSource { - names: Arc>>, - values: &'static [(&'static str, &'static str)], - api_base: String, - } - - impl SecretSource for RecordingSecretSource { - fn get_secret_str<'a>( - &'a self, - name: &'a str, - ) -> BoxFuture<'a, Result, litellm_secrets::Error>> - { - self.names.lock().unwrap().push(name.to_owned()); - Box::pin(async move { - Ok(match name { - "MISTRAL_AZURE_API_BASE" => Some(self.api_base.clone()), - _ => self - .values - .iter() - .find(|(key, _)| *key == name) - .map(|(_, value)| value.to_string()), - } - .map(litellm_secrets::SecretValue::new)) - }) - } - } - - #[rstest] - #[case::mistral("mistral/model", json!({}))] - #[case::vertex("vertex_ai/mistral-ocr-latest", json!({"vertex_project":"test-project", "vertex_location":"us-central1"}))] - #[tokio::test] - async fn ocr_contract_upstream_error_preserves_status_body_and_headers( - #[case] model: &str, - #[case] options: Value, - ) { - let payload = json!({"message": format!("{} END-OF-PROVIDER-BODY", "x".repeat(4096))}); - let expected_body = serde_json::to_string(&payload).unwrap(); - let (base, seen, server) = mock_server(vec![MockResponse { - status: 422, - headers: vec![ - ("Retry-After", "17".into()), - ("X-Request-ID", "request-123".into()), - ("X-Future-Header", "retained".into()), - ], - body: payload, - }]) - .await; - let error = perform_ocr(wire_request(model, &base, options)) - .await - .unwrap_err(); - server.await.unwrap(); - assert_eq!(seen.lock().unwrap().len(), 1); - let OcrError::Provider { - status, - body, - headers, - } = error - else { - panic!("expected provider error, got {error:?}"); - }; - assert_eq!(status, 422); - for (name, value) in [ - ("retry-after", "17"), - ("x-request-id", "request-123"), - ("x-future-header", "retained"), - ] { - assert!( - headers - .iter() - .any(|(key, actual)| key.eq_ignore_ascii_case(name) && actual == value) - ); - } - assert_eq!( - body.len(), - expected_body.len(), - "provider error body was truncated" - ); - assert_eq!(body, expected_body); - } - - #[test] - fn request_boundary_selects_mistral_and_rejects_unknown_providers() { - let request = OcrWireRequest { - model: "mistral/model".into(), - document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some(litellm_auth::SecretValue::new("key")), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: json!({"extract_header":true,"unknown":42}) - .as_object() - .unwrap() - .clone(), - input_sources: Default::default(), - timeout_seconds: None, - }; - assert!(decode_request(request).is_ok()); - assert!( - decode_request(OcrWireRequest { - model: "model".into(), - document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some(litellm_auth::SecretValue::new("key")), - api_base: None, - custom_llm_provider: Some("unknown".into()), - extra_headers: None, - optional_params: serde_json::Map::new(), - input_sources: Default::default(), - timeout_seconds: None, - }) - .is_err() - ); - } - - #[tokio::test] - async fn facade_executes_direct_mistral_once() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"hello","custom":"preserved"}], - "usage_info":{"pages_processed":1} - }))]) - .await; - let result = perform_ocr(wire_request( - "mistral/model", - &base, - json!({"pages":"0,2-4","extract_header":true,"unknown":"ignored"}), - )) - .await - .unwrap(); - server.await.unwrap(); - assert_eq!(result.pages[0].markdown, "hello"); - assert_eq!(result.pages[0].extra_fields["custom"], "preserved"); - let requests = seen.lock().unwrap(); - assert_eq!(requests.len(), 1); - assert!(requests[0].starts_with("POST /v1/ocr ")); - assert!( - requests[0] - .to_ascii_lowercase() - .contains("authorization: bearer test-key\r\n") - ); - let body: Value = - serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); - assert_eq!( - body, - json!({ - "model":"model", - "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, - "pages":"0,2-4", - "extract_header":true, - "unknown":"ignored" - }) - ); - } - - #[tokio::test] - async fn facade_retains_native_response_when_requested() { - let provider_response = json!({ - "pages":[{"index":0,"markdown":"hello"}], - "usage_info":{"pages_processed":1}, - "provider_only":"preserved" - }); - let (base, _, server) = - mock_server(vec![MockResponse::json(provider_response.clone())]).await; - let response = perform_ocr(wire_request( - "mistral/model", - &base, - json!({"req_format":"native"}), - )) - .await - .unwrap(); - - server.await.unwrap(); - assert_eq!( - response.provider_native_response.map(Value::Object), - Some(provider_response) - ); - } - - #[rstest] - #[case::plain_key(&[("MISTRAL_API_KEY", "plain")], "plain")] - #[case::azure_key_wins(&[("MISTRAL_AZURE_API_KEY", "azure"), ("MISTRAL_API_KEY", "plain")], "azure")] - #[case::empty_azure_key_falls_through(&[("MISTRAL_AZURE_API_KEY", ""), ("MISTRAL_API_KEY", "plain")], "plain")] - #[tokio::test] - async fn mistral_env_fallbacks_follow_python_through_the_injected_secret_source( - #[case] secrets: &'static [(&'static str, &'static str)], - #[case] expected_key: &str, - ) { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let names = Arc::new(Mutex::new(Vec::new())); - let client = ocr_client().with_secrets(Arc::new(RecordingSecretSource { - names: names.clone(), - values: secrets, - api_base: base.clone(), - })); - let request = decode_request(OcrWireRequest { - model: "mistral/model".into(), - document: json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), - api_key: None, - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Default::default(), - input_sources: Default::default(), - timeout_seconds: Some(2.0), - }) - .unwrap(); - - crate::ocr::client::perform(&client, request).await.unwrap(); - server.await.unwrap(); - assert_eq!( - *names.lock().unwrap(), - litellm_llms::mistral::ocr::transformation::MistralOcrConfig.secret_names() - ); - assert!(seen.lock().unwrap()[0].contains(&format!("authorization: Bearer {expected_key}"))); - } - - #[tokio::test] - async fn mistral_ocr_resolves_provider_secrets_before_transformation() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let names = Arc::new(Mutex::new(Vec::new())); - let client = ocr_client().with_secrets(Arc::new(RecordingSecretSource { - names: names.clone(), - values: &[("MISTRAL_API_KEY", "source-key")], - api_base: base.clone(), - })); - let request = decode_request(OcrWireRequest { - model: "mistral/mistral-ocr-latest".into(), - document: json!({ - "type":"document_url", - "document_url":"data:application/pdf;base64,YWJj" - }), - api_key: None, - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Default::default(), - input_sources: Default::default(), - timeout_seconds: Some(2.0), - }) - .unwrap(); - - crate::ocr::client::perform(&client, request).await.unwrap(); - server.await.unwrap(); - assert_eq!( - *names.lock().unwrap(), - litellm_llms::mistral::ocr::transformation::MistralOcrConfig.secret_names() - ); - assert!(seen.lock().unwrap()[0].contains("authorization: Bearer source-key")); - } - - #[tokio::test] - async fn ocr_client_uses_the_injected_http_pool_configuration() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let settings = HttpSettings { - user_agent: Some("host-owned/1".into()), - ..HttpSettings::default() - }; - let client = OcrClient::new( - &HttpClientPool::new(Arc::new(PublicDnsResolver)), - &Resolution::from(&settings).config, - UrlPolicy::default(), - VertexAuth::default(), - OcrSettings::default(), - Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), - ) - .unwrap(); - crate::ocr::client::perform(&client, wire_request("mistral/model", &base, json!({}))) - .await - .unwrap(); - server.await.unwrap(); - assert!(seen.lock().unwrap()[0].contains("user-agent: host-owned/1")); - } - - fn event_name(event: &CallEvent) -> &'static str { - match event { - CallEvent::Started { .. } => "started", - CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response", - CallEvent::Succeeded { .. } => "success", - CallEvent::Failed { .. } => "failure", - } - } - - fn recording_host( - request: crate::ocr::types::LiteLLMOcrRequest, - events: Arc>>, - block: bool, - ) -> LocalOcrHost { - let before_send_events = events.clone(); - LocalOcrHost::new(request) - .with_before_send(move |wire, _| { - before_send_events.lock().unwrap().push("before_send"); - if block { - return Err(OcrError::InvalidRequest("blocked".into())); - } - Ok(wire) - }) - .with_observer(move |event| events.lock().unwrap().push(event_name(event))) - } - - #[tokio::test] - async fn lifecycle_sends_headers_returned_by_the_before_send_operation() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))) - .with_before_send(|mut wire, _| { - wire.headers - .push(("x-core-callback".into(), "edited".into())); - Ok(wire) - }); - - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - - assert!(seen.lock().unwrap()[0].contains("x-core-callback: edited")); - } - - #[tokio::test] - async fn before_send_context_names_the_route_and_its_secrets() { - let (base, _, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let observed = Arc::new(Mutex::new(None)); - let captured = observed.clone(); - let host = LocalOcrHost::new(wire_request( - "mistral/model", - &base, - json!({"pages": [0], "req_format": "native"}), - )) - .with_before_send(move |wire, context| { - *captured.lock().unwrap() = Some((wire.clone(), context.clone())); - Ok(wire) - }); - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - let (wire, context) = observed.lock().unwrap().take().unwrap(); - assert_eq!(context.custom_llm_provider, "mistral"); - assert_eq!(context.model, "model"); - assert_eq!(wire.body["pages"], json!([0])); - assert!(context.secret_fields.is_empty()); - assert_eq!(context.optional_params["req_format"], "native"); - - let (base, _, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let observed = Arc::new(Mutex::new(None)); - let captured = observed.clone(); - let request = wire_request( - "azure_ai/model", - &base, - json!({"client_secret": "shh", "tenant_id": "t"}), - ); - let request = request.with_document(crate::ocr::types::OcrDocumentInput::Bytes { - bytes: b"abc".as_slice().into(), - file_name: None, - mime_type: Some("application/pdf".into()), - }); - let host = LocalOcrHost::new(request).with_before_send(move |wire, context| { - *captured.lock().unwrap() = Some(context.clone()); - Ok(wire) - }); - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - let context = observed.lock().unwrap().take().unwrap(); - assert_eq!(context.secret_fields, ["client_secret"]); - } - - #[tokio::test] - async fn lifecycle_orders_hooks_and_emits_one_success() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let events = Arc::new(Mutex::new(Vec::new())); - let host = recording_host( - wire_request("mistral/model", &base, json!({})), - events.clone(), - false, - ); - perform_ocr_with(host).await.unwrap(); - server.await.unwrap(); - assert_eq!( - *events.lock().unwrap(), - ["started", "before_send", "response", "success"] - ); - assert_eq!(seen.lock().unwrap().len(), 1); - } - - #[tokio::test] - async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() { - let events = Arc::new(Mutex::new(Vec::new())); - let host = recording_host( - wire_request("mistral/model", "http://127.0.0.1:1", json!({})), - events.clone(), - true, - ); - let error = perform_ocr_with(host).await.unwrap_err(); - assert!(matches!(error, OcrError::InvalidRequest(message) if message == "blocked")); - assert_eq!( - *events.lock().unwrap(), - ["started", "before_send", "failure"] - ); - } - - #[tokio::test] - async fn upstream_failure_emits_one_terminal_failure() { - let (base, seen, server) = mock_server(vec![MockResponse { - status: 500, - headers: vec![], - body: json!({"error":"failed"}), - }]) - .await; - let events = Arc::new(Mutex::new(Vec::new())); - let host = recording_host( - wire_request("mistral/model", &base, json!({})), - events.clone(), - false, - ); - assert!(perform_ocr_with(host).await.is_err()); - server.await.unwrap(); - assert_eq!( - *events.lock().unwrap(), - ["started", "before_send", "failure"] - ); - assert_eq!(seen.lock().unwrap().len(), 1); - } - - /// Drives the machine by hand, answering every op through `host` except `before_send`, - /// which `intercept` answers so a test can fail or cancel exactly there. - async fn drive_until( - client: OcrClient, - host: &LocalOcrHost, - mut intercept: impl FnMut(WireRequest) -> Result>, - ) -> ( - Result, - Vec<&'static str>, - crate::ocr::route::OcrMachine, - ) { - let mut machine = ocr_machine(client); - let mut result = None; - let mut ops = Vec::new(); - let outcome = loop { - let op = match machine.resume(result.take()).await { - Ok(MachineStep::Host(op)) => op, - Ok(MachineStep::Complete(response)) => break Ok(response), - Err(error) => break Err(error), - }; - let answer = match op { - HostOp::Route(op) => { - ops.push(match op { - OcrOp::ProjectRequest => "ProjectRequest", - OcrOp::ReadDocument => "ReadDocument", - OcrOp::AcquireAzureAdToken => "AcquireAzureAdToken", - }); - host.route(op) - .await - .map(HostResult::Route) - .map_err(HostFailure::Error) - } - HostOp::BeforeSend { wire, .. } => { - ops.push("BeforeSend"); - intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire))) - } - HostOp::Emit(event) => { - let event = CallEvent::Machine(event); - ops.push(event_name(&event)); - host.emit(&event) - .await - .map(|()| HostResult::Emitted) - .map_err(HostFailure::Error) - } - }; - match answer { - Ok(answer) => result = Some(answer), - Err(failure) => break machine.interrupt(failure).await, - } - }; - (outcome, ops, machine) - } - - #[tokio::test] - async fn failed_before_send_does_not_replay_or_reach_transport() { - let host = LocalOcrHost::new(wire_request( - "mistral/model", - "http://127.0.0.1:1", - json!({}), - )); - let (outcome, ops, mut machine) = drive_until(ocr_client(), &host, |_| { - Err(HostFailure::Error(OcrError::InvalidRequest( - "before_send failed".into(), - ))) - }) - .await; - assert!( - matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "before_send failed") - ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(None).await.is_err()); - } - - #[tokio::test] - async fn invalid_provider_response_emits_response_received_before_normalization_failure() { - let (base, seen, server) = - mock_server(vec![MockResponse::json(json!({"pages":"invalid"}))]).await; - let responses_received = Arc::new(Mutex::new(Vec::new())); - let observed = responses_received.clone(); - let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))) - .with_observer(move |event| { - if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { - observed.lock().unwrap().push(raw.body.clone()); - } - }); - let error = perform_ocr_with(host).await.unwrap_err(); - server.await.unwrap(); - assert!(matches!(error, OcrError::ResponseField { .. })); - assert_eq!(seen.lock().unwrap().len(), 1); - assert_eq!( - *responses_received.lock().unwrap(), - [r#"{"pages":"invalid"}"#] - ); - } - - #[tokio::test] - async fn direct_native_host_drives_the_same_state_machine() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"native"}] - }))]) - .await; - let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))); - let (outcome, ops, mut machine) = drive_until(ocr_client(), &host, Ok).await; - server.await.unwrap(); - assert_eq!(outcome.unwrap().pages[0].markdown, "native"); - assert_eq!(seen.lock().unwrap().len(), 1); - assert_eq!(ops, ["ProjectRequest", "BeforeSend", "response"]); - assert!(matches!( - machine.resume(None).await, - Err(OcrError::InvalidRequest(_)) - )); - } - - async fn drive_native_file_call( - request: crate::ocr::types::LiteLLMOcrRequest, - content: Result, - ) -> (Result, usize) { - let reads = Arc::new(Mutex::new(0)); - let counted = reads.clone(); - let content = Mutex::new(Some(content)); - let host = LocalOcrHost::new(request).with_reader(move || { - *counted.lock().unwrap() += 1; - content.lock().unwrap().take().unwrap() - }); - let outcome = perform_ocr_with(host).await; - let reads = *reads.lock().unwrap(); - (outcome, reads) - } - - #[tokio::test] - async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_encoded() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"file"}] - }))]) - .await; - let request = wire_request("mistral/model", &base, json!({})).with_document( - crate::ocr::types::OcrDocumentInput::HostReader { - mime_type: Some("application/pdf".into()), - }, - ); - let (response, reads) = drive_native_file_call( - request, - Ok(crate::ocr::types::OcrFileContent { - bytes: b"abc".as_slice().into(), - file_name: Some("scan.png".into()), - }), - ) - .await; - server.await.unwrap(); - assert_eq!(response.unwrap().pages[0].markdown, "file"); - assert_eq!(reads, 1); - assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj")); - } - - #[tokio::test] - async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called() { - let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let failure = OcrError::InvalidRequest("reader exploded".into()); - let (response, reads) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Err(failure.clone()), - ) - .await; - assert!( - matches!(response.unwrap_err(), OcrError::InvalidRequest(message) if message == "reader exploded") - ); - assert_eq!(reads, 1); - - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Ok(crate::ocr::types::OcrFileContent { - bytes: Default::default(), - file_name: None, - }), - ) - .await; - assert!(matches!(response.unwrap_err(), OcrError::EmptyFile)); - assert!(seen.lock().unwrap().is_empty()); - } - - #[tokio::test] - async fn path_documents_are_read_by_core_without_a_host_operation() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"path"}] - }))]) - .await; - let dir = std::env::temp_dir().join(format!("litellm-ocr-{}", rand::random::())); - std::fs::create_dir_all(&dir).unwrap(); - let path = dir.join("scan.png"); - std::fs::write(&path, b"abc").unwrap(); - let request = wire_request("mistral/model", &base, json!({})).with_document( - crate::ocr::types::OcrDocumentInput::Path { - path: path.clone(), - mime_type: None, - }, - ); - let (response, reads) = - drive_native_file_call(request, Err(OcrError::InvalidRequest("unused".into()))).await; - server.await.unwrap(); - std::fs::remove_dir_all(&dir).unwrap(); - assert_eq!(response.unwrap().pages[0].markdown, "path"); - assert_eq!(reads, 0); - assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj")); - - let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request.with_document(crate::ocr::types::OcrDocumentInput::Path { - path: path.clone(), - mime_type: None, - }), - Err(OcrError::InvalidRequest("unused".into())), - ) - .await; - assert!(matches!( - response.unwrap_err(), - OcrError::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound - )); - assert!(seen.lock().unwrap().is_empty()); - } - - #[tokio::test] - async fn cancellation_at_before_send_prevents_execution_and_further_resumption() { - let host = LocalOcrHost::new(wire_request( - "mistral/model", - "http://127.0.0.1:1", - json!({}), - )); - let (outcome, ops, mut machine) = drive_until(ocr_client(), &host, |_| { - Err(HostFailure::Cancelled(OcrError::InvalidRequest( - "cancelled".into(), - ))) - }) - .await; - assert!( - matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "cancelled") - ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(Some(HostResult::Emitted)).await.is_err()); - } - - #[tokio::test] - async fn missing_host_result_preserves_pending_operation() { - let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})); - let mut machine = ocr_machine(ocr_client()); - assert!(matches!( - machine.resume(None).await.unwrap(), - MachineStep::Host(HostOp::Route(OcrOp::ProjectRequest)) - )); - assert!(machine.resume(None).await.is_err()); - assert!(matches!( - machine - .resume(Some(HostResult::Route(OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }))) - .await - .unwrap(), - MachineStep::Host(HostOp::BeforeSend { .. }) - )); - } - - async fn read_bounded_response( - response: Vec, - limit: usize, - ) -> Result { - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let address = listener.local_addr().unwrap(); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.unwrap(); - let mut request = [0; 4096]; - assert!(socket.read(&mut request).await.unwrap() > 0); - socket.write_all(&response).await.unwrap(); - std::future::pending::<()>().await; - }); - let response = reqwest::Client::new() - .get(format!("http://{address}")) - .send() - .await - .unwrap(); - let result = tokio::time::timeout( - std::time::Duration::from_secs(2), - litellm_llms::base_llm::ocr::handler::read_response_bytes(response, limit), - ) - .await; - server.abort(); - let _ = server.await; - result.expect("bounded reads must finish without waiting for the rest of an oversized body") - } - - #[tokio::test] - async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() { - use litellm_llms::base_llm::ocr::error::Error; - - for response in [ - "HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh", - "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n", - ] { - assert_eq!( - read_bounded_response(response.as_bytes().to_vec(), 8) - .await - .unwrap(), - "abcdefgh" - ); - } - for response in [ - "HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n", - "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n", - ] { - assert!(matches!( - read_bounded_response(response.as_bytes().to_vec(), 8).await, - Err(Error::TooLarge { limit: 8 }) - )); - } - } - - #[rstest] - #[case::declared("Content-Length: 1000000")] - #[case::chunked("Transfer-Encoding: chunked")] - #[tokio::test] - async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_draining( - #[case] headers: &str, - ) { - let prefix = "x".repeat(4096); - let body = if headers.starts_with("Transfer") { - format!("{:x}\r\n{prefix}\r\n", prefix.len()) - } else { - prefix.clone() - }; - let response = format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}"); - let error = read_bounded_response(response.into_bytes(), prefix.len()) - .await - .unwrap_err(); - match error { - OcrError::Transport(litellm_http::transport::Error::Http { status, body }) => { - assert_eq!(status, 429); - assert_eq!(body, prefix); - } - error => panic!("unexpected error: {error}"), - } - } - - #[test] - fn response_limit_is_validated_and_not_forwarded_to_the_provider() { - let request = wire_request( - "mistral/model", - "http://localhost", - json!({"max_response_bytes": 123}), - ); - assert_eq!(request.transport.max_response_bytes, 123); - assert!(!request.optional_params.contains_key("max_response_bytes")); - for value in [ - json!(0), - json!(-1), - json!(true), - json!("123"), - json!(1.5), - json!(OCR_RESPONSE_MAX_BYTES + 1), - Value::Null, - ] { - let wire = serde_json::from_value(json!({ - "model": "mistral/model", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "optional_params": {"max_response_bytes": value} - })).unwrap(); - let Err(error) = decode_request(wire) else { - panic!("invalid response limit accepted") - }; - assert!(error.to_string().contains("max_response_bytes")); - } - } - - #[derive(Debug)] - struct PendingToken { - entered: Arc, - dropped: Arc, - } - - struct TokenFutureDrop(Arc); - - impl Drop for TokenFutureDrop { - fn drop(&mut self) { - self.0.store(true, std::sync::atomic::Ordering::SeqCst); - } - } - - impl litellm_auth::TokenProvider for PendingToken { - fn acquire(&self) -> litellm_auth::TokenFuture<'_> { - Box::pin(async move { - let _guard = TokenFutureDrop(self.dropped.clone()); - self.entered.notify_one(); - std::future::pending().await - }) - } - } - - #[tokio::test] - async fn interrupt_drops_provider_captures_before_returning() { - use std::sync::atomic::{AtomicBool, Ordering}; - - let entered = Arc::new(tokio::sync::Notify::new()); - let dropped = Arc::new(AtomicBool::new(false)); - let request = wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({})); - let request = crate::ocr::types::LiteLLMOcrRequest { - transport: OcrTransportConfig { - extra_headers: vec![("authorization".into(), "Bearer test-key".into())], - ..request.transport - }, - azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new( - PendingToken { - entered: entered.clone(), - dropped: dropped.clone(), - }, - ))), - ..request - }; - let host = LocalOcrHost::new(request); - let mut machine = ocr_machine(ocr_client()); - let mut result = None; - tokio::time::timeout(std::time::Duration::from_secs(2), async { - loop { - tokio::select! { - _ = entered.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => { - HostResult::BeforeSend(wire) - } - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, - MachineStep::Complete(_) => panic!("pending provider completed"), - }); - } - } - } - }) - .await - .unwrap(); - assert!(!dropped.load(Ordering::SeqCst)); - let selected = OcrError::InvalidRequest("cancelled".into()); - let acknowledgement = machine.interrupt(HostFailure::Cancelled(selected.clone())); - assert!( - dropped.load(Ordering::SeqCst), - "interrupt returned while provider captures were still alive" - ); - assert!( - matches!(acknowledgement.await, Err(OcrError::InvalidRequest(message)) if message == "cancelled") - ); - } - - struct CallerTokenHost { - request: Mutex>, - trace: Mutex>, - } - - impl Host for CallerTokenHost { - async fn route(&self, op: OcrOp) -> Result { - match op { - OcrOp::ProjectRequest => { - self.trace.lock().unwrap().push("project".into()); - Ok(OcrOpResult::Request { - request: Box::new(self.request.lock().unwrap().take().unwrap()), - caller_token: true, - }) - } - OcrOp::AcquireAzureAdToken => { - self.trace.lock().unwrap().push("token".into()); - Ok(OcrOpResult::AzureAdToken( - litellm_auth::ResolvedCredential::Static(litellm_auth::SecretValue::new( - "caller-token", - )), - )) - } - OcrOp::ReadDocument => Err(OcrError::InvalidRequest("no reader".into())), - } - } - - async fn before_send( - &self, - wire: WireRequest, - _: &litellm_host::event::RequestContext, - ) -> Result { - let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization"); - let authorization = wire - .headers - .iter() - .find(|(name, _)| is_authorization(name)) - .map(|(_, value)| value.clone()) - .unwrap_or_default(); - self.trace - .lock() - .unwrap() - .push(format!("before_send:{authorization}")); - let headers = wire - .headers - .into_iter() - .map(|(name, value)| match is_authorization(&name) { - true => (name, "Bearer edited".to_string()), - false => (name, value), - }) - .collect(); - Ok(WireRequest { headers, ..wire }) - } - } - - #[tokio::test] - async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let mut request = wire_request("azure_ai/model", &base, json!({})); - request.credentials.api_key = None; - let host = CallerTokenHost { - request: Mutex::new(Some(request)), - trace: Mutex::new(Vec::new()), - }; - - litellm_host::run::run(ocr_machine(ocr_client()), &host) - .await - .unwrap(); - server.await.unwrap(); - - assert_eq!( - *host.trace.lock().unwrap(), - ["project", "token", "before_send:Bearer caller-token"] - ); - assert!( - seen.lock().unwrap()[0] - .to_ascii_lowercase() - .contains("authorization: bearer edited\r\n") - ); - } - - #[tokio::test] - async fn interrupting_an_in_flight_provider_request_closes_its_connection() { - use tokio::io::AsyncReadExt; - - let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); - let base = format!("http://{}", listener.local_addr().unwrap()); - let received = Arc::new(tokio::sync::Notify::new()); - let server_received = received.clone(); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.unwrap(); - let mut request = Vec::new(); - let mut buffer = [0u8; 4096]; - while !request.windows(4).any(|window| window == b"\r\n\r\n") { - let read = socket.read(&mut buffer).await.unwrap(); - request.extend_from_slice(&buffer[..read]); - } - server_received.notify_one(); - loop { - if socket.read(&mut buffer).await.unwrap() == 0 { - break; - } - } - }); - let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))); - let mut machine = ocr_machine(ocr_client()); - let mut result = None; - tokio::time::timeout(std::time::Duration::from_secs(2), async { - loop { - tokio::select! { - _ = received.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => HostResult::BeforeSend(wire), - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, - MachineStep::Complete(_) => panic!("the stalled provider completed"), - }); - } - } - } - }) - .await - .unwrap(); - - let cancelled = OcrError::InvalidRequest("cancelled".into()); - assert!( - machine - .interrupt(HostFailure::Cancelled(cancelled)) - .await - .is_err() - ); - tokio::time::timeout(std::time::Duration::from_secs(1), server) - .await - .expect("the provider connection stayed open after the interrupt") - .unwrap(); - } -} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 59c9cec8da9..20a21e43676 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -25,9 +25,6 @@ pub enum OcrDocumentInput { file_name: Option, mime_type: Option, }, - HostReader { - mime_type: Option, - }, } impl From for OcrDocumentInput { @@ -45,12 +42,6 @@ impl From for OcrDocumentInput { } } -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct OcrFileContent { - pub bytes: Bytes, - pub file_name: Option, -} - /// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the /// shape hosts receive them: JSON-ish headers, optional timeout, optional /// credentials, and per-field provenance in `input_sources`. diff --git a/litellm-rust/crates/core/tests/audio_transcription.rs b/litellm-rust/crates/core/tests/audio_transcription.rs index aa5aef0149b..196f085a6c3 100644 --- a/litellm-rust/crates/core/tests/audio_transcription.rs +++ b/litellm-rust/crates/core/tests/audio_transcription.rs @@ -1,50 +1,250 @@ -use std::{ - io::{Read, Write}, - net::TcpListener, - thread, +use litellm_core::audio_transcription::{ + Error, audio_transcription, types::AudioTranscriptionRequest, }; +use rstest::{fixture, rstest}; +use serde_json::{Map, Value, json}; +use wiremock::ResponseTemplate; -use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest}; -use serde_json::{Map, json}; +mod support; +use support::*; -#[tokio::test] -async fn bedrock_request_is_signed_and_contains_audio() { - let listener = TcpListener::bind("127.0.0.1:0").expect("listener"); - let address = listener.local_addr().expect("address"); - let server = thread::spawn(move || { - let (mut stream, _) = listener.accept().expect("connection"); - let mut request = Vec::new(); - let mut buffer = [0_u8; 16_384]; - let count = stream.read(&mut buffer).expect("request"); - request.extend_from_slice(&buffer[..count]); - let request = String::from_utf8_lossy(&request); - assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse")); - assert!(request.contains("authorization: AWS4-HMAC-SHA256")); - assert!(request.contains("x-amz-date:")); - assert!(request.contains("\"bytes\":\"AQI=\"")); - assert!(request.contains("Transcribe the audio. Respond with only the transcript.")); - let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}"; - stream.write_all(response).expect("response"); - }); +const MODEL: &str = "mistral.voxtral-mini-3b-2507"; - let optional_params = Map::from_iter([ +fn transcript_response(text: &str) -> ResponseTemplate { + json_response(json!({"output": {"message": {"content": [{"text": text}]}}})) +} + +fn aws_params(region: &str) -> Map { + Map::from_iter([ ("aws_access_key_id".to_string(), json!("access-key")), ("aws_secret_access_key".to_string(), json!("secret-key")), - ("aws_region_name".to_string(), json!("us-east-1")), - ]); - let api_base = format!("http://{address}"); - let response = audio_transcription(AudioTranscriptionRequest { - model: "mistral.voxtral-mini-3b-2507", + ("aws_region_name".to_string(), json!(region)), + ]) +} + +#[fixture] +fn request() -> AudioTranscriptionRequest<'static> { + AudioTranscriptionRequest { + model: MODEL, audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}), api_key: None, - api_base: Some(&api_base), + api_base: None, custom_llm_provider: Some("bedrock"), extra_headers: None, - optional_params, + optional_params: aws_params("us-east-1"), timeout: None, + } +} + +#[rstest] +#[case::us_east_1("us-east-1")] +#[case::eu_west_1("eu-west-1")] +#[tokio::test] +async fn bedrock_converse_request_is_signed_for_the_requested_region( + request: AudioTranscriptionRequest<'static>, + #[case] region: &str, +) { + let upstream = upstream([transcript_response("hello")]).await; + let base = upstream.uri(); + + let response = audio_transcription(AudioTranscriptionRequest { + api_base: Some(&base), + optional_params: aws_params(region), + ..request }) .await .expect("transcription"); + assert_eq!(response, json!({"text": "hello"})); - server.join().expect("server"); + let sent = only_request(&upstream).await; + assert_eq!(sent.method.as_str(), "POST"); + assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse")); + let authorization = sent.header("authorization").expect("request is signed"); + assert!( + authorization.starts_with("AWS4-HMAC-SHA256 Credential=access-key/"), + "{authorization}" + ); + assert!( + authorization.contains(&format!("/{region}/bedrock/aws4_request")), + "{authorization}" + ); + assert!(sent.header("x-amz-date").is_some()); + assert!(!sent.body_text().contains("secret-key")); +} + +#[rstest] +#[tokio::test] +async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscriptionRequest<'static>) { + let upstream = upstream([transcript_response("hello")]).await; + let base = upstream.uri(); + let model = format!("bedrock/{MODEL}"); + + audio_transcription(AudioTranscriptionRequest { + model: &model, + custom_llm_provider: None, + api_base: Some(&base), + ..request + }) + .await + .expect("transcription"); + + assert_eq!( + only_request(&upstream).await.url.path(), + format!("/model/{MODEL}/converse") + ); +} + +#[rstest] +#[tokio::test] +async fn audio_and_transcription_params_reach_the_converse_body( + request: AudioTranscriptionRequest<'static>, + #[values("wav", "mp3", "flac", "ogg")] format: &str, +) { + let upstream = upstream([transcript_response("hello")]).await; + let base = upstream.uri(); + let optional_params = aws_params("us-east-1") + .into_iter() + .chain([ + ("language".to_string(), json!("fr")), + ("temperature".to_string(), json!(0.2)), + ]) + .collect(); + + audio_transcription(AudioTranscriptionRequest { + audio: json!({"data": "AQI=", "format": format}), + api_base: Some(&base), + optional_params, + ..request + }) + .await + .expect("transcription"); + + let body = only_request(&upstream).await.json(); + let content = &body["messages"][0]["content"]; + assert_eq!( + content[0], + json!({"audio": {"format": format, "source": {"bytes": "AQI="}}}) + ); + let instruction = content[1]["text"].as_str().expect("instruction text"); + assert!(instruction.contains("fr"), "{instruction}"); + assert_eq!(body["inferenceConfig"]["temperature"], 0.2); +} + +#[rstest] +#[case::unknown_format(json!({"data": "AQI=", "format": "aac"}))] +#[case::missing_data(json!({"format": "wav"}))] +#[case::not_an_object(json!("AQI="))] +#[tokio::test] +async fn invalid_audio_is_rejected_before_sending( + request: AudioTranscriptionRequest<'static>, + #[case] audio: Value, +) { + let upstream = upstream([transcript_response("hello")]).await; + let base = upstream.uri(); + + let error = audio_transcription(AudioTranscriptionRequest { + audio, + api_base: Some(&base), + ..request + }) + .await + .expect_err("invalid audio is rejected"); + + assert!( + matches!( + error, + Error::InvalidRequest(_) | Error::MissingField(_) | Error::InvalidType { .. } + ), + "{error:?}" + ); + assert!(received(&upstream).await.is_empty()); +} + +#[rstest] +#[case::unknown_provider(MODEL, Some("openai"), "openai")] +#[case::unresolvable_model( + "no-such-model", + None, + "unable to resolve custom_llm_provider for audio transcription request" +)] +#[tokio::test] +async fn unsupported_providers_are_rejected_before_sending( + request: AudioTranscriptionRequest<'static>, + #[case] model: &'static str, + #[case] provider: Option<&'static str>, + #[case] reported: &str, +) { + let error = audio_transcription(AudioTranscriptionRequest { + model, + custom_llm_provider: provider, + api_base: Some(UNREACHABLE_BASE), + ..request + }) + .await + .expect_err("unsupported provider errors"); + + assert_eq!(error, Error::InvalidProvider(reported.into())); +} + +#[rstest] +#[tokio::test] +async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) { + let error = audio_transcription(AudioTranscriptionRequest { + extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])), + api_base: Some(UNREACHABLE_BASE), + ..request + }) + .await + .expect_err("a non-string header is rejected"); + + assert!(matches!(error, Error::Headers(_)), "{error:?}"); +} + +#[rstest] +#[case::throttled(429)] +#[case::server_error(500)] +#[tokio::test] +async fn an_upstream_error_keeps_its_status_and_body( + request: AudioTranscriptionRequest<'static>, + #[case] status: u16, +) { + let upstream = + upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await; + let base = upstream.uri(); + + let error = audio_transcription(AudioTranscriptionRequest { + api_base: Some(&base), + ..request + }) + .await + .expect_err("upstream error propagates"); + + assert_eq!( + error, + Error::Transport(litellm_http::transport::Error::Http { + status, + body: "upstream said no".into() + }) + ); +} + +#[rstest] +#[case::not_json(ResponseTemplate::new(200).set_body_string("not json"))] +#[case::no_output(json_response(json!({"unexpected": true})))] +#[tokio::test] +async fn an_unreadable_success_body_is_an_invalid_response( + request: AudioTranscriptionRequest<'static>, + #[case] response: ResponseTemplate, +) { + let upstream = upstream([response]).await; + let base = upstream.uri(); + + let error = audio_transcription(AudioTranscriptionRequest { + api_base: Some(&base), + ..request + }) + .await + .expect_err("an unreadable body fails"); + + assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}"); } diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs new file mode 100644 index 00000000000..ae96509fe2e --- /dev/null +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -0,0 +1,320 @@ +use std::time::Duration; + +use litellm_core::chat_completions::{ + Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest, +}; +use litellm_http::transport::Error as TransportError; +use rstest::{fixture, rstest}; +use serde_json::{Map, Value, json}; +use wiremock::ResponseTemplate; + +mod support; +use support::*; + +const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; + +fn object(value: Value) -> Map { + let Value::Object(map) = value else { + panic!("expected a json object, got {value}"); + }; + map +} + +fn anthropic_response(body: &str) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_raw(body, "application/json") +} + +fn hi() -> Value { + json!([{"role": "user", "content": "hi"}]) +} + +#[fixture] +fn request() -> ChatCompletionsRequest<'static> { + ChatCompletionsRequest { + model: "anthropic/claude-sonnet-4-5", + messages: hi(), + optional_params: object(json!({"max_tokens": 16})), + api_key: Some("sk-test"), + api_base: None, + custom_llm_provider: None, + extra_headers: None, + timeout: Some(Duration::from_secs(10)), + } +} + +#[rstest] +#[tokio::test] +async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_response( + request: ChatCompletionsRequest<'static>, +) { + let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; + let base = upstream.uri(); + + let response = chat_completions(ChatCompletionsRequest { + messages: json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]), + api_base: Some(&base), + ..request + }) + .await + .expect("call succeeds"); + + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/v1/messages"); + assert_eq!(sent.header_values("x-api-key"), ["sk-test"]); + let body = sent.json(); + assert_eq!(body["model"], "claude-sonnet-4-5"); + assert_eq!( + body["messages"], + json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) + ); + assert_eq!( + body["system"], + json!([{"type": "text", "text": "be terse"}]) + ); + assert_eq!(body["max_tokens"], 16); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello") + ); + assert_eq!(response.usage.total_tokens, 15); +} + +#[rstest] +#[tokio::test] +async fn the_deployment_key_replaces_a_caller_supplied_x_api_key( + request: ChatCompletionsRequest<'static>, +) { + let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; + let base = upstream.uri(); + + chat_completions(ChatCompletionsRequest { + api_base: Some(&base), + extra_headers: Some(object( + json!({"x-api-key": "caller-key", "x-trace": "kept"}), + )), + ..request + }) + .await + .expect("call succeeds"); + + let sent = only_request(&upstream).await; + assert_eq!(sent.header_values("x-api-key"), ["sk-test"]); + assert_eq!(sent.header("x-trace"), Some("kept")); +} + +#[rstest] +#[tokio::test] +async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsRequest<'static>) { + let upstream = upstream([json_response(json!({ + "output": {"message": {"role": "assistant", "content": [{"text": "hello"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15} + }))]) + .await; + let base = upstream.uri(); + + let response = chat_completions(ChatCompletionsRequest { + model: "bedrock/anthropic.claude-sonnet-4-5", + optional_params: object(json!({ + "aws_access_key_id": "access-key", + "aws_secret_access_key": "secret-key", + "aws_region_name": "eu-west-1" + })), + api_key: None, + api_base: Some(&base), + ..request + }) + .await + .expect("call succeeds"); + + let sent = only_request(&upstream).await; + assert_eq!( + sent.url.path(), + "/model/anthropic.claude-sonnet-4-5/converse" + ); + let authorization = sent.header("authorization").expect("request is signed"); + assert!( + authorization.contains("/eu-west-1/bedrock/aws4_request"), + "{authorization}" + ); + assert_eq!( + sent.json()["messages"], + json!([{"role": "user", "content": [{"text": "hi"}]}]) + ); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello") + ); + assert_eq!(response.usage.total_tokens, 15); +} + +/// The provider already answered and billed these, so the host must not retry them on +/// its own path: they surface as `InvalidResponse`, never as a pre-send decline. +#[rstest] +#[case::missing_usage( + r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"# +)] +#[case::tool_use_block(r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#)] +#[case::not_json("not json")] +#[tokio::test] +async fn a_response_it_cannot_normalize_is_reported_as_already_sent( + request: ChatCompletionsRequest<'static>, + #[case] body: &str, +) { + let upstream = upstream([anthropic_response(body)]).await; + let base = upstream.uri(); + + let error = chat_completions(ChatCompletionsRequest { + api_base: Some(&base), + ..request + }) + .await + .expect_err("response cannot be normalized"); + + assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}"); +} + +#[rstest] +#[case::rate_limited(429)] +#[case::server_error(500)] +#[tokio::test] +async fn an_upstream_error_status_keeps_its_code_and_body( + request: ChatCompletionsRequest<'static>, + #[case] status: u16, +) { + let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await; + let base = upstream.uri(); + + let error = chat_completions(ChatCompletionsRequest { + api_base: Some(&base), + ..request + }) + .await + .expect_err("upstream rejects"); + + assert_eq!( + error, + Error::Transport(TransportError::Http { + status, + body: "slow down".into() + }) + ); +} + +/// Nothing was sent, so nothing was billed and the host can still serve the request. +#[rstest] +#[tokio::test] +async fn a_connection_that_is_never_established_declines_instead_of_failing( + request: ChatCompletionsRequest<'static>, +) { + let error = chat_completions(ChatCompletionsRequest { + api_base: Some(UNREACHABLE_BASE), + ..request + }) + .await + .expect_err("nothing is listening"); + + assert!( + matches!(error, Error::Transport(TransportError::Connect(_))), + "{error:?}" + ); +} + +#[rstest] +#[tokio::test] +async fn a_timeout_after_sending_is_not_a_pre_send_decline( + request: ChatCompletionsRequest<'static>, +) { + let upstream = + upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await; + let base = upstream.uri(); + + let error = chat_completions(ChatCompletionsRequest { + api_base: Some(&base), + timeout: Some(Duration::from_millis(100)), + ..request + }) + .await + .expect_err("the call times out"); + + assert!( + matches!(error, Error::Transport(TransportError::Network(_))), + "{error:?}" + ); +} + +#[rstest] +#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)] +#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)] +#[case::unknown_provider( + "gpt-4o", + Some("openai"), + hi(), + json!({}), + Some("provider is not on the rust chat completions path") +)] +#[case::unreadable_messages( + "anthropic/claude-sonnet-4-5", + None, + json!("hi"), + json!({}), + Some("unreadable message list") +)] +#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))] +#[case::streaming( + "anthropic/claude-sonnet-4-5", + None, + hi(), + json!({"stream": true}), + Some("streaming") +)] +#[case::unrecognized_param( + "anthropic/claude-sonnet-4-5", + None, + hi(), + json!({"not_a_param": 1}), + Some("unrecognized request parameter") +)] +#[case::opens_on_assistant_turn( + "anthropic/claude-sonnet-4-5", + None, + json!([{"role": "assistant", "content": "hi"}]), + json!({}), + Some("conversation does not open on a user turn") +)] +fn decline_reason_names_why_the_core_would_not_serve_the_request( + #[case] model: &str, + #[case] provider: Option<&str>, + #[case] messages: Value, + #[case] params: Value, + #[case] reason: Option<&str>, +) { + assert_eq!( + chat_completions_decline_reason(model, provider, messages, &object(params)), + reason + ); +} + +/// A request the decline check accepts must not be declined by the call itself. +#[rstest] +#[tokio::test] +async fn a_declined_request_fails_the_call_before_sending( + request: ChatCompletionsRequest<'static>, +) { + let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; + let base = upstream.uri(); + + let error = chat_completions(ChatCompletionsRequest { + optional_params: object(json!({"stream": true})), + api_base: Some(&base), + ..request + }) + .await + .expect_err("streaming is declined"); + + assert_eq!(error, Error::Unsupported("streaming")); + assert!(received(&upstream).await.is_empty()); +} diff --git a/litellm-rust/crates/core/tests/messages.rs b/litellm-rust/crates/core/tests/messages.rs deleted file mode 100644 index 18af8a7d619..00000000000 --- a/litellm-rust/crates/core/tests/messages.rs +++ /dev/null @@ -1,471 +0,0 @@ -use std::{sync::Arc, time::Duration}; - -use futures_util::future::BoxFuture; -use litellm_core::messages::{ - Error, messages, - route::{LocalMessagesHost, MessagesCall, messages_machine}, - types::{MessagesRequest, MessagesShaping}, -}; -use litellm_secrets::{SecretValue, source::SecretSource}; -use serde_json::{Map, Value, json}; -use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - net::{TcpListener, TcpStream}, -}; - -struct RecordingSecrets { - values: Vec<(&'static str, String)>, - fails: bool, - requested: std::sync::Mutex>, -} - -impl RecordingSecrets { - fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self { - Self { - values, - fails, - requested: std::sync::Mutex::new(Vec::new()), - } - } -} - -impl SecretSource for RecordingSecrets { - fn get_secret_str<'a>( - &'a self, - name: &'a str, - ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { - Box::pin(async move { - self.requested.lock().unwrap().push(name.to_string()); - if self.fails { - return Err(litellm_secrets::Error::ManagedSecretMissing); - } - Ok(self - .values - .iter() - .find(|(key, _)| *key == name) - .map(|(_, value)| SecretValue::new(value.clone()))) - }) - } -} - -fn secrets_call() -> MessagesCall { - let Value::Object(body) = json!({ - "model": "claude-sonnet-4-5", - "max_tokens": 16, - "messages": [{"role": "user", "content": "hi"}] - }) else { - unreachable!("literal object") - }; - MessagesCall { - model: "claude-sonnet-4-5".into(), - body, - api_key: None, - api_base: None, - custom_llm_provider: Some("anthropic".into()), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - } -} - -#[tokio::test] -async fn route_surfaces_a_secret_manager_failure_before_the_call() { - let Err(error) = litellm_host::run::run( - messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))), - &LocalMessagesHost::new(secrets_call()), - ) - .await - else { - panic!("a secret manager failure fails the call"); - }; - assert!( - matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)), - "{error:?}" - ); -} - -async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - } - String::from_utf8(request).expect("request is utf8") -} - -fn write_response(body: &str) -> String { - format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - body.len(), - body - ) -} - -#[tokio::test] -async fn messages_round_trip_builds_azure_request_and_passes_response_through() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#; - socket - .write_all(write_response(response_body).as_bytes()) - .await - .expect("writes response"); - request - }); - - let response = messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({ - "model": "claude-sonnet-4-5", - "max_tokens": 1024, - "messages": [{ - "role": "user", - "content": [{ - "type": "text", - "text": "hi", - "cache_control": {"type": "ephemeral", "scope": "global"} - }] - }] - }), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect("messages request succeeds"); - - assert_eq!(response.content[0]["text"], "hi"); - assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); - - let request = server.await.expect("server task completes"); - let (head, body) = request.split_once("\r\n\r\n").expect("has body"); - assert!(head.starts_with("POST /anthropic/v1/messages "), "{head}"); - let head_lower = head.to_ascii_lowercase(); - assert!(head_lower.contains("x-api-key: sk-azure"), "{head}"); - assert!( - head_lower.contains("anthropic-version: 2023-06-01"), - "{head}" - ); - assert!( - head_lower.contains("content-type: application/json"), - "{head}" - ); - - let sent_body: Value = serde_json::from_str(body).expect("body is json"); - assert_eq!( - sent_body["messages"][0]["content"][0]["cache_control"], - json!({"type": "ephemeral"}) - ); -} - -#[tokio::test] -async fn messages_round_trip_builds_native_anthropic_request() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#; - socket - .write_all(write_response(response_body).as_bytes()) - .await - .expect("writes response"); - request - }); - - let response = messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({ - "model": "claude-sonnet-4-5", - "max_tokens": 1024, - "messages": [{"role": "user", "content": "hi"}] - }), - api_key: Some("sk-ant"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("anthropic"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect("messages request succeeds"); - - assert_eq!(response.content[0]["text"], "hi"); - assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); - - let request = server.await.expect("server task completes"); - let (head, _) = request.split_once("\r\n\r\n").expect("has body"); - assert!(head.starts_with("POST /v1/messages "), "{head}"); - let head_lower = head.to_ascii_lowercase(); - assert!(head_lower.contains("x-api-key: sk-ant"), "{head}"); - assert!( - head_lower.contains("anthropic-version: 2023-06-01"), - "{head}" - ); -} - -#[tokio::test] -async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = - r#"{"id":"msg_2","type":"message","role":"assistant","content":[],"model":"m"}"#; - socket - .write_all(write_response(response_body).as_bytes()) - .await - .expect("writes response"); - request - }); - - let mut headers = Map::new(); - headers.insert( - "x-api-key".to_string(), - Value::String("from-python".to_string()), - ); - headers.insert( - "anthropic-beta".to_string(), - Value::String("token-efficient-tools-2025-02-19".to_string()), - ); - - messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("rust-fallback-key"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: Some(headers), - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect("messages request succeeds"); - - let request = server.await.expect("server task completes"); - let head = request - .split_once("\r\n\r\n") - .expect("has body") - .0 - .to_ascii_lowercase(); - let api_key_count = head - .lines() - .filter(|line| line.starts_with("x-api-key:")) - .count(); - assert_eq!(api_key_count, 1, "{head}"); - assert!(head.contains("x-api-key: from-python"), "{head}"); - assert!( - head.contains("anthropic-beta: token-efficient-tools-2025-02-19"), - "{head}" - ); - assert!(!head.contains("rust-fallback-key"), "{head}"); -} - -#[tokio::test] -async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = - r#"{"id":"msg_3","type":"message","role":"assistant","content":[],"model":"m"}"#; - socket - .write_all(write_response(response_body).as_bytes()) - .await - .expect("writes response"); - request - }); - - let mut headers = Map::new(); - headers.insert( - "Authorization".to_string(), - Value::String("Bearer entra-token".to_string()), - ); - - messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: None, - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: Some(headers), - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect("entra id request succeeds without api key"); - - let request = server.await.expect("server task completes"); - let head = request - .split_once("\r\n\r\n") - .expect("has body") - .0 - .to_ascii_lowercase(); - assert!(head.contains("authorization: bearer entra-token"), "{head}"); - assert!(!head.contains("x-api-key"), "{head}"); -} - -#[tokio::test] -async fn messages_requires_auth_when_no_key_and_no_header() { - let err = messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: None, - api_base: Some("http://127.0.0.1:1"), - custom_llm_provider: Some("azure_ai"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_millis(50)), - shaping: MessagesShaping::default(), - }) - .await - .expect_err("missing auth errors"); - - assert!(matches!(err, Error::Auth(_))); -} - -#[tokio::test] -async fn messages_ignores_malformed_authorization_and_uses_api_key() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let request = read_http_request(&mut socket).await; - let response_body = - r#"{"id":"msg_4","type":"message","role":"assistant","content":[],"model":"m"}"#; - socket - .write_all(write_response(response_body).as_bytes()) - .await - .expect("writes response"); - request - }); - - let mut headers = Map::new(); - headers.insert( - "Authorization".to_string(), - Value::String("Bearer ".to_string()), - ); - - messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: Some(headers), - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect("falls back to api key"); - - let request = server.await.expect("server task completes"); - let head = request - .split_once("\r\n\r\n") - .expect("has body") - .0 - .to_ascii_lowercase(); - assert!(head.contains("x-api-key: sk-azure"), "{head}"); -} - -#[tokio::test] -async fn messages_maps_provider_error_status_to_http_error() { - let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); - let addr = listener.local_addr().expect("addr"); - - tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts request"); - let _ = read_http_request(&mut socket).await; - let body = "unauthorized"; - let response = format!( - "HTTP/1.1 401 Unauthorized\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - }); - - let err = messages(MessagesRequest { - model: "claude-sonnet-4-5", - body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), - api_key: Some("sk-azure"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_secs(5)), - shaping: MessagesShaping::default(), - }) - .await - .expect_err("provider error propagates"); - - assert!(matches!( - err, - Error::Transport(litellm_http::transport::Error::Http { status: 401, .. }) - )); -} - -#[tokio::test] -async fn messages_rejects_unsupported_provider() { - let err = messages(MessagesRequest { - model: "claude-3-5-sonnet", - body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}), - api_key: Some("sk"), - api_base: Some("http://127.0.0.1:1"), - custom_llm_provider: Some("openai"), - extra_headers: None, - provider_specific_header: None, - timeout: Some(Duration::from_millis(50)), - shaping: MessagesShaping::default(), - }) - .await - .expect_err("unsupported provider errors"); - - assert!(matches!(err, Error::InvalidProvider(provider) if provider == "openai")); -} diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs new file mode 100644 index 00000000000..ca2aece5ebd --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -0,0 +1,210 @@ +use std::{convert::Infallible, sync::Mutex}; + +use litellm_core::messages::route::Messages; +use litellm_host::{ + event::{CallEvent, MachineEvent, RequestContext, WireRequest}, + host::Host, +}; +use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities; +use rstest::rstest; + +use super::*; + +type Rewrite = Box Result + Send + Sync>; + +/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps +/// every event the driver emits. +struct RecordingHost { + call: LocalMessagesHost, + rewrite: Rewrite, + events: Mutex>, + optional_params: Mutex>, +} + +impl RecordingHost { + fn new(call: MessagesCall, rewrite: Rewrite) -> Self { + Self { + call: LocalMessagesHost::new(call), + rewrite, + events: Mutex::new(Vec::new()), + optional_params: Mutex::new(Vec::new()), + } + } + + fn passthrough(call: MessagesCall) -> Self { + Self::new(call, Box::new(Ok)) + } + + fn raw_responses(&self) -> Vec { + self.events + .lock() + .unwrap() + .iter() + .filter_map(|event| match event { + CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + Some(raw.body.clone()) + } + _ => None, + }) + .collect() + } +} + +impl Host for RecordingHost { + async fn project(&self) -> Result { + self.call.project().await + } + + async fn custom_op(&self, op: Infallible) -> Result<(), Error> { + match op {} + } + + async fn before_send( + &self, + wire: WireRequest, + context: &RequestContext, + ) -> Result { + self.optional_params + .lock() + .unwrap() + .push(context.optional_params.clone()); + (self.rewrite)(wire) + } + + async fn emit(&self, event: &CallEvent) -> Result<(), Error> { + self.events.lock().unwrap().push(event.clone()); + Ok(()) + } +} + +async fn run_through(host: &RecordingHost) -> Result { + litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await +} + +fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall { + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(api_base), + ..call + } +} + +#[rstest] +#[tokio::test] +async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + let host = RecordingHost::new( + authenticated(call, upstream.uri()), + Box::new(|wire| { + let mut body = wire.body; + body["system"] = json!("added by the host"); + Ok(WireRequest { + headers: wire + .headers + .into_iter() + .chain([("x-host".to_string(), "seen".to_string())]) + .collect(), + body, + ..wire + }) + }), + ); + + run_through(&host).await.expect("messages call succeeds"); + + let request = only_request(&upstream).await; + assert_eq!(request.json()["system"], "added by the host"); + assert_eq!(request.header("x-host"), Some("seen")); + assert_eq!(request.header("x-api-key"), Some("sk-ant")); +} + +#[rstest] +#[tokio::test] +async fn a_before_send_failure_never_sends(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + let host = RecordingHost::new( + authenticated(call, upstream.uri()), + Box::new(|_| Err(Error::InvalidRequest("vetoed by the host".into()))), + ); + + let error = run_through(&host) + .await + .err() + .expect("the host failure fails the call"); + + assert_eq!(error, Error::InvalidRequest("vetoed by the host".into())); + assert!(received(&upstream).await.is_empty()); + assert!(host.raw_responses().is_empty()); +} + +#[rstest] +#[tokio::test] +async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) { + let raw = message_body(); + let upstream = upstream([json_response(raw.clone())]).await; + let host = RecordingHost::passthrough(authenticated(call, upstream.uri())); + + let output = run_through(&host).await.expect("messages call succeeds"); + + assert!(matches!(output, MessagesOutput::Message(_))); + let [emitted] = <[String; 1]>::try_from(host.raw_responses()) + .unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len())); + assert_eq!(serde_json::from_str::(&emitted).unwrap(), raw); +} + +#[rstest] +#[case::upstream_error(ResponseTemplate::new(500).set_body_string("boom"))] +#[case::stream(ResponseTemplate::new(200).set_body_raw("event: message_stop\ndata: {}\n\n", "text/event-stream"))] +#[tokio::test] +async fn no_raw_response_is_emitted_for_a_stream_or_a_failure( + call: MessagesCall, + #[case] response: ResponseTemplate, +) { + let upstream = upstream([response]).await; + let mut body = call.body.clone(); + body.insert("stream".into(), json!(true)); + let host = + RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri())); + + let _ = run_through(&host).await; + + assert_eq!(received(&upstream).await.len(), 1); + assert!(host.raw_responses().is_empty()); +} + +/// Python logs `optional_params` as what it is about to send, so a dropped param must +/// not resurface in callbacks. +#[rstest] +#[tokio::test] +async fn the_request_context_carries_the_shaped_params_without_model_or_messages( + call: MessagesCall, +) { + let upstream = upstream([message_response()]).await; + let body: Map = call + .body + .clone() + .into_iter() + .chain([("temperature".to_string(), json!(0.2))]) + .collect(); + let host = RecordingHost::passthrough(authenticated( + MessagesCall { + body, + shaping: MessagesShaping { + capabilities: AnthropicModelCapabilities { + supports_sampling_params: false, + ..AnthropicModelCapabilities::default() + }, + drop_params: true, + ..MessagesShaping::default() + }, + ..call + }, + upstream.uri(), + )); + + run_through(&host).await.expect("messages call succeeds"); + + let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap()) + .unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len())); + assert_eq!(optional_params, json!({"max_tokens": 16})); +} diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs new file mode 100644 index 00000000000..21ee678ced3 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -0,0 +1,95 @@ +use std::{sync::Arc, time::Duration}; + +use litellm_core::messages::{ + Error, + route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}, + types::MessagesShaping, +}; +use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use rstest::fixture; +use serde_json::{Map, Value, json}; +use wiremock::ResponseTemplate; + +#[path = "../support/mod.rs"] +mod support; +use support::*; + +mod host; +mod request; +mod response; +mod secrets; +mod stream; + +const MODEL: &str = "claude-sonnet-4-5"; + +fn object(value: Value) -> Map { + let Value::Object(map) = value else { + panic!("expected a json object, got {value}"); + }; + map +} + +fn message_body() -> Value { + json!({ + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "hi"}], + "model": MODEL, + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 2} + }) +} + +fn message_response() -> ResponseTemplate { + json_response(message_body()) +} + +/// A non-streaming call with nothing that would authenticate or route it, so each test +/// states the provider, credentials, and base it depends on. +#[fixture] +fn call() -> MessagesCall { + MessagesCall { + model: MODEL.into(), + body: object(json!({ + "model": MODEL, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}] + })), + api_key: None, + api_base: None, + custom_llm_provider: Some("anthropic".into()), + extra_headers: None, + provider_specific_header: None, + timeout: Some(Duration::from_secs(5)), + shaping: MessagesShaping::default(), + } +} + +fn headers<'a>(pairs: impl IntoIterator) -> Option> { + Some( + pairs + .into_iter() + .map(|(name, value)| (name.to_string(), Value::from(value))) + .collect(), + ) +} + +async fn run_with( + secrets: Arc, + call: MessagesCall, +) -> Result { + litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await +} + +/// Runs the route with a secret source that knows nothing, so no environment leaks in. +async fn run(call: MessagesCall) -> Result { + run_with(Arc::new(RecordingSecrets::empty()), call).await +} + +async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse { + match run(call).await.expect("messages call succeeds") { + MessagesOutput::Message(message) => *message, + MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"), + } +} diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs new file mode 100644 index 00000000000..2927356b773 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -0,0 +1,675 @@ +use litellm_llms::anthropic::common_utils::{ + ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities, + SupportedEffortTiers, beta, +}; +use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use rstest::rstest; + +use super::*; + +#[rstest] +#[case::anthropic_key("anthropic", Some("sk-ant"), &[], ("x-api-key", "sk-ant"), &["authorization"])] +#[case::azure_key("azure_ai", Some("sk-azure"), &[], ("x-api-key", "sk-azure"), &["authorization"])] +#[case::caller_x_api_key_wins( + "azure_ai", + Some("rust-fallback-key"), + &[("x-api-key", "from-python")], + ("x-api-key", "from-python"), + &["authorization"] +)] +#[case::entra_bearer_without_key( + "azure_ai", + None, + &[("Authorization", "Bearer entra-token")], + ("authorization", "Bearer entra-token"), + &["x-api-key"] +)] +#[case::empty_bearer_falls_back_to_key( + "azure_ai", + Some("sk-azure"), + &[("Authorization", "Bearer ")], + ("x-api-key", "sk-azure"), + &[] +)] +#[case::anthropic_forwards_caller_authorization( + "anthropic", + Some("sk-ant"), + &[("Authorization", "Bearer caller")], + ("authorization", "Bearer caller"), + &["x-api-key"] +)] +#[case::anthropic_oauth_key_becomes_bearer( + "anthropic", + Some("sk-ant-oat01-token"), + &[], + ("authorization", "Bearer sk-ant-oat01-token"), + &["x-api-key"] +)] +#[tokio::test] +async fn credentials_become_exactly_one_auth_header( + call: MessagesCall, + #[case] provider: &str, + #[case] api_key: Option<&str>, + #[case] extra_headers: &[(&str, &str)], + #[case] expected: (&str, &str), + #[case] absent: &[&str], +) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + custom_llm_provider: Some(provider.into()), + api_key: api_key.map(Into::into), + api_base: Some(upstream.uri()), + extra_headers: headers(extra_headers.iter().copied()), + ..call + }) + .await; + + let request = only_request(&upstream).await; + let (name, value) = expected; + assert_eq!(request.header_values(name), [value]); + for name in absent { + assert_eq!(request.header(name), None, "{name} must not be sent"); + } +} + +#[rstest] +#[case::anthropic("anthropic")] +#[case::azure_ai("azure_ai")] +#[tokio::test] +async fn a_call_without_credentials_fails_before_sending( + call: MessagesCall, + #[case] provider: &str, +) { + let upstream = upstream([message_response()]).await; + + let error = run(MessagesCall { + custom_llm_provider: Some(provider.into()), + api_base: Some(upstream.uri()), + ..call + }) + .await + .err() + .expect("a call without credentials fails"); + + assert!( + matches!( + error, + Error::Auth(litellm_auth::Error::MissingApiKey { .. }) + ), + "{error:?}" + ); + assert!(received(&upstream).await.is_empty()); +} + +#[rstest] +#[case::anthropic(MODEL, Some("anthropic"), "", "/v1/messages")] +#[case::anthropic_base_with_trailing_slash(MODEL, Some("anthropic"), "/", "/v1/messages")] +#[case::anthropic_base_with_the_messages_path( + MODEL, + Some("anthropic"), + "/v1/messages", + "/v1/messages" +)] +#[case::azure_ai(MODEL, Some("azure_ai"), "", "/anthropic/v1/messages")] +#[case::provider_from_model_prefix("anthropic/claude-sonnet-4-5", None, "", "/v1/messages")] +#[tokio::test] +async fn each_provider_posts_to_its_messages_endpoint( + call: MessagesCall, + #[case] model: &str, + #[case] provider: Option<&str>, + #[case] base_suffix: &str, + #[case] path: &str, +) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + model: model.into(), + custom_llm_provider: provider.map(Into::into), + api_key: Some("sk".into()), + api_base: Some(format!("{}{base_suffix}", upstream.uri())), + ..call + }) + .await; + + let request = only_request(&upstream).await; + assert_eq!(request.method.as_str(), "POST"); + assert_eq!(request.url.path(), path); + assert_eq!(request.json()["model"], MODEL); + assert_eq!(request.header_values("anthropic-version"), ["2023-06-01"]); + assert_eq!(request.header_values("content-type"), ["application/json"]); +} + +#[rstest] +#[case::unknown_provider(MODEL, Some("openai"), "openai")] +#[case::unresolvable_model( + "no-such-model", + None, + "unable to resolve custom_llm_provider for messages request" +)] +#[tokio::test] +async fn unsupported_providers_are_rejected_before_sending( + call: MessagesCall, + #[case] model: &str, + #[case] provider: Option<&str>, + #[case] reported: &str, +) { + let error = run(MessagesCall { + model: model.into(), + custom_llm_provider: provider.map(Into::into), + api_key: Some("sk".into()), + api_base: Some(UNREACHABLE_BASE.into()), + ..call + }) + .await + .err() + .expect("unsupported provider errors"); + + assert_eq!(error, Error::InvalidProvider(reported.into())); +} + +#[rstest] +#[tokio::test] +async fn caller_headers_and_provider_scoped_headers_are_forwarded(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + let scoped = |provider: &str, value: &str| ProviderSpecificHeader { + custom_llm_provider: provider.into(), + extra_headers: object(json!({"x-scoped": value})), + }; + + run_message(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + extra_headers: headers([("anthropic-beta", "token-efficient-tools-2025-02-19")]), + provider_specific_header: Some(ProviderSpecificHeaders::Many(vec![ + scoped("bedrock", "other-provider"), + scoped("azure_ai, anthropic", "this-provider"), + ])), + ..call + }) + .await; + + let request = only_request(&upstream).await; + assert_eq!( + request.header("anthropic-beta"), + Some("token-efficient-tools-2025-02-19") + ); + assert_eq!(request.header_values("x-scoped"), ["this-provider"]); +} + +#[rstest] +#[tokio::test] +async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + custom_llm_provider: Some("azure_ai".into()), + api_key: Some("sk-azure".into()), + api_base: Some(upstream.uri()), + body: object(json!({ + "model": MODEL, + "max_tokens": 16, + "messages": [{ + "role": "user", + "content": [{ + "type": "text", + "text": "hi", + "cache_control": {"type": "ephemeral", "scope": "global"} + }] + }] + })), + ..call + }) + .await; + + assert_eq!( + only_request(&upstream).await.json()["messages"][0]["content"][0]["cache_control"], + json!({"type": "ephemeral"}) + ); +} + +#[rstest] +#[tokio::test] +async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + let mut body = call.body.clone(); + body.insert("temperature".into(), json!(0.5)); + body.insert("top_k".into(), json!(3)); + + run_message(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + body, + shaping: MessagesShaping { + additional_drop_params: vec!["temperature".into()], + ..MessagesShaping::default() + }, + ..call + }) + .await; + + let sent = only_request(&upstream).await.json(); + assert_eq!(sent.get("temperature"), None); + assert_eq!(sent["top_k"], 3); +} + +fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall { + let body: Map = call.body.into_iter().chain(object(fields)).collect(); + MessagesCall { body, ..call } +} + +fn sent_betas(request: &wiremock::Request) -> Vec { + let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta")) + .unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}")); + header + .split(',') + .map(str::trim) + .map(str::to_string) + .collect() +} + +#[rstest] +#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])] +#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])] +#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])] +#[case::context_management_edits( + json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}), + &[beta::CONTEXT_MANAGEMENT_2025_06_27] +)] +#[case::per_message_output_config( + json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}), + &[beta::PER_TURN_CONTROL_2026_07_01] +)] +#[case::advisor_tool( + json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}), + &[beta::ADVISOR_TOOL_2026_03_01] +)] +#[case::several_features_at_once( + json!({"speed": "fast", "output_format": {"type": "json_schema"}}), + &[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01] +)] +#[tokio::test] +async fn feature_betas_join_the_callers_betas_in_one_sorted_header( + call: MessagesCall, + #[case] fields: Value, + #[case] features: &[&str], +) { + let upstream = upstream([message_response()]).await; + let capabilities = AnthropicModelCapabilities { + supports_speed: true, + ..AnthropicModelCapabilities::default() + }; + + run_message(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + extra_headers: headers([("Anthropic-Beta", "caller-beta-2025-01-01")]), + shaping: MessagesShaping { + capabilities, + ..MessagesShaping::default() + }, + ..call + }, + fields, + )) + .await; + + let sent = sent_betas(&only_request(&upstream).await); + let mut expected: Vec = features + .iter() + .map(|feature| feature.to_string()) + .chain(["caller-beta-2025-01-01".to_string()]) + .collect(); + expected.sort(); + assert_eq!(sent, expected); +} + +#[rstest] +#[tokio::test] +async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + api_key: Some("sk-ant-oat01-token".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await; + + let request = only_request(&upstream).await; + assert_eq!( + request.header("anthropic-dangerous-direct-browser-access"), + Some("true") + ); + assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]); + assert_eq!(request.header("x-api-key"), None); +} + +#[rstest] +#[case::anthropic("anthropic")] +#[case::azure_ai("azure_ai")] +#[tokio::test] +async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[case] provider: &str) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + custom_llm_provider: Some(provider.into()), + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + extra_headers: headers([ + ("Anthropic-Version", "2024-01-01"), + ("Content-Type", "application/json; charset=utf-8"), + ]), + ..call + }) + .await; + + let request = only_request(&upstream).await; + assert_eq!(request.header_values("anthropic-version"), ["2024-01-01"]); + assert_eq!( + request.header_values("content-type"), + ["application/json; charset=utf-8"] + ); +} + +fn sampling_removed() -> AnthropicModelCapabilities { + AnthropicModelCapabilities { + supports_sampling_params: false, + ..AnthropicModelCapabilities::default() + } +} + +#[rstest] +#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")] +#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")] +#[tokio::test] +async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it( + call: MessagesCall, + #[case] capabilities: AnthropicModelCapabilities, + #[case] fields: Value, + #[case] dropped: &[&str], + #[case] rejected_as: &str, +) { + let upstream = upstream([message_response(), message_response()]).await; + let shaped = |drop_params: bool| { + with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + shaping: MessagesShaping { + capabilities: capabilities.clone(), + drop_params, + ..MessagesShaping::default() + }, + body: call.body.clone(), + custom_llm_provider: call.custom_llm_provider.clone(), + extra_headers: None, + provider_specific_header: None, + model: call.model.clone(), + timeout: call.timeout, + }, + fields.clone(), + ) + }; + + let error = run(shaped(false)) + .await + .err() + .expect("an unsupported param is rejected without drop_params"); + assert!( + matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)), + "{error:?}" + ); + assert!(received(&upstream).await.is_empty()); + + run_message(shaped(true)).await; + let sent = only_request(&upstream).await.json(); + for name in dropped { + assert_eq!(sent.get(*name), None, "{name} must be dropped"); + } + assert_eq!(sent["max_tokens"], 16); +} + +#[rstest] +#[case::adaptive_thinking(json!({"type": "adaptive"}), json!({"type": "adaptive", "display": "summarized"}))] +#[case::disabled_thinking(json!({"type": "disabled"}), json!({"type": "disabled"}))] +#[tokio::test] +async fn reasoning_auto_summary_marks_active_thinking_on_the_wire( + call: MessagesCall, + #[case] thinking: Value, + #[case] expected: Value, +) { + let upstream = upstream([message_response()]).await; + + run_message(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + shaping: MessagesShaping { + capabilities: AnthropicModelCapabilities { + supports_reasoning: true, + supports_adaptive_thinking: true, + ..AnthropicModelCapabilities::default() + }, + reasoning_auto_summary: true, + ..MessagesShaping::default() + }, + ..call + }, + json!({"thinking": thinking}), + )) + .await; + + assert_eq!(only_request(&upstream).await.json()["thinking"], expected); +} + +#[rstest] +#[case::reasoning_effort_on_an_adaptive_model( + AnthropicModelCapabilities { + supports_reasoning: true, + supports_adaptive_thinking: true, + supports_output_config: true, + effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() }, + ..AnthropicModelCapabilities::default() + }, + json!({"reasoning_effort": "high"}), + json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}}) +)] +#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens( + AnthropicModelCapabilities { + supports_reasoning: true, + ..AnthropicModelCapabilities::default() + }, + json!({"reasoning_effort": "high"}), + json!({"thinking": {"type": "enabled", "budget_tokens": 2999}}) +)] +#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget( + AnthropicModelCapabilities { + supports_reasoning: true, + ..AnthropicModelCapabilities::default() + }, + json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}), + json!({"thinking": {"type": "enabled", "budget_tokens": 2999}}) +)] +#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped( + AnthropicModelCapabilities::default(), + json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}), + json!({}) +)] +#[tokio::test] +async fn reasoning_is_translated_by_the_model_capabilities( + call: MessagesCall, + #[case] capabilities: AnthropicModelCapabilities, + #[case] fields: Value, + #[case] expected: Value, +) { + let upstream = upstream([message_response()]).await; + + run_message(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + shaping: MessagesShaping { + capabilities, + ..MessagesShaping::default() + }, + ..call + }, + [("max_tokens".to_string(), json!(3000))] + .into_iter() + .chain(object(fields)) + .collect(), + )) + .await; + + let sent = only_request(&upstream).await.json(); + assert_eq!(sent.get("reasoning_effort"), None); + assert_eq!(sent.get("temperature"), None); + let reasoning: Map = ["thinking", "output_config"] + .into_iter() + .filter_map(|name| Some((name.to_string(), sent.get(name)?.clone()))) + .collect(); + assert_eq!(Value::Object(reasoning), expected); +} + +#[rstest] +#[case::empty_text_blocks( + json!([{"role": "assistant", "content": [{"type": "text", "text": " "}, {"type": "text", "text": "kept"}]}]), + json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}]) +)] +#[case::provider_specific_fields( + json!([{"role": "assistant", "content": [{"type": "text", "text": "kept", "provider_specific_fields": {"x": 1}}]}]), + json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}]) +)] +#[case::unencrypted_web_search_results_become_text( + json!([{"role": "assistant", "content": [{ + "type": "web_search_tool_result", + "tool_use_id": "srvtoolu_1", + "content": [{"type": "web_search_result", "title": "T", "url": "https://e.x", "page_age": null}] + }]}]), + json!([{"role": "assistant", "content": [{"type": "text", "text": "Web search results:\n\nTitle: T\nURL: https://e.x"}]}]) +)] +#[tokio::test] +async fn replayed_history_is_cleaned_before_sending( + call: MessagesCall, + #[case] history: Value, + #[case] expected: Value, +) { + let upstream = upstream([message_response()]).await; + + run_message(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }, + json!({"messages": history}), + )) + .await; + + assert_eq!(only_request(&upstream).await.json()["messages"], expected); +} + +#[rstest] +#[tokio::test] +async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + + run_message(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }, + json!({"metadata": {"user_id": "u-1", "trace_id": "internal", "tags": ["a"]}}), + )) + .await; + + assert_eq!( + only_request(&upstream).await.json()["metadata"], + json!({"user_id": "u-1"}) + ); +} + +#[rstest] +#[case::numeric_user_id(json!({"metadata": {"user_id": 7}}))] +#[case::missing_max_tokens(json!({"max_tokens": null}))] +#[tokio::test] +async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fields: Value) { + let upstream = upstream([message_response()]).await; + + let error = run(with_fields( + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }, + fields, + )) + .await + .err() + .expect("the request is rejected"); + + assert!(error.is_request(), "{error:?}"); + assert!(received(&upstream).await.is_empty()); +} + +#[rstest] +#[tokio::test] +async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + + run_message(with_fields( + MessagesCall { + custom_llm_provider: Some("azure_ai".into()), + api_key: Some("sk-azure".into()), + api_base: Some(upstream.uri()), + ..call + }, + json!({ + "system": "top level", + "messages": [ + {"role": "system", "content": "from a message"}, + {"role": "user", "content": "hi"} + ] + }), + )) + .await; + + let sent = only_request(&upstream).await.json(); + assert_eq!( + sent["system"], + json!([ + {"type": "text", "text": "top level"}, + {"type": "text", "text": "from a message"} + ]) + ); + assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}])); +} + +#[rstest] +#[case::bare_model(MODEL, MODEL)] +#[case::one_prefix("anthropic/claude-sonnet-4-5", MODEL)] +#[case::doubled_prefix_loses_one_segment( + "anthropic/anthropic/claude-sonnet-4-5", + "anthropic/claude-sonnet-4-5" +)] +#[tokio::test] +async fn the_provider_prefix_is_stripped_exactly_once( + call: MessagesCall, + #[case] model: &str, + #[case] sent_model: &str, +) { + let upstream = upstream([message_response()]).await; + + run_message(MessagesCall { + model: model.into(), + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await; + + assert_eq!(only_request(&upstream).await.json()["model"], sent_model); +} diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs new file mode 100644 index 00000000000..133b7d2b162 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -0,0 +1,221 @@ +use litellm_core::messages::{messages, types::MessagesRequest}; +use litellm_http::transport::Error as TransportError; +use rstest::rstest; + +use super::*; + +#[rstest] +#[case::anthropic("anthropic")] +#[case::azure_ai("azure_ai")] +#[tokio::test] +async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider: &str) { + let upstream = upstream([message_response()]).await; + + let message = run_message(MessagesCall { + custom_llm_provider: Some(provider.into()), + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await; + + assert_eq!(message.id, "msg_1"); + assert_eq!(message.content, [json!({"type": "text", "text": "hi"})]); + assert_eq!(message.stop_reason.as_deref(), Some("end_turn")); +} + +/// A refusal and fields the route does not model come back exactly as the provider sent +/// them, since the Python side returns the raw message and the router decides what to do. +#[rstest] +#[tokio::test] +async fn the_message_passes_through_losslessly(call: MessagesCall) { + let upstream_body = json!({ + "id": "msg_2", + "type": "message", + "role": "assistant", + "model": MODEL, + "content": [ + {"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "q"}}, + {"type": "text", "text": "no", "citations": [{"type": "web_search_result_location", "url": "https://e.x"}]} + ], + "stop_reason": "refusal", + "stop_sequence": null, + "stop_details": {"type": "safeguard", "safeguard_types": ["dangerous_tool_use"]}, + "container": {"id": "container_1", "expires_at": "2026-01-01T00:00:00Z"}, + "context_management": {"applied_edits": []}, + "usage": {"input_tokens": 1, "output_tokens": 2, "server_tool_use": {"web_search_requests": 1}}, + "unknown_future_field": {"nested": true} + }); + let upstream = upstream([json_response(upstream_body.clone())]).await; + + let message = run_message(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await; + + assert_eq!(message.stop_reason.as_deref(), Some("refusal")); + assert_eq!(serde_json::to_value(&message).unwrap(), upstream_body); +} + +#[rstest] +#[tokio::test] +async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) { + let envelope = + json!({"type": "error", "error": {"type": "invalid_request_error", "message": "bad"}}); + let upstream = upstream([status_response(400, envelope.clone())]).await; + + let error = run(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await + .err() + .expect("upstream error propagates"); + + let Error::Transport(TransportError::Http { status, body }) = error else { + panic!("{error:?}"); + }; + assert_eq!(status, 400); + assert_eq!(serde_json::from_str::(&body).unwrap(), envelope); +} + +#[rstest] +#[tokio::test] +async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall) { + let long = "x".repeat(600); + let upstream = upstream([ResponseTemplate::new(500).set_body_string(long.clone())]).await; + + let error = run(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await + .err() + .expect("upstream error propagates"); + + assert_eq!( + error, + Error::Transport(TransportError::Http { + status: 500, + body: format!("{}... (truncated)", &long[..256]) + }) + ); +} + +#[rstest] +#[case::bad_request(400)] +#[case::unauthorized(401)] +#[case::rate_limited(429)] +#[case::server_error(500)] +#[case::overloaded(529)] +#[tokio::test] +async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case] status: u16) { + let upstream = + upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await; + + let error = run(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await + .err() + .expect("upstream error propagates"); + + assert_eq!( + error, + Error::Transport(TransportError::Http { + status, + body: "upstream said no".into() + }) + ); +} + +#[rstest] +#[case::not_json(ResponseTemplate::new(200).set_body_string("not json"))] +#[case::not_a_message(json_response(json!({"unexpected": true})))] +#[tokio::test] +async fn an_unreadable_success_body_is_an_invalid_response( + call: MessagesCall, + #[case] response: ResponseTemplate, +) { + let upstream = upstream([response]).await; + + let error = run(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }) + .await + .err() + .expect("an unreadable body fails"); + + assert!(error.is_response(), "{error:?}"); +} + +#[rstest] +#[tokio::test] +async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) { + let upstream = upstream([message_response().set_delay(Duration::from_secs(5))]).await; + + let error = run(MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + timeout: Some(Duration::from_millis(100)), + ..call + }) + .await + .err() + .expect("the call times out"); + + assert!(matches!(error, Error::Transport(_)), "{error:?}"); +} + +fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> { + MessagesRequest { + model: MODEL, + body, + api_key: Some("sk-ant"), + api_base: Some(api_base), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + provider_specific_header: None, + timeout: Some(Duration::from_secs(5)), + shaping: MessagesShaping::default(), + } +} + +#[tokio::test] +async fn the_facade_runs_the_route_in_process() { + let upstream = upstream([message_response()]).await; + let base = upstream.uri(); + + let message = messages(facade_request( + json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}), + &base, + )) + .await + .expect("messages request succeeds"); + + assert_eq!(message.id, "msg_1"); + assert_eq!( + only_request(&upstream).await.header("x-api-key"), + Some("sk-ant") + ); +} + +#[tokio::test] +async fn the_facade_rejects_a_body_that_is_not_an_object() { + let error = messages(facade_request(json!([]), UNREACHABLE_BASE)) + .await + .expect_err("a non-object body is rejected"); + + assert_eq!( + error, + Error::InvalidRequest("messages body must be an object".into()) + ); +} diff --git a/litellm-rust/crates/core/tests/messages/secrets.rs b/litellm-rust/crates/core/tests/messages/secrets.rs new file mode 100644 index 00000000000..55e510d00d3 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/secrets.rs @@ -0,0 +1,200 @@ +use rstest::rstest; + +use super::*; + +#[rstest] +#[case::anthropic( + "anthropic", + "ANTHROPIC_API_KEY", + "ANTHROPIC_BASE_URL", + "/v1/messages", + &["ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"] +)] +#[case::azure_ai( + "azure_ai", + "AZURE_API_KEY", + "AZURE_API_BASE", + "/anthropic/v1/messages", + &["AZURE_API_KEY", "AZURE_API_BASE"] +)] +#[tokio::test] +async fn the_credential_and_base_come_from_the_secret_source( + call: MessagesCall, + #[case] provider: &str, + #[case] key_name: &str, + #[case] base_name: &str, + #[case] path: &str, + #[case] looked_up: &[&str], +) { + let upstream = upstream([message_response()]).await; + let base = upstream.uri(); + let secrets = Arc::new(RecordingSecrets::new([ + (key_name, "sk-from-manager"), + (base_name, base.as_str()), + ])); + + let output = run_with( + secrets.clone(), + MessagesCall { + custom_llm_provider: Some(provider.into()), + ..call + }, + ) + .await + .expect("messages call succeeds"); + + assert!(matches!(output, MessagesOutput::Message(_))); + let request = only_request(&upstream).await; + assert_eq!(request.url.path(), path); + assert_eq!(request.header("x-api-key"), Some("sk-from-manager")); + assert_eq!(secrets.requested(), looked_up); +} + +#[rstest] +#[tokio::test] +async fn call_arguments_win_over_the_secret_source(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + let secrets = Arc::new(RecordingSecrets::new([ + ("ANTHROPIC_API_KEY", "sk-from-manager"), + ("ANTHROPIC_BASE_URL", UNREACHABLE_BASE), + ])); + + run_with( + secrets, + MessagesCall { + api_key: Some("sk-from-call".into()), + api_base: Some(upstream.uri()), + ..call + }, + ) + .await + .expect("messages call succeeds"); + + assert_eq!( + only_request(&upstream).await.header("x-api-key"), + Some("sk-from-call") + ); +} + +#[rstest] +#[tokio::test] +async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCall) { + let upstream = upstream([message_response()]).await; + + let error = run_with( + Arc::new(RecordingSecrets::failing()), + MessagesCall { + api_key: Some("sk".into()), + api_base: Some(upstream.uri()), + ..call + }, + ) + .await + .err() + .expect("a secret manager failure fails the call"); + + assert!( + matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)), + "{error:?}" + ); + assert!(received(&upstream).await.is_empty()); +} + +#[derive(Clone, Copy)] +enum Base { + Upstream, + Unreachable, + Blank, + Absent, +} + +fn base_value(base: Base, upstream: &str) -> Option { + match base { + Base::Upstream => Some(upstream.to_string()), + Base::Unreachable => Some(UNREACHABLE_BASE.to_string()), + Base::Blank => Some(" ".to_string()), + Base::Absent => None, + } +} + +#[rstest] +#[case::api_base_beats_base_url(Base::Upstream, Base::Unreachable)] +#[case::blank_api_base_falls_through_to_base_url(Base::Blank, Base::Upstream)] +#[case::base_url_alone(Base::Absent, Base::Upstream)] +#[tokio::test] +async fn the_anthropic_base_env_precedence_picks_the_upstream( + call: MessagesCall, + #[case] api_base: Base, + #[case] base_url: Base, +) { + let upstream = upstream([message_response()]).await; + let uri = upstream.uri(); + let values: Vec<(&str, &str)> = [ + ("ANTHROPIC_API_KEY", Some("sk-env".to_string())), + ("ANTHROPIC_API_BASE", base_value(api_base, &uri)), + ("ANTHROPIC_BASE_URL", base_value(base_url, &uri)), + ] + .iter() + .filter_map(|(name, value)| Some((*name, value.as_deref()?))) + .map(|(name, value)| (name, Box::leak(value.to_string().into_boxed_str()) as &str)) + .collect(); + + run_with(Arc::new(RecordingSecrets::new(values)), call) + .await + .expect("messages call reaches the upstream the precedence picks"); + + assert_eq!(only_request(&upstream).await.url.path(), "/v1/messages"); +} + +#[rstest] +#[case::auth_token_alone( + &[("ANTHROPIC_AUTH_TOKEN", "tok")], + ("authorization", "Bearer tok"), + "x-api-key" +)] +#[case::api_key_beats_the_auth_token( + &[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "tok")], + ("x-api-key", "sk-env"), + "authorization" +)] +#[tokio::test] +async fn the_auth_token_env_is_a_bearer_only_without_a_key( + call: MessagesCall, + #[case] values: &[(&str, &str)], + #[case] expected: (&str, &str), + #[case] absent: &str, +) { + let upstream = upstream([message_response()]).await; + + run_with( + Arc::new(RecordingSecrets::new(values.iter().copied())), + MessagesCall { + api_base: Some(upstream.uri()), + ..call + }, + ) + .await + .expect("messages call succeeds"); + + let request = only_request(&upstream).await; + let (name, value) = expected; + assert_eq!(request.header_values(name), [value]); + assert_eq!(request.header(absent), None); +} + +#[rstest] +#[tokio::test] +async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall) { + let error = run_with( + Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])), + MessagesCall { + custom_llm_provider: Some("azure_ai".into()), + ..call + }, + ) + .await + .err() + .expect("azure needs a base"); + + assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase)); +} diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs new file mode 100644 index 00000000000..c4be3127d66 --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -0,0 +1,266 @@ +use std::{convert::Infallible, sync::Mutex}; + +use bytes::Bytes; +use litellm_core::messages::route::{Messages, MessagesStreamHead}; +use litellm_host::host::{Demand, Host}; +use rstest::rstest; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, +}; + +use super::*; + +const UPSTREAM_HEADERS: [(&str, &str); 2] = [ + ("request-id", "req_upstream_123"), + ("anthropic-ratelimit-requests-remaining", "41"), +]; + +const SSE_BODY: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + +enum Seen { + Open(Vec<(String, String)>), + Deliver(Bytes), +} + +/// Projects like `LocalMessagesHost`, records every stream op in the order the route +/// performs it, and detaches after `detach_after` ops. +struct RecordingStreamHost { + call: LocalMessagesHost, + detach_after: usize, + seen: Mutex>, +} + +impl RecordingStreamHost { + fn new(call: MessagesCall, detach_after: usize) -> Self { + Self { + call: LocalMessagesHost::new(call), + detach_after, + seen: Mutex::new(Vec::new()), + } + } + + fn record(&self, op: Seen) -> Demand { + let mut seen = self.seen.lock().unwrap(); + seen.push(op); + match seen.len() < self.detach_after { + true => Demand::More, + false => Demand::Detached, + } + } +} + +impl Host for RecordingStreamHost { + async fn project(&self) -> Result { + self.call.project().await + } + + async fn custom_op(&self, op: Infallible) -> Result<(), Error> { + match op {} + } + + async fn open(&self, head: MessagesStreamHead) -> Result { + Ok(self.record(Seen::Open(head.headers))) + } + + async fn deliver(&self, chunk: Bytes) -> Result { + Ok(self.record(Seen::Deliver(chunk))) + } +} + +fn streaming(call: MessagesCall, api_base: String) -> MessagesCall { + let mut body = call.body.clone(); + body.insert("stream".into(), json!(true)); + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(api_base), + body, + ..call + } +} + +fn sse_response() -> ResponseTemplate { + UPSTREAM_HEADERS.iter().fold( + ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream"), + |response, (name, value)| response.insert_header(*name, *value), + ) +} + +async fn stream_through(host: &RecordingStreamHost) -> Result { + litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await +} + +#[rstest] +#[tokio::test] +async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: MessagesCall) { + let upstream = upstream([sse_response()]).await; + let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX); + + let outcome = stream_through(&host).await.expect("streamed call succeeds"); + + assert!(matches!(outcome, MessagesOutput::Streamed)); + let seen = host.seen.into_inner().unwrap(); + let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else { + panic!("the stream opens before any chunk is delivered"); + }; + let surfaced: Vec<(&str, &str)> = headers + .iter() + .filter(|(name, _)| { + UPSTREAM_HEADERS + .iter() + .any(|(upstream, _)| upstream == name) + }) + .map(|(name, value)| (name.as_str(), value.as_str())) + .collect(); + assert_eq!(surfaced, UPSTREAM_HEADERS); + let delivered: Vec = chunks + .iter() + .flat_map(|step| match step { + Seen::Deliver(chunk) => chunk.to_vec(), + Seen::Open(_) => panic!("the stream opens exactly once"), + }) + .collect(); + assert_eq!(delivered, SSE_BODY.as_bytes()); +} + +#[rstest] +#[case::at_open(1)] +#[case::after_the_first_chunk(2)] +#[tokio::test] +async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] detach_after: usize) { + let upstream = upstream([sse_response()]).await; + let host = RecordingStreamHost::new(streaming(call, upstream.uri()), detach_after); + + let outcome = stream_through(&host) + .await + .expect("a detached stream still completes"); + + assert!(matches!(outcome, MessagesOutput::Streamed)); + assert_eq!(host.seen.into_inner().unwrap().len(), detach_after); +} + +#[rstest] +#[case::text_body(ResponseTemplate::new(429).set_body_string("slow down"), "slow down")] +#[case::json_envelope( + status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})), + r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"# +)] +#[tokio::test] +async fn an_upstream_error_fails_the_call_without_opening_the_stream( + call: MessagesCall, + #[case] response: ResponseTemplate, + #[case] body: &str, +) { + let upstream = upstream([response]).await; + let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX); + + let error = stream_through(&host) + .await + .err() + .expect("upstream error propagates"); + + assert_eq!( + error, + Error::Transport(litellm_http::transport::Error::Http { + status: 429, + body: body.into() + }) + ); + assert!(host.seen.into_inner().unwrap().is_empty()); +} + +/// The native route relays bytes as they are. Python's synthetic `api_error` for a stream +/// that never reaches `message_stop` lives in its SSE wrapper, above this route. +#[rstest] +#[tokio::test] +async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: MessagesCall) { + const INCOMPLETE: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\n"; + let upstream = + upstream([ResponseTemplate::new(200).set_body_raw(INCOMPLETE, "text/event-stream")]).await; + let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX); + + stream_through(&host).await.expect("streamed call succeeds"); + + let delivered: Vec = host + .seen + .into_inner() + .unwrap() + .iter() + .flat_map(|step| match step { + Seen::Deliver(chunk) => chunk.to_vec(), + Seen::Open(_) => Vec::new(), + }) + .collect(); + assert_eq!(delivered, INCOMPLETE.as_bytes()); +} + +/// Serves one SSE chunk and then holds the connection open without ever finishing. +async fn stalling_upstream() -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = vec![0; 4096]; + let _ = socket.read(&mut request).await; + socket + .write_all( + b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n\ + 1f\r\nevent: message_start\ndata: {}\n\n\r\n", + ) + .await + .unwrap(); + std::future::pending::<()>().await; + }); + base +} + +#[rstest] +#[tokio::test] +async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) { + let base = stalling_upstream().await; + let host = RecordingStreamHost::new( + MessagesCall { + timeout: Some(Duration::from_millis(300)), + ..streaming(call, base) + }, + usize::MAX, + ); + + let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host)) + .await + .expect("the stalled stream gives up within the timeout") + .err() + .expect("a stalled body fails the call"); + + assert!(matches!(error, Error::Transport(_)), "{error:?}"); + let seen = host.seen.into_inner().unwrap(); + assert!( + matches!(seen.as_slice(), [Seen::Open(_), Seen::Deliver(chunk)] if chunk.as_ref() == b"event: message_start\ndata: {}\n\n"), + "the chunk before the stall reached the caller, saw {} ops", + seen.len() + ); +} + +#[rstest] +#[tokio::test] +async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) { + let upstream = upstream([sse_response()]).await; + let host = RecordingStreamHost::new( + MessagesCall { + custom_llm_provider: Some("azure_ai".into()), + ..streaming(call, upstream.uri()) + }, + usize::MAX, + ); + + let error = stream_through(&host) + .await + .err() + .expect("azure streaming is refused"); + + assert_eq!( + error, + Error::Unsupported("streaming messages for this provider") + ); + assert!(received(&upstream).await.is_empty()); +} diff --git a/litellm-rust/crates/core/tests/ocr/aws_textract.rs b/litellm-rust/crates/core/tests/ocr/aws_textract.rs new file mode 100644 index 00000000000..790e16a95ec --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/aws_textract.rs @@ -0,0 +1,173 @@ +use std::{collections::BTreeMap, time::SystemTime}; + +use litellm_auth_aws::{Credentials, aws_signature_headers, sign_post}; +use rstest::rstest; +use time::{PrimitiveDateTime, format_description}; +use wiremock::Request; + +use super::*; + +const ACCESS_KEY_ID: &str = "AKIDEXAMPLE"; +const SECRET_ACCESS_KEY: &str = "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"; +const DETECT: &str = "aws_textract/detect-document-text"; +const ANALYZE: &str = "aws_textract/analyze-document"; + +fn textract_request(model: &str, base: &str) -> LiteLLMOcrRequest { + ocr_request_with_document( + model, + &format!("{base}/"), + json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}), + json!({ + "aws_access_key_id": ACCESS_KEY_ID, + "aws_secret_access_key": SECRET_ACCESS_KEY, + "aws_region_name": "eu-west-1" + }), + ) +} + +fn textract_response() -> ResponseTemplate { + json_response(json!({ + "DocumentMetadata": {"Pages": 1}, + "Blocks": [{"BlockType": "PAGE"}, {"BlockType": "LINE", "Text": "Invoice 12345"}] + })) +} + +/// Recomputes SigV4 over the request the upstream received, at the time the client claimed. +fn expected_authorization(url: &str, sent: &Request) -> String { + let format = + format_description::parse_borrowed::<2>("[year][month][day]T[hour][minute][second]Z") + .unwrap(); + let signed_at: SystemTime = + PrimitiveDateTime::parse(sent.header("x-amz-date").unwrap(), &format) + .unwrap() + .assume_utc() + .into(); + let headers: BTreeMap = ["content-type", "x-amz-target"] + .into_iter() + .map(|name| (name.to_string(), sent.header(name).unwrap().to_string())) + .collect(); + sign_post( + url, + &sent.body, + &aws_signature_headers(&headers), + "eu-west-1", + "textract", + &Credentials::new(ACCESS_KEY_ID, SECRET_ACCESS_KEY, None, None, "test"), + signed_at, + ) + .unwrap()["Authorization"] + .clone() +} + +/// The recorded URL names wiremock's host, not the address the client signed for. +fn assert_signed(upstream: &MockServer, sent: &Request) { + let url = format!("{}/", upstream.uri()); + assert_eq!( + sent.header("authorization"), + Some(expected_authorization(&url, sent).as_str()) + ); +} + +#[tokio::test] +async fn detect_document_text_is_signed_and_lines_become_the_page() { + let upstream = upstream([textract_response()]).await; + + let response = perform_with(LocalOcrHost::new(textract_request(DETECT, &upstream.uri()))) + .await + .unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!( + sent.header("x-amz-target"), + Some("Textract.DetectDocumentText") + ); + assert_eq!( + sent.header("content-type"), + Some("application/x-amz-json-1.1") + ); + assert_eq!(sent.json(), json!({"Document": {"Bytes": "b3JpZ2luYWw="}})); + assert_signed(&upstream, &sent); + assert_eq!(response.pages[0].markdown, "Invoice 12345"); + assert_eq!(response.usage_info.unwrap().pages_processed, Some(1)); +} + +#[tokio::test] +async fn a_body_rewritten_by_before_send_is_what_gets_signed_and_sent() { + let upstream = upstream([textract_response()]).await; + let host = LocalOcrHost::new(textract_request(DETECT, &upstream.uri())).with_before_send( + |mut wire, _| { + assert!( + !wire + .headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("authorization")), + "the hook ran after signing" + ); + wire.body["Document"]["Bytes"] = Value::from("cmVkYWN0ZWQ="); + Ok(wire) + }, + ); + + perform_with(host).await.unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!(sent.json(), json!({"Document": {"Bytes": "cmVkYWN0ZWQ="}})); + assert_signed(&upstream, &sent); +} + +#[tokio::test] +async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() { + let upstream = upstream([json_response(json!({ + "DocumentMetadata": {"Pages": 1}, + "Blocks": [ + {"Id": "l1", "BlockType": "LINE", "Text": "Quarterly Report"}, + {"Id": "t", "BlockType": "LAYOUT_TITLE", + "Relationships": [{"Type": "CHILD", "Ids": ["l1"]}]} + ] + }))]) + .await; + + let response = perform_with(LocalOcrHost::new(textract_request( + ANALYZE, + &upstream.uri(), + ))) + .await + .unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!( + sent.header("x-amz-target"), + Some("Textract.AnalyzeDocument") + ); + assert_eq!(sent.json()["FeatureTypes"], json!(["LAYOUT", "TABLES"])); + assert_signed(&upstream, &sent); + assert_eq!(response.pages[0].markdown, "# Quarterly Report"); +} + +#[rstest] +#[case::detect(DETECT)] +#[case::analyze(ANALYZE)] +#[tokio::test] +async fn a_multi_page_rejection_reaches_the_caller_with_the_single_page_limit(#[case] model: &str) { + let upstream = upstream([status_response( + 400, + json!({ + "__type": "UnsupportedDocumentException", + "Message": "Request has unsupported document format" + }), + )]) + .await; + + let error = perform_with(LocalOcrHost::new(textract_request(model, &upstream.uri()))) + .await + .unwrap_err(); + + let Error::Provider { status, body, .. } = error else { + panic!("expected a provider error, got {error:?}"); + }; + assert_eq!(status, 400); + assert!( + body.contains("multi-page documents are not supported"), + "{body}" + ); +} diff --git a/litellm-rust/crates/core/tests/ocr/azure_ai.rs b/litellm-rust/crates/core/tests/ocr/azure_ai.rs new file mode 100644 index 00000000000..0d6ba024c15 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/azure_ai.rs @@ -0,0 +1,270 @@ +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; + +use litellm_auth::{ + ResolvedCredential, SecretValue, TokenFuture, TokenProvider, TokenProviderHandle, +}; +use rstest::rstest; + +use super::*; + +#[derive(Debug)] +struct CountingToken { + token: fn(usize) -> String, + calls: AtomicUsize, +} + +impl CountingToken { + fn new(token: fn(usize) -> String) -> Arc { + Arc::new(Self { + token, + calls: AtomicUsize::new(0), + }) + } + + fn calls(&self) -> usize { + self.calls.load(Ordering::SeqCst) + } +} + +impl TokenProvider for CountingToken { + fn acquire(&self) -> TokenFuture<'_> { + let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1; + let token = SecretValue::new((self.token)(call)); + Box::pin(async move { + Ok(ResolvedCredential::AccessToken { + token, + expires_on: None, + }) + }) + } +} + +fn numbered_token(call: usize) -> String { + format!("callback-{call}") +} + +fn azure_request( + provider: &Arc, + api_base: Option<&str>, + api_key: Option<&str>, + extra_headers: Value, + optional_params: Value, +) -> LiteLLMOcrRequest { + let wire = serde_json::from_value(json!({ + "model": "azure_ai/mistral-ocr-latest", + "document": {"type": "document_url", "document_url": INLINE_PDF}, + "api_key": api_key, + "api_base": api_base, + "custom_llm_provider": null, + "extra_headers": extra_headers, + "optional_params": optional_params, + "timeout_seconds": 2.0 + })) + .unwrap(); + let mut request = decode_request(wire).unwrap(); + request.azure_ad_token_provider = Some(TokenProviderHandle::new(provider.clone())); + request +} + +fn ocr_page() -> ResponseTemplate { + json_response(json!({"pages": [{"index": 0, "markdown": "hello"}]})) +} + +#[tokio::test] +async fn mistral_on_azure_sends_the_prepared_bearer_and_the_mistral_body() { + let upstream = upstream([json_response(json!({ + "pages": [{"index": 0, "markdown": "hello"}], + "usage_info": {"pages_processed": 1} + }))]) + .await; + let request = with_headers( + without_api_key(ocr_request( + "azure_ai/model", + &upstream.uri(), + json!({"include_image_base64": true}), + )), + &[("Authorization", "Bearer python-prepared-token")], + ); + + let result = perform(request).await.unwrap(); + + assert_eq!(result.pages[0].markdown, "hello"); + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/providers/mistral/azure/ocr"); + assert_eq!( + sent.header("authorization"), + Some("Bearer python-prepared-token") + ); + assert_eq!( + sent.json(), + json!({ + "model": "model", + "document": {"type": "document_url", "document_url": INLINE_PDF}, + "include_image_base64": true + }) + ); +} + +#[tokio::test] +async fn a_static_entra_token_becomes_the_bearer() { + let upstream = upstream([pages_response()]).await; + let request = without_api_key(ocr_request( + "azure_ai/model", + &upstream.uri(), + json!({"azure_ad_token": "rust-owned-token"}), + )); + + perform(request).await.unwrap(); + + assert_eq!( + only_request(&upstream).await.header("authorization"), + Some("Bearer rust-owned-token") + ); +} + +#[tokio::test] +async fn a_guardrail_that_swaps_in_a_remote_document_is_rejected() { + let host = LocalOcrHost::new(ocr_request("azure_ai/model", UNREACHABLE_BASE, json!({}))) + .with_before_send(|mut wire, _| { + wire.body["document"] = json!({ + "type": "document_url", + "document_url": "https://example.com/not-inline.pdf" + }); + Ok(wire) + }); + + let error = perform_with(host).await.unwrap_err(); + + assert!(error.to_string().contains("data URI"), "{error}"); +} + +#[tokio::test] +async fn the_token_provider_is_the_bearer_and_is_acquired_for_each_request() { + let provider = CountingToken::new(numbered_token); + let upstream = upstream([ocr_page(), ocr_page()]).await; + let base = upstream.uri(); + + for _ in 0..2 { + perform(azure_request( + &provider, + Some(&base), + None, + Value::Null, + json!({}), + )) + .await + .unwrap(); + } + + assert_eq!(provider.calls(), 2); + let authorizations: Vec = received(&upstream) + .await + .iter() + .map(|request| { + request + .header("authorization") + .unwrap_or_default() + .to_string() + }) + .collect(); + assert_eq!(authorizations, ["Bearer callback-1", "Bearer callback-2"]); +} + +#[rstest] +#[case::api_key_skips_provider(Some("resource-key"), Value::Null, json!({}), "Bearer resource-key", 0)] +#[case::provider_beats_static_token( + None, + Value::Null, + json!({"azure_ad_token": "static-token"}), + "Bearer callback-1", + 1 +)] +#[case::header_wins_on_the_wire_but_provider_still_runs( + None, + json!({"Authorization": "Bearer override"}), + json!({}), + "Bearer override", + 1 +)] +#[tokio::test] +async fn credential_precedence( + #[case] api_key: Option<&str>, + #[case] extra_headers: Value, + #[case] optional_params: Value, + #[case] expected_authorization: &str, + #[case] expected_calls: usize, +) { + let provider = CountingToken::new(numbered_token); + let upstream = upstream([ocr_page()]).await; + + perform(azure_request( + &provider, + Some(&upstream.uri()), + api_key, + extra_headers, + optional_params, + )) + .await + .unwrap(); + + assert_eq!(provider.calls(), expected_calls); + assert_eq!( + only_request(&upstream).await.header_values("authorization"), + [expected_authorization] + ); +} + +#[rstest] +#[case::missing_api_base( + false, + json!({}), + numbered_token, + |error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase { + provider: "Azure AI", + environment_variable: "AZURE_AI_API_BASE", + })), + 0 +)] +#[case::unsupported_oidc_reference( + true, + json!({"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"}), + numbered_token, + |error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)), + 0 +)] +#[case::empty_provider_token_ignores_static_token( + true, + json!({"azure_ad_token": "static-token"}), + |_| String::new(), + |error: &Error| matches!(error, Error::MissingAzureAiCredentials), + 1 +)] +#[tokio::test] +async fn credential_failures_send_no_provider_request( + #[case] with_api_base: bool, + #[case] optional_params: Value, + #[case] token: fn(usize) -> String, + #[case] expected: fn(&Error) -> bool, + #[case] expected_calls: usize, +) { + let provider = CountingToken::new(token); + let upstream = upstream([ocr_page()]).await; + let base = upstream.uri(); + + let error = perform(azure_request( + &provider, + with_api_base.then_some(base.as_str()), + None, + Value::Null, + optional_params, + )) + .await + .unwrap_err(); + + assert!(expected(&error), "unexpected error: {error:?}"); + assert_eq!(provider.calls(), expected_calls); + assert!(received(&upstream).await.is_empty()); +} diff --git a/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs b/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs new file mode 100644 index 00000000000..1921176d158 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/azure_document_intelligence.rs @@ -0,0 +1,441 @@ +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_host::event::{CallEvent, MachineEvent}; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use rstest::rstest; + +use super::*; + +const MODEL: &str = "azure_ai/doc-intelligence/prebuilt-read"; + +fn read_request(base: &str, options: Value) -> LiteLLMOcrRequest { + ocr_request(MODEL, base, options) +} + +#[tokio::test] +async fn pages_features_and_extra_options_map_to_the_analyze_call() { + let upstream = upstream([json_response(json!({ + "status": "succeeded", + "analyzeResult": {"pages": []} + }))]) + .await; + let request = read_request( + &upstream.uri(), + json!({ + "pages": [2, 0, 0, 1], + "features": ["keyValuePairs", "languages"], + "future_option": {"nested": null}, + "extra_body": {"provider_option": false} + }), + ) + .with_document( + document( + json!({"type": "document_url", "document_url": "https://example.com/document.pdf"}), + ) + .into(), + ); + + perform(request).await.unwrap(); + + let sent = only_request(&upstream).await; + assert!( + sent.url.path().ends_with("/prebuilt-read:analyze"), + "{}", + sent.url + ); + assert_eq!(sent.query("pages").as_deref(), Some("1,2,3")); + assert_eq!( + sent.query("features").as_deref(), + Some("keyValuePairs,languages") + ); + assert_eq!( + sent.json(), + json!({ + "urlSource": "https://example.com/document.pdf", + "future_option": {"nested": null}, + "provider_option": false + }) + ); +} + +#[rstest] +#[case(json!({"pages": [true]}), Error::Pages("expected only integers or only strings".into()))] +#[case(json!({"pages": [1, "2"]}), Error::Pages("expected only integers or only strings".into()))] +#[case(json!({"pages": [-1]}), Error::Pages("negative page index".into()))] +#[case(json!({"pages": "1&&features=bad"}), Error::Pages("invalid native page range".into()))] +#[case(json!({"features": "languages&pages=1"}), Error::Features)] +#[case(json!({"req_format": "azure"}), Error::RequestFormat)] +#[tokio::test] +async fn invalid_pages_features_and_format_are_rejected_before_sending( + #[case] options: Value, + #[case] expected: Error, +) { + let upstream = upstream([json_response(json!({}))]).await; + + let result = match decode_request(wire( + MODEL, + &upstream.uri(), + json!({"type": "document_url", "document_url": "https://example.com/a.pdf"}), + options.clone(), + )) { + Ok(request) => perform(request).await, + Err(error) => Err(error), + }; + + assert!( + received(&upstream).await.is_empty(), + "sent invalid options: {options}" + ); + let error = result.unwrap_err(); + assert_eq!( + std::mem::discriminant(&error), + std::mem::discriminant(&expected) + ); + assert_eq!(error.http_status_code(), Some(400)); + assert_eq!(error.to_string(), expected.to_string()); +} + +#[rstest] +#[case::no_options(json!({}))] +#[case::litellm_format(json!({"req_format": "litellm"}))] +#[tokio::test] +async fn an_inline_document_is_sent_as_base64_and_only_page_text_is_kept(#[case] options: Value) { + let upstream = upstream([json_response(json!({ + "status": "succeeded", + "analyzeResult": {"pages": [{"pageNumber": 1, "lines": [{"content": "hello"}]}]} + }))]) + .await; + + let response = perform(read_request(&upstream.uri(), options)) + .await + .unwrap(); + + assert_eq!(response.pages.len(), 1); + assert_eq!(response.pages[0].index, 0); + assert_eq!(response.pages[0].markdown, "hello"); + assert_eq!(response.provider_native_response, None); + let serialized = response.into_json(); + for field in ["content", "tables", "keyValuePairs"] { + assert_eq!(serialized.get(field), Some(&Value::Null), "{field}"); + } + let sent = only_request(&upstream).await; + for field in ["pages", "features", "req_format"] { + assert_eq!(sent.query(field), None, "{field}"); + } + assert_eq!(sent.json(), json!({"base64Source": "YWJj"})); +} + +#[tokio::test] +async fn native_format_normalizes_pages_and_keeps_the_provider_response() { + let operation = json!({ + "status": "succeeded", + "operationExtension": 42, + "analyzeResult": { + "content": "A\n\nB", + "tables": [{"cells": []}], + "keyValuePairs": [{"key": {"content": "A"}}], + "pages": [{ + "pageNumber": "2", + "width": "8.5", + "height": 11, + "unit": "inch", + "lines": [{"content": "A"}, {"content": null}, {"content": "B"}] + }] + } + }); + let upstream = upstream([json_response(operation.clone())]).await; + + let result = perform(read_request( + &upstream.uri(), + json!({"req_format": "native"}), + )) + .await + .unwrap(); + + assert_eq!(result.pages[0].index, 1); + assert_eq!(result.pages[0].markdown, "A\n\nB"); + assert_eq!( + serde_json::to_value(&result.pages[0].dimensions).unwrap(), + json!({"width": 816, "height": 1056, "dpi": 96}) + ); + assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1)); + let serialized = result.clone().into_json(); + assert_eq!(serialized["content"], "A\n\nB"); + assert_eq!(serialized["tables"], json!([{"cells": []}])); + assert_eq!( + serialized["keyValuePairs"], + json!([{"key": {"content": "A"}}]) + ); + assert!(serialized.get("key_value_pairs").is_none()); + assert_eq!( + result.provider_native_response.map(Value::Object), + Some(operation) + ); +} + +#[tokio::test] +async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() { + let upstream = upstream([json_response(json!({ + "status": "succeeded", + "analyzeResult": {"pages": [{"pageNumber": 1, "width": 8.5, "height": 11, "unit": "inch"}]} + }))]) + .await; + let client = ocr_client().with_settings(OcrSettings { + document_intelligence_api_version: "2099-01-01".into(), + document_intelligence_dpi: 72, + ..OcrSettings::default() + }); + + let result = + litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))) + .await + .unwrap(); + + assert_eq!( + only_request(&upstream) + .await + .query("api-version") + .as_deref(), + Some("2099-01-01") + ); + assert_eq!( + serde_json::to_value(&result.pages[0].dimensions).unwrap(), + json!({"width": 612, "height": 792, "dpi": 72}) + ); +} + +#[tokio::test] +async fn an_accepted_response_polls_to_success_with_only_credentials() { + let operation = json!({"status": "succeeded", "analyzeResult": {"pages": []}}); + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({})), + json_response(json!({"status": "running"})).insert_header("Retry-After", "0"), + json_response(operation.clone()), + ], + ) + .await; + let request = with_headers( + read_request(&upstream.uri(), json!({"req_format": "native"})), + &[("X-Trace", "initial-only")], + ); + + let result = perform(request).await.unwrap(); + + assert_eq!( + result.provider_native_response.map(Value::Object), + Some(operation) + ); + let requests = received(&upstream).await; + assert_eq!(requests.len(), 3); + assert_eq!(requests[0].header("x-trace"), Some("initial-only")); + for poll in &requests[1..] { + assert_eq!(poll.method.as_str(), "GET"); + assert_eq!(poll.url.path(), "/operation"); + assert_eq!(poll.header("x-trace"), None); + assert_eq!(poll.header("ocp-apim-subscription-key"), Some("test-key")); + } +} + +#[tokio::test] +async fn polling_forwards_bearer_credentials() { + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({})), + json_response(json!({"status": "succeeded"})), + ], + ) + .await; + let request = with_headers( + without_api_key(read_request(&upstream.uri(), json!({}))), + &[("Authorization", "Bearer token")], + ); + + perform(request).await.unwrap(); + + assert_eq!( + received(&upstream).await[1].header("authorization"), + Some("Bearer token") + ); +} + +#[tokio::test] +async fn response_received_fires_for_the_submission_and_the_completed_poll() { + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({"submitted": true})), + json_response(json!({"status": "succeeded"})), + ], + ) + .await; + let observed = Arc::new(Mutex::new(Vec::new())); + let recorder = observed.clone(); + let host = + LocalOcrHost::new(read_request(&upstream.uri(), json!({}))).with_observer(move |event| { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { + recorder.lock().unwrap().push(raw.body.clone()); + } + }); + + perform_with(host).await.unwrap(); + + assert_eq!(received(&upstream).await.len(), 2); + assert_eq!( + *observed.lock().unwrap(), + [r#"{"submitted":true}"#, r#"{"status":"succeeded"}"#] + ); +} + +#[tokio::test] +async fn polling_does_not_follow_redirects() { + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({})), + ResponseTemplate::new(302) + .insert_header("Location", format!("{}/redirected", upstream.uri())), + json_response(json!({"status": "succeeded"})), + ], + ) + .await; + + let error = perform(read_request(&upstream.uri(), json!({}))) + .await + .unwrap_err(); + + assert!(error.to_string().contains("status 302"), "{error}"); + assert_eq!(received(&upstream).await.len(), 2); +} + +#[tokio::test] +async fn a_failed_operation_is_an_error() { + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({})), + json_response(json!({"status": "failed"})), + ], + ) + .await; + + let error = perform(read_request(&upstream.uri(), json!({}))) + .await + .unwrap_err(); + + assert!(error.to_string().contains("status failed"), "{error}"); +} + +#[tokio::test] +async fn the_polling_deadline_bounds_the_retry_delay() { + let upstream = MockServer::start().await; + respond_in_order( + &upstream, + [ + accepted(&upstream, json!({})), + json_response(json!({"status": "notStarted"})).insert_header("Retry-After", "9999"), + ], + ) + .await; + let client = ocr_client().with_settings(OcrSettings { + poll_timeout: Duration::from_millis(100), + ..OcrSettings::default() + }); + + let error = tokio::time::timeout( + Duration::from_secs(1), + litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))), + ) + .await + .expect("the deadline cuts the retry delay short") + .unwrap_err(); + + assert!(error.to_string().contains("timed out"), "{error}"); +} + +#[rstest] +#[case::null_pages(json!({"pages": null}), "pages")] +#[case::null_page(json!({"pages": [null]}), "pages[0]")] +#[case::null_lines(json!({"pages": [{"lines": null}]}), "lines")] +#[case::bad_width(json!({"pages": [{"width": "bad"}]}), "width")] +#[tokio::test] +async fn malformed_provider_pages_report_the_response_path( + #[case] analysis: Value, + #[case] path: &str, +) { + let upstream = upstream([json_response(json!({ + "status": "succeeded", + "analyzeResult": analysis + }))]) + .await; + + let error = perform(read_request(&upstream.uri(), json!({}))) + .await + .unwrap_err(); + + assert!(error.to_string().contains(path), "{error}"); +} + +#[rstest] +#[case::missing(None)] +#[case::relative(Some("/relative"))] +#[case::cross_origin(Some("http://example.com/operation"))] +#[case::with_userinfo(Some("http://user:password@127.0.0.1/operation"))] +#[tokio::test] +async fn an_unusable_operation_location_is_rejected(#[case] location: Option<&str>) { + let response = location + .into_iter() + .fold(ResponseTemplate::new(202), |response, location| { + response.insert_header("Operation-Location", location) + }); + let upstream = upstream([response]).await; + + let error = perform(read_request(&upstream.uri(), json!({}))) + .await + .unwrap_err(); + + assert!(error.to_string().contains("operation-location"), "{error}"); + assert_eq!(received(&upstream).await.len(), 1); +} + +#[tokio::test] +async fn the_model_id_is_percent_encoded() { + let upstream = upstream([json_response(json!({"status": "succeeded"}))]).await; + + perform(ocr_request( + "azure_ai/doc-intelligence/a ?#é", + &upstream.uri(), + json!({}), + )) + .await + .unwrap(); + + let sent = only_request(&upstream).await; + assert!( + sent.url.path().ends_with("/a%20%3F%23%C3%A9:analyze"), + "{}", + sent.url + ); +} + +#[rstest] +#[case::dot("azure_ai/doc-intelligence/.")] +#[case::dot_dot("azure_ai/doc-intelligence/..")] +#[tokio::test] +async fn dot_segment_model_ids_are_rejected(#[case] model: &str) { + let error = perform(ocr_request(model, UNREACHABLE_BASE, json!({}))) + .await + .unwrap_err(); + + assert!(error.to_string().contains("dot segment"), "{error}"); +} diff --git a/litellm-rust/crates/core/tests/ocr/cohere.rs b/litellm-rust/crates/core/tests/ocr/cohere.rs new file mode 100644 index 00000000000..007aa49a2fd --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/cohere.rs @@ -0,0 +1,42 @@ +use rstest::rstest; + +use super::*; + +#[rstest] +#[case::cohere("cohere/parse-v5.0", "/v2/parse")] +#[case::azure_ai("azure_ai/Cohere-parse-v5.0", "/providers/cohere/v2/parse")] +#[tokio::test] +async fn an_image_goes_to_the_parse_endpoint_with_the_bearer_key( + #[case] model: &str, + #[case] path: &str, +) { + let upstream = upstream([pages_response()]).await; + let request = ocr_request_with_document( + model, + &upstream.uri(), + json!({"type": "image_url", "image_url": "data:image/png;base64,YWJj"}), + json!({}), + ); + + perform(request).await.unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!(sent.method.as_str(), "POST"); + assert_eq!(sent.url.path(), path); + assert_eq!(sent.header("authorization"), Some("Bearer test-key")); +} + +#[rstest] +#[tokio::test] +async fn a_non_image_document_is_rejected_before_sending( + #[values("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")] model: &str, +) { + let upstream = upstream([pages_response()]).await; + + let error = perform(ocr_request(model, &upstream.uri(), json!({}))) + .await + .unwrap_err(); + + assert!(matches!(error, Error::CohereImageOnly), "{error:?}"); + assert!(received(&upstream).await.is_empty()); +} diff --git a/litellm-rust/crates/core/tests/ocr/documents.rs b/litellm-rust/crates/core/tests/ocr/documents.rs new file mode 100644 index 00000000000..e29ff3e9ee9 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/documents.rs @@ -0,0 +1,182 @@ +use base64::Engine; +use litellm_core::ocr::types::OcrDocumentInput; +use litellm_host::event::WireRequest; +use rstest::rstest; +use wiremock::{Mock, matchers::any}; + +use super::*; + +const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document"; +const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ="; + +#[derive(Clone, Copy, Debug)] +enum Route { + Mistral, + AzureAi, + VertexMistral, + AzureCohereParse, + Cohere, +} + +impl Route { + fn model(self) -> &'static str { + match self { + Self::Mistral => "mistral/model", + Self::AzureAi => "azure_ai/model", + Self::VertexMistral => "vertex_ai/mistral-ocr-maas", + Self::AzureCohereParse => "azure_ai/cohere-parse", + Self::Cohere => "cohere/model", + } + } + + fn document_type(self) -> &'static str { + match self { + Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url", + Self::AzureCohereParse | Self::Cohere => "image_url", + } + } + + fn options(self) -> Value { + match self { + Self::Mistral | Self::AzureAi => json!({"pages": [0]}), + Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}), + Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}), + } + } +} + +/// What the host does to the wire request in `before_send`. +#[derive(Clone, Copy, Debug)] +enum Guardrail { + Detached, + ReplacesDocument, +} + +impl Guardrail { + fn before_send(self, wire: WireRequest) -> WireRequest { + let Value::Object(fields) = wire.body else { + return wire; + }; + let body = fields + .into_iter() + .map(|(name, value)| match self { + Self::ReplacesDocument if name == "document" => { + let document_type = value["type"].clone(); + let key = document_type.as_str().unwrap_or_default().to_string(); + (name, json!({"type": document_type, key: REPLACED_DOCUMENT})) + } + Self::Detached | Self::ReplacesDocument => (name, value), + }) + .collect(); + WireRequest { + body: Value::Object(body), + ..wire + } + } +} + +/// Serves [`SERVED_DOCUMENT`] as `image/png` to every request. +async fn document_server() -> MockServer { + let server = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(200).set_body_raw(SERVED_DOCUMENT, "image/png")) + .mount(&server) + .await; + server +} + +/// Sends a remote document through `route` and returns the document the provider saw. +async fn provider_document(route: Route, guardrail: Guardrail) -> Value { + let documents = document_server().await; + let upstream = upstream([pages_response()]).await; + let document_type = route.document_type(); + let request = ocr_request_with_document( + route.model(), + &upstream.uri(), + json!({"type": document_type, document_type: format!("{}/scan.png", documents.uri())}), + route.options(), + ); + let host = + LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(guardrail.before_send(wire))); + + perform_with(host).await.unwrap(); + + only_request(&upstream).await.json()["document"][document_type].clone() +} + +#[rstest] +#[case::azure_ai(Route::AzureAi)] +#[case::vertex_mistral(Route::VertexMistral)] +#[case::azure_cohere_parse(Route::AzureCohereParse)] +#[tokio::test] +async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) { + let expected = format!( + "data:image/png;base64,{}", + base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT) + ); + + assert_eq!( + provider_document(route, Guardrail::Detached).await, + expected + ); +} + +#[rstest] +#[tokio::test] +async fn a_document_replaced_by_the_host_reaches_the_provider( + #[values( + Route::Mistral, + Route::AzureAi, + Route::VertexMistral, + Route::AzureCohereParse, + Route::Cohere + )] + route: Route, +) { + assert_eq!( + provider_document(route, Guardrail::ReplacesDocument).await, + REPLACED_DOCUMENT + ); +} + +#[tokio::test] +async fn an_empty_byte_document_fails_before_sending() { + let upstream = upstream([pages_response()]).await; + let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document( + OcrDocumentInput::Bytes { + bytes: Default::default(), + file_name: None, + mime_type: None, + }, + ); + + let error = perform(request).await.unwrap_err(); + + assert!(matches!(error, Error::EmptyFile), "{error:?}"); + assert!(received(&upstream).await.is_empty()); +} + +#[tokio::test] +async fn a_missing_path_document_fails_before_sending() { + let upstream = upstream([pages_response()]).await; + let path = + std::env::temp_dir().join(format!("litellm-ocr-missing-{}.png", rand::random::())); + let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document( + OcrDocumentInput::Path { + path: path.clone(), + mime_type: None, + }, + ); + + let error = perform(request).await.unwrap_err(); + + assert!( + matches!( + &error, + Error::FileRead { path: failed, source } + if *failed == path && source.kind() == std::io::ErrorKind::NotFound + ), + "{error:?}" + ); + assert!(received(&upstream).await.is_empty()); +} diff --git a/litellm-rust/crates/core/tests/ocr/lifecycle.rs b/litellm-rust/crates/core/tests/ocr/lifecycle.rs new file mode 100644 index 00000000000..65e64cce79b --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/lifecycle.rs @@ -0,0 +1,269 @@ +use std::sync::{Arc, Mutex}; + +use litellm_core::ocr::{ + route::{Ocr, OcrOp, OcrProjection, ocr_machine}, + types::OcrDocumentInput, +}; +use litellm_host::{ + event::{CallEvent, MachineEvent, RequestContext, WireRequest}, + host::Host, +}; +use rstest::rstest; + +use super::*; + +pub(crate) fn event_name(event: &CallEvent) -> &'static str { + match event { + CallEvent::Started { .. } => "started", + CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response", + CallEvent::Succeeded { .. } => "success", + CallEvent::Failed { .. } => "failure", + } +} + +fn recording_host( + request: LiteLLMOcrRequest, + events: Arc>>, + block: bool, +) -> LocalOcrHost { + let before_send_events = events.clone(); + LocalOcrHost::new(request) + .with_before_send(move |wire, _| { + before_send_events.lock().unwrap().push("before_send"); + match block { + true => Err(Error::InvalidRequest("blocked".into())), + false => Ok(wire), + } + }) + .with_observer(move |event| events.lock().unwrap().push(event_name(event))) +} + +#[tokio::test] +async fn hooks_run_in_order_and_one_success_is_emitted() { + let upstream = upstream([pages_response()]).await; + let events = Arc::new(Mutex::new(Vec::new())); + + perform_with(recording_host( + ocr_request("mistral/model", &upstream.uri(), json!({})), + events.clone(), + false, + )) + .await + .unwrap(); + + assert_eq!( + *events.lock().unwrap(), + ["started", "before_send", "response", "success"] + ); + assert_eq!(received(&upstream).await.len(), 1); +} + +#[tokio::test] +async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() { + let upstream = upstream([pages_response()]).await; + let events = Arc::new(Mutex::new(Vec::new())); + + let error = perform_with(recording_host( + ocr_request("mistral/model", &upstream.uri(), json!({})), + events.clone(), + true, + )) + .await + .unwrap_err(); + + assert!( + matches!(&error, Error::InvalidRequest(message) if message == "blocked"), + "{error:?}" + ); + assert_eq!( + *events.lock().unwrap(), + ["started", "before_send", "failure"] + ); + assert!(received(&upstream).await.is_empty()); +} + +#[tokio::test] +async fn an_upstream_failure_emits_one_terminal_failure() { + let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await; + let events = Arc::new(Mutex::new(Vec::new())); + + let result = perform_with(recording_host( + ocr_request("mistral/model", &upstream.uri(), json!({})), + events.clone(), + false, + )) + .await; + + assert!(result.is_err()); + assert_eq!( + *events.lock().unwrap(), + ["started", "before_send", "failure"] + ); + assert_eq!(received(&upstream).await.len(), 1); +} + +#[tokio::test] +async fn an_invalid_provider_response_is_observed_before_normalization_fails() { + let upstream = upstream([json_response(json!({"pages": "invalid"}))]).await; + let observed = Arc::new(Mutex::new(Vec::new())); + let recorder = observed.clone(); + let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({}))) + .with_observer(move |event| { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { + recorder.lock().unwrap().push(raw.body.clone()); + } + }); + + let error = perform_with(host).await.unwrap_err(); + + assert!(matches!(error, Error::ResponseField { .. }), "{error:?}"); + assert_eq!(*observed.lock().unwrap(), [r#"{"pages":"invalid"}"#]); +} + +#[tokio::test] +async fn headers_returned_by_before_send_are_sent() { + let upstream = upstream([pages_response()]).await; + let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({}))) + .with_before_send(|mut wire, _| { + wire.headers + .push(("x-core-callback".into(), "edited".into())); + Ok(wire) + }); + + perform_with(host).await.unwrap(); + + assert_eq!( + only_request(&upstream).await.header("x-core-callback"), + Some("edited") + ); +} + +async fn before_send_context(request: LiteLLMOcrRequest) -> (WireRequest, RequestContext) { + let observed = Arc::new(Mutex::new(None)); + let captured = observed.clone(); + let host = LocalOcrHost::new(request).with_before_send(move |wire, context| { + *captured.lock().unwrap() = Some((wire.clone(), context.clone())); + Ok(wire) + }); + perform_with(host).await.unwrap(); + let context = observed.lock().unwrap().take(); + context.expect("before_send ran") +} + +#[tokio::test] +async fn before_send_sees_the_route_its_params_and_the_body() { + let upstream = upstream([pages_response()]).await; + + let (wire, context) = before_send_context(ocr_request( + "mistral/model", + &upstream.uri(), + json!({"pages": [0], "req_format": "native"}), + )) + .await; + + assert_eq!(context.custom_llm_provider, "mistral"); + assert_eq!(context.model, "model"); + assert_eq!(context.optional_params["req_format"], "native"); + assert!(context.secret_fields.is_empty()); + assert_eq!(wire.body["pages"], json!([0])); +} + +#[rstest] +#[case::client_secret(json!({"client_secret": "shh", "tenant_id": "t"}), &["client_secret"])] +#[case::no_secrets(json!({"tenant_id": "t"}), &[])] +#[tokio::test] +async fn before_send_names_the_secret_params(#[case] options: Value, #[case] secrets: &[&str]) { + let upstream = upstream([pages_response()]).await; + let request = ocr_request("azure_ai/model", &upstream.uri(), options).with_document( + OcrDocumentInput::Bytes { + bytes: b"abc".as_slice().into(), + file_name: None, + mime_type: Some("application/pdf".into()), + }, + ); + + let (_, context) = before_send_context(request).await; + + assert_eq!(context.secret_fields, secrets); +} + +/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_send`. +struct CallerTokenHost { + request: Mutex>, + trace: Mutex>, +} + +impl Host for CallerTokenHost { + async fn project(&self) -> Result { + self.trace.lock().unwrap().push("project".into()); + Ok(OcrProjection { + request: self.request.lock().unwrap().take().unwrap(), + caller_token: true, + }) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { + match op { + OcrOp::AcquireAzureAdToken(reply) => { + self.trace.lock().unwrap().push("token".into()); + reply.send(litellm_auth::ResolvedCredential::Static( + litellm_auth::SecretValue::new("caller-token"), + )); + Ok(()) + } + } + } + + async fn before_send( + &self, + wire: WireRequest, + _: &RequestContext, + ) -> Result { + let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization"); + let authorization = wire + .headers + .iter() + .find(|(name, _)| is_authorization(name)) + .map(|(_, value)| value.clone()) + .unwrap_or_default(); + self.trace + .lock() + .unwrap() + .push(format!("before_send:{authorization}")); + let headers = wire + .headers + .into_iter() + .map(|(name, value)| match is_authorization(&name) { + true => (name, "Bearer edited".to_string()), + false => (name, value), + }) + .collect(); + Ok(WireRequest { headers, ..wire }) + } +} + +#[tokio::test] +async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() { + let upstream = upstream([pages_response()]).await; + let host = CallerTokenHost { + request: Mutex::new(Some(without_api_key(ocr_request( + "azure_ai/model", + &upstream.uri(), + json!({}), + )))), + trace: Mutex::new(Vec::new()), + }; + + litellm_host::run::run(ocr_machine(ocr_client()), &host) + .await + .unwrap(); + + assert_eq!( + *host.trace.lock().unwrap(), + ["project", "token", "before_send:Bearer caller-token"] + ); + assert_eq!( + only_request(&upstream).await.header_values("authorization"), + ["Bearer edited"] + ); +} diff --git a/litellm-rust/crates/core/tests/ocr/machine.rs b/litellm-rust/crates/core/tests/ocr/machine.rs new file mode 100644 index 00000000000..073ce67e74b --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/machine.rs @@ -0,0 +1,284 @@ +use std::{ + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use litellm_core::ocr::{ + route::{OcrMachine, OcrOp, OcrProjection}, + types::OcrDocumentInput, +}; +use litellm_host::{ + event::{CallEvent, WireRequest}, + host::{Host, HostOp}, + machine::{HostFailure, Machine, MachineStep}, +}; +use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig; +use rstest::rstest; +use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify}; + +use super::{lifecycle::event_name, *}; + +/// Drives the machine by hand, answering every op through `host` except `before_send`, +/// which `intercept` answers so a test can fail or cancel exactly there. +async fn drive_until( + host: &LocalOcrHost, + mut intercept: impl FnMut(WireRequest) -> Result>, +) -> ( + Result, + Vec<&'static str>, + OcrMachine, +) { + let mut machine = ocr_machine(ocr_client()); + let mut ops = Vec::new(); + let outcome = loop { + let op = match machine.resume().await { + Ok(MachineStep::Host(op)) => op, + Ok(MachineStep::Complete(response)) => break Ok(response), + Err(error) => break Err(error), + }; + let answer = match op { + HostOp::Project(reply) => { + ops.push("Project"); + host.project() + .await + .map(|projection| reply.send(projection)) + .map_err(HostFailure::Error) + } + HostOp::Custom(op) => { + ops.push(match op { + OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken", + }); + host.custom_op(op).await.map_err(HostFailure::Error) + } + HostOp::BeforeSend { wire, reply, .. } => { + ops.push("BeforeSend"); + intercept(*wire).map(|wire| reply.send(wire)) + } + HostOp::Emit(event, reply) => { + let event = CallEvent::Machine(event); + ops.push(event_name(&event)); + host.emit(&event) + .await + .map(|()| reply.send(())) + .map_err(HostFailure::Error) + } + }; + if let Err(failure) = answer { + break machine.interrupt(failure).await; + } + }; + (outcome, ops, machine) +} + +/// Answers every op until `stop` fires, leaving the machine suspended mid-call. +async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, stop: &Notify) { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + tokio::select! { + _ = stop.notified() => break, + step = machine.resume() => { + match step.unwrap() { + MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), + MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), + MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), + MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), + MachineStep::Complete(_) => panic!("the stalled call completed"), + } + } + } + } + }) + .await + .expect("the call reached the stall point"); +} + +#[tokio::test] +async fn a_hand_driven_machine_performs_the_same_call() { + let upstream = upstream([json_response(json!({ + "pages": [{"index": 0, "markdown": "native"}] + }))]) + .await; + let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({}))); + + let (outcome, ops, mut machine) = drive_until(&host, Ok).await; + + assert_eq!(outcome.unwrap().pages[0].markdown, "native"); + assert_eq!(received(&upstream).await.len(), 1); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); + assert!(matches!( + machine.resume().await, + Err(Error::InvalidRequest(_)) + )); +} + +#[tokio::test] +async fn a_path_document_is_read_by_core_without_a_host_operation() { + let upstream = upstream([json_response(json!({ + "pages": [{"index": 0, "markdown": "path"}] + }))]) + .await; + let dir = std::env::temp_dir().join(format!("litellm-ocr-{}", rand::random::())); + std::fs::create_dir_all(&dir).unwrap(); + let path = dir.join("scan.png"); + std::fs::write(&path, b"abc").unwrap(); + let request = ocr_request("mistral/model", &upstream.uri(), json!({})).with_document( + OcrDocumentInput::Path { + path, + mime_type: None, + }, + ); + + let (response, ops, _) = drive_until(&LocalOcrHost::new(request), Ok).await; + std::fs::remove_dir_all(&dir).unwrap(); + + assert_eq!(response.unwrap().pages[0].markdown, "path"); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); + assert_eq!( + only_request(&upstream).await.json()["document"]["image_url"], + "data:image/png;base64,YWJj" + ); +} + +#[rstest] +#[case::failed(HostFailure::Error(Error::InvalidRequest("before_send failed".into())), "before_send failed")] +#[case::cancelled(HostFailure::Cancelled(Error::InvalidRequest("cancelled".into())), "cancelled")] +#[tokio::test] +async fn a_before_send_failure_ends_the_call_without_reaching_transport( + #[case] failure: HostFailure, + #[case] message: &str, +) { + let upstream = upstream([pages_response()]).await; + let host = LocalOcrHost::new(ocr_request("mistral/model", &upstream.uri(), json!({}))); + let failure = Arc::new(std::sync::Mutex::new(Some(failure))); + + let (outcome, ops, mut machine) = drive_until(&host, |_| { + Err(failure + .lock() + .unwrap() + .take() + .expect("before_send is asked once")) + }) + .await; + + assert!( + matches!(&outcome, Err(Error::InvalidRequest(actual)) if actual == message), + "{outcome:?}" + ); + assert_eq!(ops, ["Project", "BeforeSend"]); + assert!(machine.resume().await.is_err()); + assert!(received(&upstream).await.is_empty()); +} + +#[tokio::test] +async fn resuming_before_answering_keeps_the_pending_operation() { + let request = ocr_request("mistral/model", UNREACHABLE_BASE, json!({})); + let mut machine = ocr_machine(ocr_client()); + let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else { + panic!("expected the projection op first"); + }; + + assert!(machine.resume().await.is_err()); + reply.send(OcrProjection { + request, + caller_token: false, + }); + assert!(matches!( + machine.resume().await, + Ok(MachineStep::Host(HostOp::BeforeSend { .. })) + )); +} + +#[derive(Debug)] +struct PendingToken { + entered: Arc, + dropped: Arc, +} + +struct TokenFutureDrop(Arc); + +impl Drop for TokenFutureDrop { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +impl litellm_auth::TokenProvider for PendingToken { + fn acquire(&self) -> litellm_auth::TokenFuture<'_> { + Box::pin(async move { + let _guard = TokenFutureDrop(self.dropped.clone()); + self.entered.notify_one(); + std::future::pending().await + }) + } +} + +#[tokio::test] +async fn interrupt_drops_provider_captures_before_returning() { + let entered = Arc::new(Notify::new()); + let dropped = Arc::new(AtomicBool::new(false)); + let mut request = ocr_request("azure_ai/mistral-ocr", "https://example.invalid", json!({})); + request.transport = OcrTransportConfig { + extra_headers: vec![("authorization".into(), "Bearer test-key".into())], + ..request.transport + }; + request.azure_ad_token_provider = Some(litellm_auth::TokenProviderHandle::new(Arc::new( + PendingToken { + entered: entered.clone(), + dropped: dropped.clone(), + }, + ))); + let host = LocalOcrHost::new(request); + let mut machine = ocr_machine(ocr_client()); + + drive_until_notified(&mut machine, &host, &entered).await; + assert!(!dropped.load(Ordering::SeqCst)); + let acknowledgement = machine.interrupt(HostFailure::Cancelled(Error::InvalidRequest( + "cancelled".into(), + ))); + + assert!( + dropped.load(Ordering::SeqCst), + "interrupt returned while provider captures were still alive" + ); + assert!( + matches!(acknowledgement.await, Err(Error::InvalidRequest(message)) if message == "cancelled") + ); +} + +#[tokio::test] +async fn interrupting_an_in_flight_provider_request_closes_its_connection() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + let received = Arc::new(Notify::new()); + let server_received = received.clone(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = Vec::new(); + let mut buffer = [0u8; 4096]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = socket.read(&mut buffer).await.unwrap(); + request.extend_from_slice(&buffer[..read]); + } + server_received.notify_one(); + while socket.read(&mut buffer).await.unwrap() != 0 {} + }); + let host = LocalOcrHost::new(ocr_request("mistral/model", &base, json!({}))); + let mut machine = ocr_machine(ocr_client()); + + drive_until_notified(&mut machine, &host, &received).await; + let cancelled = Error::InvalidRequest("cancelled".into()); + + assert!( + machine + .interrupt(HostFailure::Cancelled(cancelled)) + .await + .is_err() + ); + tokio::time::timeout(Duration::from_secs(1), server) + .await + .expect("the provider connection stayed open after the interrupt") + .unwrap(); +} diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs new file mode 100644 index 00000000000..1a915389b20 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -0,0 +1,125 @@ +use litellm_core::ocr::{ + document::prepare_document, + route::{LocalOcrHost, ocr_machine}, + types::LiteLLMOcrRequest, + wire::{OcrWireRequest, decode_request}, +}; +use litellm_llms::base_llm::ocr::{ + error::Error, + handler::OcrClient, + transformation::{LiteLLMOcrResponse, OcrDocument}, +}; +use serde_json::{Map, Value, json}; +use wiremock::{MockServer, ResponseTemplate}; + +#[path = "../support/mod.rs"] +mod support; +use support::*; + +mod aws_textract; +mod azure_ai; +mod azure_document_intelligence; +mod cohere; +mod documents; +mod lifecycle; +mod machine; +mod mistral; +mod reducto; +mod vertex_ai; + +const INLINE_PDF: &str = "data:application/pdf;base64,YWJj"; + +fn object(value: Value) -> Map { + let Value::Object(map) = value else { + panic!("expected a json object, got {value}"); + }; + map +} + +fn ocr_client() -> OcrClient { + let document_http = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("test document client builds"); + OcrClient::for_test(reqwest::Client::new(), document_http) +} + +async fn perform(request: LiteLLMOcrRequest) -> Result { + litellm_core::ocr::client::perform(&ocr_client(), request).await +} + +async fn perform_with(host: LocalOcrHost) -> Result { + litellm_host::run::run(ocr_machine(ocr_client()), &host).await +} + +fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest { + OcrWireRequest { + model: model.into(), + document, + api_key: Some(litellm_auth::SecretValue::new("test-key")), + api_base: Some(base.into()), + custom_llm_provider: None, + extra_headers: None, + optional_params: object(options), + input_sources: Default::default(), + timeout_seconds: Some(2.0), + } +} + +/// A request for an inline PDF, authenticated with `test-key`. +fn ocr_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest { + ocr_request_with_document( + model, + base, + json!({"type": "document_url", "document_url": INLINE_PDF}), + options, + ) +} + +fn ocr_request_with_document( + model: &str, + base: &str, + document: Value, + options: Value, +) -> LiteLLMOcrRequest { + decode_request(wire(model, base, document, options)).expect("request decodes") +} + +fn document(value: Value) -> OcrDocument { + serde_json::from_value(value).expect("document parses") +} + +/// Points the request's resolved document at `source`, keeping its type. +fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest { + let resolved = request + .map_document(prepare_document) + .expect("document resolves"); + let document = resolved.document.clone().with_source(source.into()); + resolved.with_document(document.into()) +} + +fn with_headers(request: LiteLLMOcrRequest, headers: &[(&str, &str)]) -> LiteLLMOcrRequest { + let mut request = request; + request.transport.extra_headers = headers + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect(); + request +} + +fn without_api_key(request: LiteLLMOcrRequest) -> LiteLLMOcrRequest { + let mut request = request; + request.credentials.api_key = None; + request +} + +fn pages_response() -> ResponseTemplate { + json_response(json!({"pages": []})) +} + +/// An Azure Document Intelligence 202 whose operation lives on `server`. +fn accepted(server: &MockServer, body: Value) -> ResponseTemplate { + ResponseTemplate::new(202) + .insert_header("Operation-Location", format!("{}/operation", server.uri())) + .set_body_json(body) +} diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs new file mode 100644 index 00000000000..f80e564b03f --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -0,0 +1,248 @@ +use std::sync::Arc; + +use litellm_auth_gcp::VertexAuth; +use litellm_http::{ + HttpClientPool, HttpSettings, Resolution, + media::{PublicDnsResolver, UrlPolicy}, +}; +use litellm_llms::{ + base_llm::ocr::{ + settings::OcrSettings, + transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES}, + }, + mistral::ocr::transformation::MistralOcrConfig, +}; +use rstest::rstest; + +use super::*; + +#[tokio::test] +async fn direct_mistral_sends_one_request_with_every_option() { + let upstream = upstream([json_response(json!({ + "pages": [{"index": 0, "markdown": "hello", "custom": "preserved"}], + "usage_info": {"pages_processed": 1} + }))]) + .await; + + let result = perform(ocr_request( + "mistral/model", + &upstream.uri(), + json!({"pages": "0,2-4", "extract_header": true, "unknown": "ignored"}), + )) + .await + .unwrap(); + + assert_eq!(result.pages[0].markdown, "hello"); + assert_eq!(result.pages[0].extra_fields["custom"], "preserved"); + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/v1/ocr"); + assert_eq!(sent.header("authorization"), Some("Bearer test-key")); + assert_eq!( + sent.json(), + json!({ + "model": "model", + "document": {"type": "document_url", "document_url": INLINE_PDF}, + "pages": "0,2-4", + "extract_header": true, + "unknown": "ignored" + }) + ); +} + +#[rstest] +#[case::litellm_format(json!({}), false)] +#[case::native_format(json!({"req_format": "native"}), true)] +#[tokio::test] +async fn the_native_response_is_kept_only_when_requested( + #[case] options: Value, + #[case] kept: bool, +) { + let provider_response = json!({ + "pages": [{"index": 0, "markdown": "hello"}], + "usage_info": {"pages_processed": 1}, + "provider_only": "preserved" + }); + let upstream = upstream([json_response(provider_response.clone())]).await; + + let response = perform(ocr_request("mistral/model", &upstream.uri(), options)) + .await + .unwrap(); + + assert_eq!( + response.provider_native_response.map(Value::Object), + kept.then_some(provider_response) + ); +} + +#[rstest] +#[case::mistral("mistral/model", json!({}))] +#[case::vertex( + "vertex_ai/mistral-ocr-latest", + json!({"vertex_project": "test-project", "vertex_location": "us-central1"}) +)] +#[tokio::test] +async fn an_upstream_error_keeps_its_status_whole_body_and_headers( + #[case] model: &str, + #[case] options: Value, +) { + let payload = json!({"message": format!("{} END-OF-PROVIDER-BODY", "x".repeat(4096))}); + let expected_body = serde_json::to_string(&payload).unwrap(); + let upstream = upstream([status_response(422, payload) + .insert_header("Retry-After", "17") + .insert_header("X-Request-ID", "request-123") + .insert_header("X-Future-Header", "retained")]) + .await; + + let error = perform(ocr_request(model, &upstream.uri(), options)) + .await + .unwrap_err(); + + assert_eq!(received(&upstream).await.len(), 1); + let Error::Provider { + status, + body, + headers, + } = error + else { + panic!("expected provider error, got {error:?}"); + }; + assert_eq!(status, 422); + for (name, value) in [ + ("retry-after", "17"), + ("x-request-id", "request-123"), + ("x-future-header", "retained"), + ] { + assert!( + headers + .iter() + .any(|(key, actual)| key.eq_ignore_ascii_case(name) && actual == value), + "{name} missing from {headers:?}" + ); + } + assert_eq!(body, expected_body); +} + +#[rstest] +#[case::mistral_prefix("mistral/model", None, true)] +#[case::unknown_provider("model", Some("unknown"), false)] +fn decoding_accepts_known_providers_and_rejects_unknown_ones( + #[case] model: &str, + #[case] provider: Option<&str>, + #[case] accepted: bool, +) { + let request = OcrWireRequest { + custom_llm_provider: provider.map(Into::into), + ..wire( + model, + "https://example.com", + json!({"type": "document_url", "document_url": "https://example.com/doc.pdf"}), + json!({"extract_header": true, "unknown": 42}), + ) + }; + + assert_eq!(decode_request(request).is_ok(), accepted); +} + +#[rstest] +#[case::plain_key(&[("MISTRAL_API_KEY", "plain")], "plain")] +#[case::azure_key_wins(&[("MISTRAL_AZURE_API_KEY", "azure"), ("MISTRAL_API_KEY", "plain")], "azure")] +#[case::empty_azure_key_falls_through(&[("MISTRAL_AZURE_API_KEY", ""), ("MISTRAL_API_KEY", "plain")], "plain")] +#[tokio::test] +async fn missing_credentials_come_from_the_injected_secret_source( + #[case] secrets: &[(&str, &str)], + #[case] expected_key: &str, +) { + let upstream = upstream([pages_response()]).await; + let base = upstream.uri(); + let source = Arc::new(RecordingSecrets::new( + secrets + .iter() + .copied() + .chain([("MISTRAL_AZURE_API_BASE", base.as_str())]), + )); + let client = ocr_client().with_secrets(source.clone()); + let request = decode_request(OcrWireRequest { + api_key: None, + api_base: None, + ..wire( + "mistral/model", + &base, + json!({"type": "document_url", "document_url": INLINE_PDF}), + json!({}), + ) + }) + .unwrap(); + + litellm_core::ocr::client::perform(&client, request) + .await + .unwrap(); + + assert_eq!(source.requested(), MistralOcrConfig.secret_names()); + assert_eq!( + only_request(&upstream).await.header("authorization"), + Some(format!("Bearer {expected_key}").as_str()) + ); +} + +#[tokio::test] +async fn the_client_uses_the_injected_http_pool_configuration() { + let upstream = upstream([pages_response()]).await; + let settings = HttpSettings { + user_agent: Some("host-owned/1".into()), + ..HttpSettings::default() + }; + let client = OcrClient::new( + &HttpClientPool::new(Arc::new(PublicDnsResolver)), + &Resolution::from(&settings).config, + UrlPolicy::default(), + VertexAuth::default(), + OcrSettings::default(), + Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), + ) + .unwrap(); + + litellm_core::ocr::client::perform( + &client, + ocr_request("mistral/model", &upstream.uri(), json!({})), + ) + .await + .unwrap(); + + assert_eq!( + only_request(&upstream).await.header("user-agent"), + Some("host-owned/1") + ); +} + +#[test] +fn a_valid_response_limit_is_consumed_and_not_forwarded() { + let request = ocr_request( + "mistral/model", + UNREACHABLE_BASE, + json!({"max_response_bytes": 123}), + ); + + assert_eq!(request.transport.max_response_bytes, 123); + assert!(!request.optional_params.contains_key("max_response_bytes")); +} + +#[rstest] +#[case::zero(json!(0))] +#[case::negative(json!(-1))] +#[case::boolean(json!(true))] +#[case::string(json!("123"))] +#[case::fraction(json!(1.5))] +#[case::above_the_cap(json!(OCR_RESPONSE_MAX_BYTES + 1))] +#[case::null(Value::Null)] +fn an_invalid_response_limit_is_rejected(#[case] limit: Value) { + let Err(error) = decode_request(wire( + "mistral/model", + UNREACHABLE_BASE, + json!({"type": "document_url", "document_url": INLINE_PDF}), + json!({"max_response_bytes": limit}), + )) else { + panic!("invalid response limit {limit} accepted"); + }; + + assert!(error.to_string().contains("max_response_bytes"), "{error}"); +} diff --git a/litellm-rust/crates/core/tests/ocr/reducto.rs b/litellm-rust/crates/core/tests/ocr/reducto.rs new file mode 100644 index 00000000000..8ccab27e58d --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/reducto.rs @@ -0,0 +1,321 @@ +use std::sync::{Arc, Mutex}; + +use litellm_host::event::{CallEvent, MachineEvent, WireRequest}; +use rstest::rstest; + +use super::*; + +fn upload_response() -> ResponseTemplate { + json_response(json!({"file_id": "reducto://uploaded.pdf"})) +} + +fn chunks_response(chunks: Value) -> ResponseTemplate { + json_response(json!({"result": {"chunks": chunks}})) +} + +fn source_field(model: &str) -> &'static str { + match model.ends_with("parse-legacy") { + true => "document_url", + false => "input", + } +} + +#[rstest] +#[case::v3( + "reducto/parse-v3", + json!({ + "formatting": {"table_output_format": "html"}, + "retrieval": {"chunk_mode": "section"}, + "settings": {"ocr_system": "standard"}, + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }), + "reducto://already.pdf", + json!({ + "input": "reducto://already.pdf", + "formatting": {"table_output_format": "html"}, + "retrieval": {"chunk_mode": "section"}, + "settings": {"ocr_system": "standard"}, + "future_ocr_option": true, + "provider_option": "value" + }) +)] +#[case::legacy( + "reducto/parse-legacy", + json!({ + "enhance": {"agentic": [{"type": "table"}]}, + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }), + "reducto://legacy.pdf", + json!({ + "document_url": "reducto://legacy.pdf", + "options": {"enhance": {"agentic": [{"type": "table"}]}}, + "future_ocr_option": true, + "provider_option": "value" + }) +)] +#[tokio::test] +async fn an_uploaded_document_is_parsed_with_mapped_options( + #[case] model: &str, + #[case] options: Value, + #[case] source: &str, + #[case] expected: Value, +) { + let upstream = upstream([chunks_response(json!([]))]).await; + + perform(with_source( + ocr_request(model, &upstream.uri(), options), + source, + )) + .await + .unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/parse"); + assert_eq!(sent.json(), expected); +} + +#[rstest] +#[tokio::test] +async fn an_inline_document_is_uploaded_as_multipart_then_parsed( + #[values("parse-v3", "parse-legacy")] model: &str, + #[values("application/pdf", "image/png")] mime_type: &str, +) { + let upstream = upstream([ + upload_response(), + chunks_response(json!([{"content": "hello"}])), + ]) + .await; + let data_uri = format!("data:{mime_type};base64,YWJj"); + let document = match mime_type.starts_with("image/") { + true => json!({"type": "image_url", "image_url": data_uri}), + false => json!({"type": "document_url", "document_url": data_uri}), + }; + let request = with_headers( + ocr_request_with_document( + &format!("reducto/{model}"), + &upstream.uri(), + document, + json!({}), + ), + &[ + ("Content-Type", "application/json"), + ("X-Trace", "upload-test"), + ], + ); + + let response = perform(request).await.unwrap(); + + assert_eq!(response.pages[0].markdown, "hello"); + let requests = received(&upstream).await; + let [upload, parse] = requests.as_slice() else { + panic!( + "expected an upload and a parse, got {} requests", + requests.len() + ); + }; + assert_eq!(upload.url.path(), "/upload"); + assert!( + upload + .header("content-type") + .is_some_and(|value| value.starts_with("multipart/form-data; boundary=")), + "{:?}", + upload.header("content-type") + ); + assert_eq!(upload.header("x-trace"), Some("upload-test")); + let multipart = upload.body_text(); + assert!( + multipart.contains(&format!("Content-Type: {mime_type}\r\n")), + "{multipart}" + ); + assert!(multipart.contains("\r\n\r\nabc\r\n--"), "{multipart}"); + assert_eq!(parse.url.path(), "/parse"); + assert_eq!( + parse.json(), + json!({source_field(model): "reducto://uploaded.pdf"}) + ); + for request in &requests { + assert_eq!(request.header("authorization"), Some("Bearer test-key")); + } +} + +#[tokio::test] +async fn response_received_fires_once_for_the_parse_response() { + let upstream = upstream([upload_response(), chunks_response(json!([]))]).await; + let observed = Arc::new(Mutex::new(Vec::new())); + let recorder = observed.clone(); + let host = LocalOcrHost::new(ocr_request("reducto/parse-v3", &upstream.uri(), json!({}))) + .with_observer(move |event| { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { + recorder.lock().unwrap().push(raw.body.clone()); + } + }); + + perform_with(host).await.unwrap(); + + assert_eq!(received(&upstream).await.len(), 2); + assert_eq!(*observed.lock().unwrap(), [r#"{"result":{"chunks":[]}}"#]); +} + +#[rstest] +#[case::empty_id(json_response(json!({"file_id": ""})))] +#[case::missing_id(json_response(json!({})))] +#[case::null_id(json_response(json!({"file_id": null})))] +#[case::upload_failure(status_response(503, json!({"error": "unavailable"})))] +#[tokio::test] +async fn a_failed_upload_stops_before_parse(#[case] upload: ResponseTemplate) { + let upstream = upstream([upload]).await; + + let result = perform(ocr_request("reducto/parse-v3", &upstream.uri(), json!({}))).await; + + assert!(result.is_err()); + assert_eq!(received(&upstream).await.len(), 1); +} + +#[rstest] +#[case::remote_url("https://example.com/a.pdf", Error::ReductoSource)] +#[case::empty_file_id("reducto://", Error::RequestField { path: "document file id".into() })] +#[case::data_uri_without_payload("data:application/pdf;base64", Error::InvalidDataUri)] +#[case::invalid_base64("data:application/pdf;base64,INVALID!", Error::InvalidDataUri)] +#[tokio::test] +async fn invalid_document_sources_are_rejected_before_sending( + #[case] source: &str, + #[case] expected: Error, +) { + let upstream = upstream([json_response(json!({}))]).await; + + let result = perform(with_source( + ocr_request("reducto/parse-v3", &upstream.uri(), json!({})), + source, + )) + .await; + + assert!( + received(&upstream).await.is_empty(), + "sent invalid source: {source}" + ); + let error = result.unwrap_err(); + assert_eq!( + std::mem::discriminant(&error), + std::mem::discriminant(&expected) + ); + assert_eq!(error.http_status_code(), Some(400)); + assert_eq!(error.to_string(), expected.to_string()); +} + +#[tokio::test] +async fn a_forwarded_authorization_wins_and_the_native_response_is_omitted_by_default() { + let upstream = upstream([json_response( + json!({"job_id": "job-1", "result": {"chunks": []}}), + )]) + .await; + let request = with_headers( + with_source( + ocr_request("reducto/parse-v3", &upstream.uri(), json!({})), + "reducto://ready.pdf", + ), + &[("authorization", "Bearer existing")], + ); + + let response = perform(request).await.unwrap(); + + assert_eq!(response.provider_native_response, None); + assert_eq!( + only_request(&upstream).await.header_values("authorization"), + ["Bearer existing"] + ); +} + +#[tokio::test] +async fn native_format_retains_the_provider_response() { + let raw = json!({ + "result": {"chunks": [{"content": "native OCR response"}]}, + "usage": {"num_pages": 1} + }); + let upstream = upstream([json_response(raw.clone())]).await; + + let response = perform(with_source( + ocr_request( + "reducto/parse-v3", + &upstream.uri(), + json!({"req_format": "native"}), + ), + "reducto://ready.pdf", + )) + .await + .unwrap(); + + assert_eq!(response.pages[0].markdown, "native OCR response"); + assert_eq!( + response.provider_native_response.map(Value::Object), + Some(raw) + ); +} + +#[tokio::test] +async fn an_unknown_model_reaches_parse_and_keeps_its_name() { + let upstream = upstream([chunks_response( + json!([{"content": "future model response"}]), + )]) + .await; + + let response = perform(with_source( + ocr_request("reducto/future-parse-model", &upstream.uri(), json!({})), + "reducto://ready.pdf", + )) + .await + .unwrap(); + + assert_eq!(response.model, "future-parse-model"); + assert_eq!(response.pages[0].markdown, "future model response"); + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/parse"); + assert_eq!(sent.json(), json!({"input": "reducto://ready.pdf"})); +} + +#[tokio::test] +async fn a_guardrail_can_replace_the_document_before_upload() { + let upstream = upstream([chunks_response(json!([]))]).await; + let host = LocalOcrHost::new(ocr_request("reducto/parse-v3", &upstream.uri(), json!({}))) + .with_before_send(|wire, _| { + assert_eq!(wire.body["document_url"], INLINE_PDF); + Ok(WireRequest { + body: json!({"type": "document_url", "document_url": "reducto://guarded.pdf"}), + ..wire + }) + }); + + perform_with(host).await.unwrap(); + + let sent = only_request(&upstream).await; + assert_eq!(sent.url.path(), "/parse"); + assert_eq!(sent.json(), json!({"input": "reducto://guarded.pdf"})); +} + +#[rstest] +#[tokio::test] +async fn guardrail_headers_reach_both_upload_and_parse( + #[values("reducto/parse-v3", "reducto/parse-legacy")] model: &str, +) { + let upstream = upstream([upload_response(), chunks_response(json!([]))]).await; + let request = with_headers( + ocr_request(model, &upstream.uri(), json!({})), + &[("authorization", "Bearer original")], + ); + let host = LocalOcrHost::new(request).with_before_send(|wire, _| { + Ok(WireRequest { + headers: vec![("authorization".into(), "Bearer guarded".into())], + ..wire + }) + }); + + perform_with(host).await.unwrap(); + + let requests = received(&upstream).await; + let paths: Vec<&str> = requests.iter().map(|request| request.url.path()).collect(); + assert_eq!(paths, ["/upload", "/parse"]); + for request in &requests { + assert_eq!(request.header_values("authorization"), ["Bearer guarded"]); + } +} diff --git a/litellm-rust/crates/core/tests/ocr/vertex_ai.rs b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs new file mode 100644 index 00000000000..f0b2488e494 --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/vertex_ai.rs @@ -0,0 +1,184 @@ +use litellm_auth::{InputSource, Sourced}; +use litellm_core::ocr::arguments::is_supported_request; +use litellm_llms::base_llm::ocr::settings::OcrSettings; +use rstest::rstest; + +use super::*; + +#[tokio::test] +async fn mistral_is_served_at_the_resolved_project_and_location() { + let upstream = upstream([json_response(json!({ + "pages": [{"index": 0, "markdown": "hello"}], + "usage_info": {"pages_processed": 1} + }))]) + .await; + + let response = perform(ocr_request( + "vertex_ai/mistral-ocr-maas", + &upstream.uri(), + json!({ + "vertex_project": "project-1", + "vertex_location": "europe-west4", + "extract_footer": true + }), + )) + .await + .unwrap(); + + assert_eq!(response.pages[0].markdown, "hello"); + let sent = only_request(&upstream).await; + assert_eq!( + sent.url.path(), + "/v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict" + ); + assert_eq!(sent.header("authorization"), Some("Bearer test-key")); + assert_eq!( + sent.json(), + json!({ + "model": "mistral-ocr-maas", + "document": {"type": "document_url", "document_url": INLINE_PDF}, + "extract_footer": true + }) + ); +} + +#[tokio::test] +async fn configured_project_and_location_apply_when_the_call_sets_neither() { + let upstream = upstream([pages_response()]).await; + let client = ocr_client().with_settings(OcrSettings { + vertex_project: Some("configured-project".into()), + vertex_location: Some("europe-west4".into()), + ..OcrSettings::default() + }); + + litellm_core::ocr::client::perform( + &client, + ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})), + ) + .await + .unwrap(); + + assert_eq!( + only_request(&upstream).await.url.path(), + "/v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict" + ); +} + +#[tokio::test] +async fn a_supplied_authorization_is_forwarded_without_a_static_token() { + let upstream = upstream([pages_response()]).await; + let request = with_headers( + without_api_key(ocr_request( + "vertex_ai/model", + &upstream.uri(), + json!({"vertex_project": "project-1"}), + )), + &[("authorization", "Bearer supplied")], + ); + + perform(request).await.unwrap(); + + assert_eq!( + only_request(&upstream).await.header_values("authorization"), + ["Bearer supplied"] + ); +} + +#[tokio::test] +async fn invalid_credentials_fail_before_sending() { + let error = perform(ocr_request( + "vertex_ai/model", + UNREACHABLE_BASE, + json!({"vertex_credentials": true}), + )) + .await + .unwrap_err(); + + assert!(error.to_string().contains("vertex_credentials"), "{error}"); +} + +#[rstest] +#[tokio::test] +async fn a_request_controlled_api_base_is_rejected_before_vertex_auth( + #[values("vertex_ai/mistral-ocr-maas", "vertex_ai/deepseek-ocr-maas")] model: &str, +) { + let mut request = ocr_request( + model, + "https://caller.example", + json!({"vertex_project": "project-1"}), + ); + request.credentials.api_base = Some(Sourced::new( + "https://caller.example".into(), + InputSource::Request, + )); + + let error = perform(request).await.unwrap_err(); + + assert!( + error + .to_string() + .contains("request-controlled Vertex AI endpoint"), + "{error}" + ); +} + +#[tokio::test] +async fn deepseek_is_served_at_the_openai_compatible_endpoint() { + let upstream = upstream([json_response(json!({ + "choices": [{"message": {"content": "recognized"}}], + "usage": {"prompt_tokens": 1} + }))]) + .await; + let request = with_source( + ocr_request( + "vertex_ai/deepseek-ocr-maas", + &upstream.uri(), + json!({ + "vertex_project": "project-1", + "vertex_location": "europe-west4", + "temperature": 0.1, + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }), + ), + "gs://bucket/document.pdf", + ); + + let response = perform(request).await.unwrap(); + + assert_eq!(response.pages[0].markdown, "recognized"); + assert_eq!( + response.usage_info.unwrap().extra_fields["prompt_tokens"], + 1 + ); + let sent = only_request(&upstream).await; + assert_eq!( + sent.url.path(), + "/v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions" + ); + assert_eq!(sent.header("authorization"), Some("Bearer test-key")); + let body = sent.json(); + assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); + assert_eq!(body["temperature"], 0.1); + assert_eq!(body["future_ocr_option"], true); + assert_eq!(body["provider_option"], "value"); + assert!(body.get("vertex_project").is_none()); + assert!(body.get("extra_body").is_none()); + assert_eq!( + body["messages"][0]["content"][0], + json!({"type": "image_url", "image_url": "gs://bucket/document.pdf"}) + ); +} + +#[rstest] +#[case::deepseek("deepseek-ocr-maas", Some("vertex_ai"), true)] +#[case::mistral("mistral-ocr-maas", Some("vertex_ai"), true)] +#[case::prefixed("vertex_ai/mistral-ocr-maas", None, true)] +#[case::unknown_provider("model", Some("unknown"), false)] +fn supported_requests_follow_the_registered_configs( + #[case] model: &str, + #[case] provider: Option<&str>, + #[case] supported: bool, +) { + assert_eq!(is_supported_request(model, provider), supported); +} diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs new file mode 100644 index 00000000000..4d2fe0232d0 --- /dev/null +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -0,0 +1,155 @@ +//! Shared fixtures for route integration tests: a scripted upstream and a recording +//! secret source. + +#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset + +use std::sync::Mutex; + +use futures_util::future::BoxFuture; +use litellm_secrets::{SecretValue, source::SecretSource}; +use serde_json::Value; +use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; + +/// A port nothing listens on, for calls that must fail before any request is sent. +pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1"; + +/// Starts an upstream that answers its n-th request with the n-th response and 404s after. +pub async fn upstream(responses: impl IntoIterator) -> MockServer { + let server = MockServer::start().await; + respond_in_order(&server, responses).await; + server +} + +/// Scripts responses on a started server, for responses that need its address. +pub async fn respond_in_order( + server: &MockServer, + responses: impl IntoIterator, +) { + for response in responses { + Mock::given(any()) + .respond_with(response) + .up_to_n_times(1) + .mount(server) + .await; + } +} + +pub async fn received(server: &MockServer) -> Vec { + server + .received_requests() + .await + .expect("request recording is on") +} + +pub async fn only_request(server: &MockServer) -> Request { + let [request] = <[Request; 1]>::try_from(received(server).await) + .unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len())); + request +} + +pub fn json_response(body: Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(body) +} + +pub fn status_response(status: u16, body: Value) -> ResponseTemplate { + ResponseTemplate::new(status).set_body_json(body) +} + +pub trait ReceivedRequest { + fn header(&self, name: &str) -> Option<&str>; + fn header_values(&self, name: &str) -> Vec<&str>; + fn json(&self) -> Value; + fn body_text(&self) -> String; + /// The path and query, as the request line carried them. + fn target(&self) -> String; + fn query(&self, name: &str) -> Option; +} + +impl ReceivedRequest for Request { + fn header(&self, name: &str) -> Option<&str> { + self.headers.get(name).and_then(|value| value.to_str().ok()) + } + + fn header_values(&self, name: &str) -> Vec<&str> { + self.headers + .get_all(name) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect() + } + + fn json(&self) -> Value { + serde_json::from_slice(&self.body).expect("request body is json") + } + + fn body_text(&self) -> String { + String::from_utf8_lossy(&self.body).into_owned() + } + + fn target(&self) -> String { + match self.url.query() { + Some(query) => format!("{}?{query}", self.url.path()), + None => self.url.path().to_string(), + } + } + + fn query(&self, name: &str) -> Option { + self.url + .query_pairs() + .find_map(|(key, value)| (key == name).then(|| value.into_owned())) + } +} + +/// A secret source that answers from a fixed table and records every name it was asked for. +pub struct RecordingSecrets { + values: Vec<(String, String)>, + fails: bool, + requested: Mutex>, +} + +impl RecordingSecrets { + pub fn new<'a>(values: impl IntoIterator) -> Self { + Self { + values: values + .into_iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect(), + fails: false, + requested: Mutex::new(Vec::new()), + } + } + + pub fn empty() -> Self { + Self::new([]) + } + + pub fn failing() -> Self { + Self { + fails: true, + ..Self::empty() + } + } + + pub fn requested(&self) -> Vec { + self.requested.lock().unwrap().clone() + } +} + +impl SecretSource for RecordingSecrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { + Box::pin(async move { + self.requested.lock().unwrap().push(name.to_string()); + if self.fails { + return Err(litellm_secrets::Error::ManagedSecretMissing); + } + Ok(self + .values + .iter() + .find(|(key, _)| key == name) + .map(|(_, value)| SecretValue::new(value.clone()))) + }) + } +} diff --git a/litellm-rust/crates/coroutine/AGENTS.md b/litellm-rust/crates/coroutine/AGENTS.md new file mode 100644 index 00000000000..fcb4f4df47f --- /dev/null +++ b/litellm-rust/crates/coroutine/AGENTS.md @@ -0,0 +1,31 @@ +# Requirements + +Core must pause mid-call to ask the host for things it cannot do itself (Python callbacks, secret and token reads, `before_send` rewrites, stream demand), then continue where it stopped. Any change to this crate must keep every requirement below; the alternatives section says which one each rejected design breaks + +- R1 Core never calls the host: it names an op and waits for the answer, so it stays free of PyO3 and of any other host runtime +- R2 Async host work is awaited by the host's own driver in the caller's asyncio task (`litellm/rust_bridge/lifecycle.py`), so `contextvars` writes reach the caller; a Rust-side `into_future` would run it in a copied context +- R3 The body awaits real I/O (HTTP, `spawn_blocking`, timers) between yields, so `resume` is itself a future driven by the caller's runtime +- R4 Route code stays straight-line async (`host.route(OcrOp::ReadDocument).await?`) instead of hand-written states +- R5 Each op fixes its answer type at compile time: a host cannot answer `ReadDocument` with a token, and core never matches a result variant it did not ask for +- R6 A yield the body makes while being resumed is returned by that same poll, so the host driver's inline first poll needs no extra event-loop turn per op +- R7 No task is spawned: `cancel`, or dropping the coroutine, drops the body, and nothing waits forever on an answer that cannot come +- R8 Several yields can be pending at once, since route code hands clones of its `Co` to token providers and hooks +- R9 Stable Rust + +# Other implementations and why they do not fit + +- Nightly `std::ops::Coroutine`: breaks R9, and its body cannot await futures between yields (R3) +- `genawaiter`: resumes async bodies only with a noop waker, so the body cannot await real I/O (R3) +- `simple_coro`: typestate `Coro` makes answering before resuming a compile-time rule, but its body cannot await arbitrary futures (R3) and its reply type `R` is fixed per coroutine (R5) +- `corosensei` and other stackful coroutines: sync bodies on their own stack, no async I/O inside (R3) +- A hand-written phase enum with an `advance` match (the old `HostPhase`): every await point becomes a state (R4) +- An injected host trait with `async fn`s: core would call the host itself (R1, R2) +- Sans-IO, where core does no I/O and HTTP becomes one more host op: keeps every requirement and makes `resume` a pure step function, but HTTP, streaming, retries and timeouts would move out of core into every bridge; the one real alternative, not taken +- Temporal's Rust workflow SDK (`WorkflowFuture`, `WfContext`) is the closest precedent: an `async fn` polled in place, commands sent over a channel with a oneshot to unblock them. Roles are inverted there (the language SDK owns the program, core answers), and its workflow body may not do real I/O + +# Tradeoffs accepted + +- A tokio `mpsc` channel plus a `oneshot` per yield instead of compiler-generated states +- Protocol mistakes (resuming before answering, resuming after the end) are runtime `ResumeError`s, not compile errors +- Pending yields come out one per `resume`, in the order they were made, and each reply goes back to the yield that made it (R8) +- An answer sent after its yield stopped waiting (for example the body timed out on it) is discarded, since the body already moved on diff --git a/litellm-rust/crates/coroutine/Cargo.toml b/litellm-rust/crates/coroutine/Cargo.toml new file mode 100644 index 00000000000..3ff79ac5f2c --- /dev/null +++ b/litellm-rust/crates/coroutine/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "litellm-coroutine" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +description = "Async coroutines on stable Rust whose every yield carries its own typed reply" + +[dependencies] +thiserror.workspace = true +tokio = { workspace = true, features = ["sync"] } + +[dev-dependencies] +rstest.workspace = true +tokio = { workspace = true, features = ["rt", "macros", "time"] } diff --git a/litellm-rust/crates/coroutine/src/co.rs b/litellm-rust/crates/coroutine/src/co.rs new file mode 100644 index 00000000000..d84b041931c --- /dev/null +++ b/litellm-rust/crates/coroutine/src/co.rs @@ -0,0 +1,42 @@ +use std::sync::Weak; + +use tokio::sync::mpsc; + +use crate::{Abandoned, Reply, reply}; + +pub(crate) struct Request { + pub(crate) value: Y, + pub(crate) outstanding: Weak<()>, +} + +/// The body's handle for yielding, `genawaiter`'s `Co`. +pub struct Co { + yields: mpsc::UnboundedSender>, +} + +impl Clone for Co { + fn clone(&self) -> Self { + Self { + yields: self.yields.clone(), + } + } +} + +impl Co { + pub(crate) fn new(yields: mpsc::UnboundedSender>) -> Self { + Self { yields } + } + + /// Yields the value `ask` builds around a fresh [`Reply`] and waits for its answer. + pub async fn yield_(&self, ask: impl FnOnce(Reply) -> Y) -> Result { + let (reply, answer) = reply(); + let outstanding = reply.outstanding(); + self.yields + .send(Request { + value: ask(reply), + outstanding, + }) + .map_err(|_| Abandoned)?; + answer.await + } +} diff --git a/litellm-rust/crates/coroutine/src/coroutine.rs b/litellm-rust/crates/coroutine/src/coroutine.rs new file mode 100644 index 00000000000..fa34f816a37 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/coroutine.rs @@ -0,0 +1,94 @@ +use std::{ + future::{Future, poll_fn}, + pin::Pin, + sync::Weak, + task::{Context, Poll}, +}; + +use tokio::sync::mpsc; + +use crate::{Co, ResumeError, co::Request}; + +/// What one `resume` produced, as in [`std::ops::CoroutineState`]. +#[derive(Debug, PartialEq, Eq)] +pub enum CoroutineState { + Yielded(Y), + Complete(C), +} + +type Body = Pin + Send>>; + +enum Step { + Yielded(Request), + Complete(C), +} + +fn queued( + yields: &mut mpsc::UnboundedReceiver>, + context: &mut Context<'_>, +) -> Option> { + match yields.poll_recv(context) { + Poll::Ready(request) => request, + Poll::Pending => None, + } +} + +pub struct Coroutine { + body: Option>, + yields: mpsc::UnboundedReceiver>, + outstanding: Weak<()>, +} + +impl Coroutine { + /// Builds the body from `producer`. Nothing runs until the first `resume`. + pub fn new(producer: impl FnOnce(Co) -> F) -> Self + where + F: Future + Send + 'static, + { + let (sender, yields) = mpsc::unbounded_channel(); + Self { + body: Some(Box::pin(producer(Co::new(sender)))), + yields, + outstanding: Weak::new(), + } + } + + pub async fn resume(&mut self) -> Result, ResumeError> { + let Some(body) = self.body.as_mut() else { + return Err(ResumeError::Finished); + }; + if self.outstanding.strong_count() > 0 { + return Err(ResumeError::Unanswered); + } + let yields = &mut self.yields; + let step = poll_fn(|context| { + if let Some(request) = queued(yields, context) { + return Poll::Ready(Step::Yielded(request)); + } + if let Poll::Ready(output) = body.as_mut().poll(context) { + return Poll::Ready(Step::Complete(output)); + } + queued(yields, context) + .map_or(Poll::Pending, |request| Poll::Ready(Step::Yielded(request))) + }) + .await; + match step { + Step::Yielded(Request { value, outstanding }) => { + self.outstanding = outstanding; + Ok(CoroutineState::Yielded(value)) + } + Step::Complete(output) => { + self.cancel(); + Ok(CoroutineState::Complete(output)) + } + } + } + + /// Drops the body and fails every yield still waiting, or yet to be made, with + /// [`Abandoned`](crate::Abandoned). + pub fn cancel(&mut self) { + self.body = None; + self.yields.close(); + while self.yields.try_recv().is_ok() {} + } +} diff --git a/litellm-rust/crates/coroutine/src/error.rs b/litellm-rust/crates/coroutine/src/error.rs new file mode 100644 index 00000000000..b8fded23fdb --- /dev/null +++ b/litellm-rust/crates/coroutine/src/error.rs @@ -0,0 +1,14 @@ +/// A `resume` the coroutine refused, leaving it as it was. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +pub enum ResumeError { + #[error("coroutine resumed after it finished")] + Finished, + #[error("coroutine resumed before the reply to its last yield was sent or dropped")] + Unanswered, +} + +/// No answer will come to a yield: its [`Reply`](crate::Reply) was dropped unsent, or the +/// coroutine it was sent to is gone. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +#[error("the yield was abandoned before it was answered")] +pub struct Abandoned; diff --git a/litellm-rust/crates/coroutine/src/lib.rs b/litellm-rust/crates/coroutine/src/lib.rs new file mode 100644 index 00000000000..636aaf1b2b6 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/lib.rs @@ -0,0 +1,12 @@ +//! Async coroutines on stable Rust whose every yield carries its own typed [`Reply`]. +//! See `AGENTS.md` for the requirement, the alternatives and the contracts. + +mod co; +mod coroutine; +mod error; +mod reply; + +pub use co::Co; +pub use coroutine::{Coroutine, CoroutineState}; +pub use error::{Abandoned, ResumeError}; +pub use reply::{Answer, Reply, reply}; diff --git a/litellm-rust/crates/coroutine/src/reply.rs b/litellm-rust/crates/coroutine/src/reply.rs new file mode 100644 index 00000000000..b3cb7da2e9d --- /dev/null +++ b/litellm-rust/crates/coroutine/src/reply.rs @@ -0,0 +1,60 @@ +use std::{ + fmt, + future::Future, + pin::Pin, + sync::{Arc, Weak}, + task::{Context, Poll}, +}; + +use tokio::sync::oneshot; + +use crate::Abandoned; + +/// The one way to answer a yield. Sending or dropping it settles the yield. +pub struct Reply { + slot: oneshot::Sender, + outstanding: Arc<()>, +} + +impl Reply { + /// An answer the yield no longer awaits is discarded. + pub fn send(self, answer: A) { + let _ = self.slot.send(answer); + } + + /// Alive until this reply is sent or dropped. + pub(crate) fn outstanding(&self) -> Weak<()> { + Arc::downgrade(&self.outstanding) + } +} + +impl fmt::Debug for Reply { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("Reply") + } +} + +/// The waiting end of a [`Reply`]. +pub struct Answer { + slot: oneshot::Receiver, +} + +impl Future for Answer { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + Pin::new(&mut self.slot) + .poll(context) + .map(|answer| answer.map_err(|_| Abandoned)) + } +} + +/// A reply outside any coroutine, for answering a host operation directly. +pub fn reply() -> (Reply, Answer) { + let (slot, answer) = oneshot::channel(); + let reply = Reply { + slot, + outstanding: Arc::new(()), + }; + (reply, Answer { slot: answer }) +} diff --git a/litellm-rust/crates/coroutine/tests/coroutine.rs b/litellm-rust/crates/coroutine/tests/coroutine.rs new file mode 100644 index 00000000000..91d195503df --- /dev/null +++ b/litellm-rust/crates/coroutine/tests/coroutine.rs @@ -0,0 +1,256 @@ +use std::{ + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_coroutine::{Abandoned, Co, Coroutine, CoroutineState, Reply, ResumeError, reply}; +use rstest::rstest; +use tokio::time::timeout; + +#[derive(Debug)] +enum Ask { + Name(Reply<&'static str>), + Count(Reply), +} + +type Test = Coroutine; + +fn yielded(state: Result, ResumeError>) -> Ask { + match state { + Ok(CoroutineState::Yielded(ask)) => ask, + Ok(CoroutineState::Complete(_)) => panic!("expected a yield, the body returned"), + Err(error) => panic!("expected a yield, resume failed: {error}"), + } +} + +fn complete(state: Result, ResumeError>) -> C { + match state { + Ok(CoroutineState::Complete(output)) => output, + Ok(CoroutineState::Yielded(ask)) => panic!("expected completion, got {ask:?}"), + Err(error) => panic!("expected completion, resume failed: {error}"), + } +} + +fn name(ask: Ask) -> Reply<&'static str> { + match ask { + Ask::Name(reply) => reply, + other => panic!("expected a name ask, got {other:?}"), + } +} + +fn count(ask: Ask) -> Reply { + match ask { + Ask::Count(reply) => reply, + other => panic!("expected a count ask, got {other:?}"), + } +} + +/// A body parked at one name ask, with nothing else going on. +fn suspended_once() -> Test> { + Coroutine::new(|co| async move { co.yield_(Ask::Name).await }) +} + +#[tokio::test] +async fn each_typed_answer_resumes_the_yield_that_asked_for_it() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Name).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + format!("{first}+{second}") + }); + + name(yielded(coroutine.resume().await)).send("a"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), "a+2"); +} + +/// A driver that polls `resume` once, inline, sees every yield the body makes during +/// that poll instead of being sent back to its event loop. +#[test] +fn a_yield_made_while_resuming_is_returned_by_that_same_poll() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Count).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + first + second + }); + let mut context = std::task::Context::from_waker(std::task::Waker::noop()); + let mut poll_once = + |coroutine: &mut Test| match std::pin::pin!(coroutine.resume()).poll(&mut context) { + std::task::Poll::Ready(state) => state, + std::task::Poll::Pending => panic!("resume needed a second poll"), + }; + + count(yielded(poll_once(&mut coroutine))).send(1); + count(yielded(poll_once(&mut coroutine))).send(2); + + assert_eq!(complete(poll_once(&mut coroutine)), 3); +} + +#[tokio::test] +async fn the_body_awaits_real_futures_between_yields() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(5)).await; + co.yield_(Ask::Count).await.unwrap() + }); + + count(yielded(coroutine.resume().await)).send(7); + + assert_eq!(complete(coroutine.resume().await), 7); +} + +#[tokio::test] +async fn concurrent_yields_come_out_in_order_and_are_answered_separately() { + let mut coroutine: Test<(&str, u32)> = Coroutine::new(|co| async move { + let (first, second) = tokio::join!(co.yield_(Ask::Name), co.yield_(Ask::Count)); + (first.unwrap(), second.unwrap()) + }); + + name(yielded(coroutine.resume().await)).send("one"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), ("one", 2)); +} + +#[tokio::test] +async fn resuming_before_the_reply_is_settled_is_refused_and_keeps_the_yield_waiting() { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + + assert_eq!( + coroutine.resume().await.unwrap_err(), + ResumeError::Unanswered + ); + + reply.send("real"); + assert_eq!(complete(coroutine.resume().await), Ok("real")); +} + +#[tokio::test] +async fn a_dropped_reply_abandons_its_yield() { + let mut coroutine = suspended_once(); + drop(yielded(coroutine.resume().await)); + + assert_eq!(complete(coroutine.resume().await), Err(Abandoned)); +} + +#[tokio::test] +async fn an_answer_the_yield_no_longer_awaits_is_discarded() { + let mut coroutine: Test<&str> = Coroutine::new(|co| async move { + tokio::select! { + biased; + _ = co.yield_(Ask::Name) => unreachable!("the answer comes after the body moved on"), + () = std::future::ready(()) => {} + } + co.yield_(Ask::Name).await.unwrap() + }); + let stale = name(yielded(coroutine.resume().await)); + stale.send("stale"); + + name(yielded(coroutine.resume().await)).send("fresh"); + + assert_eq!(complete(coroutine.resume().await), "fresh"); +} + +#[rstest] +#[case::returned(false)] +#[case::cancelled(true)] +#[tokio::test] +async fn a_finished_coroutine_refuses_to_resume(#[case] cancel: bool) { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + if cancel { + coroutine.cancel(); + } else { + reply.send("done"); + complete(coroutine.resume().await).unwrap(); + } + + assert_eq!(coroutine.resume().await.unwrap_err(), ResumeError::Finished); +} + +#[tokio::test] +async fn a_dropped_resume_leaves_the_coroutine_resumable() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(20)).await; + co.yield_(Ask::Count).await.unwrap() + }); + assert!( + timeout(Duration::from_millis(1), coroutine.resume()) + .await + .is_err() + ); + + count(yielded(coroutine.resume().await)).send(3); + + assert_eq!(complete(coroutine.resume().await), 3); +} + +struct Dropped(Arc>); + +impl Drop for Dropped { + fn drop(&mut self) { + *self.0.lock().unwrap() = true; + } +} + +#[tokio::test] +async fn cancel_drops_the_body() { + let dropped = Arc::new(Mutex::new(false)); + let guard = Dropped(Arc::clone(&dropped)); + let mut coroutine: Test<()> = Coroutine::new(|co| async move { + let _guard = guard; + co.yield_(Ask::Count).await.unwrap(); + }); + let _reply = yielded(coroutine.resume().await); + + coroutine.cancel(); + + assert!(*dropped.lock().unwrap()); +} + +#[rstest] +#[case::cancelled(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_co_that_escaped_the_body_is_abandoned_once_the_coroutine_ends(#[case] cancel: bool) { + let escaped: Arc>>> = Arc::default(); + let slot = Arc::clone(&escaped); + let mut coroutine: Test<()> = Coroutine::new(move |co| { + *slot.lock().unwrap() = Some(co.clone()); + async move { + co.yield_(Ask::Count).await.unwrap(); + } + }); + let _reply = yielded(coroutine.resume().await); + let co = escaped.lock().unwrap().take().unwrap(); + let waiting = tokio::spawn(async move { co.yield_(Ask::Name).await }); + tokio::task::yield_now().await; + + if cancel { + coroutine.cancel(); + } else { + drop(coroutine); + } + + let outcome = timeout(Duration::from_secs(1), waiting) + .await + .expect("an escaped yield waits forever") + .unwrap(); + assert_eq!(outcome, Err(Abandoned)); +} + +#[rstest] +#[case::sent(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_detached_reply_settles_its_answer(#[case] send: bool) { + let (reply, answer) = reply::(); + if send { + reply.send(5); + } else { + drop(reply); + } + + assert_eq!(answer.await, if send { Ok(5) } else { Err(Abandoned) }); +} diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 5aca13eeb18..7c1919f9f39 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,8 +1,8 @@ - Target invariants; implementation and runtime validation may lag these rules -- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits +- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) + - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index fb6379dc35a..c1b35c0f69d 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +bytes.workspace = true futures-util.workspace = true litellm-host.workspace = true pyo3.workspace = true diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 3a4cb49be4d..7f07475bc4c 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -1,5 +1,5 @@ use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::PyRuntimeError; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -85,7 +85,7 @@ pub trait PythonLifecycle: Send + Sync { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -/// Why a route operation the host answered did not produce a result: the route's own code +/// Why a custom operation the host answered did not produce a result: the route's own code /// rejected it, which the route classifies like any other native failure, or Python code /// raised, which reaches the caller as it was raised. #[derive(Debug)] @@ -100,45 +100,61 @@ impl From for InvokeError { } } -/// The Python side of one route: answers the route's own operations, builds the public +/// The Python side of one protocol: answers its custom operations, builds the public /// response and classifies native failures into public exceptions. -pub trait RouteHost: Send + Sync { - type Route: Route; +pub trait ProtocolHost: Send + Sync { + type Protocol: Protocol; /// The public exception a native failure maps to, kept as a value until the driver /// raises it. type Failure: Into; - /// `arguments` is the keyword view the lifecycle's `begin` produced, not the - /// caller's own dict. A route host that projects from it inherits whatever that - /// adapter rewrote. - fn invoke( + /// Projects the call's request. `arguments` is the keyword view the lifecycle's + /// `begin` produced, not the caller's own dict, so the projection inherits whatever + /// that adapter rewrote. + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: ::Op, - ) -> Result<::OpResult, InvokeError<::Error>>; + ) -> Result< + ::Projection, + InvokeError<::Error>, + >; + + /// Answers `op` through its reply. + fn invoke( + &mut self, + py: Python<'_>, + op: ::Op, + ) -> Result<(), InvokeError<::Error>>; fn complete( &mut self, py: Python<'_>, - response: ::Response, + response: ::Response, + ) -> PyResult>; + + /// What the stream carries at hand-off, as the caller's stream receives it. + fn head( + &mut self, + py: Python<'_>, + head: ::StreamHead, ) -> PyResult>; /// One streamed chunk as the caller receives it. fn chunk( &mut self, py: Python<'_>, - chunk: ::Chunk, + chunk: ::Chunk, ) -> PyResult>; fn classify( &self, py: Python<'_>, - error: ::Error, + error: ::Error, ) -> PyResult; - fn host_error(error: &PyErr) -> ::Error; + fn host_error(error: &PyErr) -> ::Error; fn close(&mut self, py: Python<'_>); diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 77a294d274b..372af2843bd 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -2,10 +2,11 @@ use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; +use litellm_host::event::WireRequest; use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_host::host::{Demand, HostOp, HostResult, HostStep}; +use litellm_host::host::{Demand, HostOp, HostStep, Reply}; use litellm_host::machine::{HostFailure, Machine, MachineStep}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -13,21 +14,21 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; -type RouteOf = ::Route; -type ErrorOf = as Route>::Error; -type ResponseOf = as Route>::Response; -type NativeStep = MachineStep, ResponseOf>; +type ProtocolOf = ::Protocol; +type ErrorOf = as Protocol>::Error; +type ResponseOf = as Protocol>::Response; +type NativeStep = MachineStep, ResponseOf>; type NativeResult = Result, ErrorOf>; -type NativeResume = Option>, HostFailure>>>; +type Interruption = Option>>; type MachineResult = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >; struct MachineState { @@ -44,12 +45,11 @@ enum Stage { Failed(Py), } -#[derive(Clone, Copy)] enum Expect { Started, Arguments, - Wire, - Emitted, + Wire(Reply), + Emitted(Reply<()>), Response, Terminal, } @@ -58,20 +58,30 @@ enum Pending { Native, Adapter(Expect), /// The stream handed to the caller waits for its next read or its close. - Consumer, + Consumer(Reply), } -enum Next { +/// A route answer as the driver resumes on it: a Python exception interrupts the call as +/// raised, a native rejection resumes the machine with it. +fn answered(answer: Result<(), InvokeError>) -> PyResult> { + match answer { + Ok(()) => Ok(Ok(())), + Err(InvokeError::Native(error)) => Ok(Err(error)), + Err(InvokeError::Python(error)) => Err(error), + } +} + +enum Next { Return(ExecutionStep), Continue(HostStep, Py>), } struct PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { - route: H, + host: H, adapter: Box, machine: Option>>>, arguments: Option>, @@ -89,17 +99,17 @@ where pub fn run_call( py: Python<'_>, machine: M, - route: H, + host: H, adapter: Box, arguments: Py, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine> + 'static, + H: ProtocolHost + 'static, + M: Machine> + 'static, { let mut driver = PythonDriver { - route, + host, adapter, machine: Some(Arc::new(Mutex::new(MachineState { machine, @@ -124,10 +134,10 @@ where } match driver.resume(None)? { ExecutionStep::Return(value) => Ok(value), - ExecutionStep::Open => py + ExecutionStep::Open(head) => py .import("litellm.rust_bridge.lifecycle")? .getattr("SyncStream")? - .call1((Py::new(py, Execution::suspended(driver))?,)) + .call1((Py::new(py, Execution::suspended(driver))?, head)) .map(Bound::unbind), ExecutionStep::Await(_) | ExecutionStep::Yield(_) => { Err(PyRuntimeError::new_err("sync call suspended")) @@ -141,8 +151,8 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool { impl PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn timing(&self) -> Timing { Timing { @@ -172,13 +182,13 @@ where self.run_steps(py, HostStep::Ready(result)) } (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), - (Some(Pending::Consumer), Some(read)) => { - let demand = if read.is_ok() { + (Some(Pending::Consumer(reply)), Some(read)) => { + reply.send(if read.is_ok() { Demand::More } else { Demand::Detached - }; - self.resume_machine(py, Some(Ok(HostResult::Demand(demand)))) + }); + self.resume_machine(py, None) } (Some(Pending::Adapter(expect)), Some(result)) => { match self.adapter.resume(py, result) { @@ -196,22 +206,24 @@ where step: LifecycleStep, expect: Expect, ) -> PyResult { + if let LifecycleStep::Await(awaitable) = step { + self.pending = Some(Pending::Adapter(expect)); + return Ok(ExecutionStep::Await(awaitable)); + } match (expect, step) { - (_, LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(expect)); - Ok(ExecutionStep::Await(awaitable)) - } (Expect::Started, LifecycleStep::Done) => self.begin(py), (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) } - (Expect::Wire, LifecycleStep::Wire(wire)) => { - self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire)))) + (Expect::Wire(reply), LifecycleStep::Wire(wire)) => { + reply.send(*wire); + self.resume_machine(py, None) } - (Expect::Emitted, LifecycleStep::Done) => { - self.resume_machine(py, Some(Ok(HostResult::Emitted))) + (Expect::Emitted(reply), LifecycleStep::Done) => { + reply.send(()); + self.resume_machine(py, None) } (Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response), (Expect::Terminal, LifecycleStep::Done) => match &self.stage { @@ -242,9 +254,9 @@ where fn resume_machine( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult { - let step = self.resume_core(py, result)?; + let step = self.resume_core(py, interruption)?; self.run_steps(py, step) } @@ -277,54 +289,72 @@ where } Err(error) => return self.machine_failed(py, error).map(Next::Return), }; - let answer = match op { - HostOp::Route(op) => { + let answered = match op { + HostOp::Project(reply) => { let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; - match self.route.invoke(py, arguments.bind(py), op) { - Ok(result) => Ok(HostResult::Route(result)), - Err(InvokeError::Native(error)) => { - return self - .resume_core(py, Some(Err(HostFailure::Error(error)))) - .map(Next::Continue); - } - Err(InvokeError::Python(error)) => Err(error), - } + let projected = self.host.project(py, arguments.bind(py)); + answered(projected.map(|projection| reply.send(projection))) } - HostOp::BeforeSend { wire, context } => { - match self.adapter.before_send(py, wire, &context) { - Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)), + HostOp::Custom(op) => answered(self.host.invoke(py, op)), + HostOp::BeforeSend { + wire, + context, + reply, + } => match self.adapter.before_send(py, wire, &context) { + Ok(LifecycleStep::Wire(wire)) => { + reply.send(*wire); + Ok(Ok(())) + } + Ok(LifecycleStep::Await(awaitable)) => { + self.pending = Some(Pending::Adapter(Expect::Wire(reply))); + return Ok(Next::Return(ExecutionStep::Await(awaitable))); + } + Ok(_) => return Err(missing_state()), + Err(error) => Err(error), + }, + HostOp::Open(head, reply) => return self.opened(py, head, reply).map(Next::Return), + HostOp::Deliver(chunk, reply) => { + return self.delivered(py, chunk, reply).map(Next::Return); + } + HostOp::Emit(event, reply) => { + match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { + Ok(LifecycleStep::Done) => { + reply.send(()); + Ok(Ok(())) + } Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Wire)); + self.pending = Some(Pending::Adapter(Expect::Emitted(reply))); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } Ok(_) => return Err(missing_state()), Err(error) => Err(error), } } - HostOp::Open(_) => return self.opened(py).map(Next::Return), - HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return), - HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { - Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), - Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Emitted)); - return Ok(Next::Return(ExecutionStep::Await(awaitable))); - } - Ok(_) => return Err(missing_state()), - Err(error) => Err(error), - }, }; - match answer { - Ok(answer) => self.resume_core(py, Some(Ok(answer))).map(Next::Continue), + match answered { + Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue), + Ok(Err(native)) => self + .resume_core(py, Some(HostFailure::Error(native))) + .map(Next::Continue), Err(error) => self.interrupt(py, error).map(Next::Return), } } - fn opened(&mut self, py: Python<'_>) -> PyResult { + fn opened( + &mut self, + py: Python<'_>, + head: as Protocol>::StreamHead, + reply: Reply, + ) -> PyResult { self.stage = Stage::Streaming; + let head = match self.host.head(py, head) { + Ok(head) => head, + Err(error) => return self.interrupt(py, error), + }; match self.adapter.opened(py) { Ok(()) => { - self.pending = Some(Pending::Consumer); - Ok(ExecutionStep::Open) + self.pending = Some(Pending::Consumer(reply)); + Ok(ExecutionStep::Open(head)) } Err(error) => self.interrupt(py, error), } @@ -333,15 +363,16 @@ where fn delivered( &mut self, py: Python<'_>, - chunk: as Route>::Chunk, + chunk: as Protocol>::Chunk, + reply: Reply, ) -> PyResult { - let chunk = match self.route.chunk(py, chunk) { + let chunk = match self.host.chunk(py, chunk) { Ok(chunk) => chunk, Err(error) => return self.interrupt(py, error), }; match self.adapter.delivered(py, &chunk) { Ok(()) => { - self.pending = Some(Pending::Consumer); + self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Yield(chunk)) } Err(error) => self.interrupt(py, error), @@ -357,25 +388,24 @@ where } else { HostFailure::Error(native) }; - self.resume_machine(py, Some(Err(failure))) + self.resume_machine(py, Some(failure)) } fn resume_core( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult, Py>> { let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?); let future = async move { let mut state = state.lock().await; - let result = match result { - Some(Err(failure)) => state + let result = match interruption { + Some(failure) => state .machine .interrupt(failure) .await .map(MachineStep::Complete), - Some(Ok(result)) => state.machine.resume(Some(result)).await, - None => state.machine.resume(None).await, + None => state.machine.resume().await, }; state.result = Some(result); Ok(()) @@ -414,7 +444,7 @@ where fn completed(&mut self, py: Python<'_>, response: ResponseOf) -> PyResult { self.ended_at = Some(epoch_seconds()); - let public = match self.route.complete(py, response) { + let public = match self.host.complete(py, response) { Ok(public) => public, Err(error) => return self.failure(py, error, FailureOrigin::Call), }; @@ -441,7 +471,7 @@ where /// fails, that failure is raised with the native error's text as its `__context__`. fn classified(&self, py: Python<'_>, error: ErrorOf) -> PyErr { let native = error.to_string(); - let classifier_error = match self.route.classify(py, error) { + let classifier_error = match self.host.classify(py, error) { Ok(failure) => return failure.into(), Err(classifier_error) => classifier_error, }; @@ -486,7 +516,7 @@ where if self.machine.take().is_some() { Python::attach(|py| { self.adapter.close(py); - self.route.close(py); + self.host.close(py); }); } } @@ -494,15 +524,15 @@ where impl ExecutionBody for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn resume(&mut self, result: Option>>) -> PyResult { Python::attach(|py| self.drive(py, result)) } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.route.traverse(visit)?; + self.host.traverse(visit)?; self.adapter.traverse(visit)?; visit.call(&self.arguments)?; visit.call(&self.interrupted)?; @@ -516,8 +546,8 @@ where impl Drop for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn drop(&mut self) { self.clear(); @@ -528,8 +558,8 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_host::event::{MachineEvent, RequestContext, WireRequest}; - use litellm_host::machine::{Interrupted, Step}; + use litellm_host::event::{MachineEvent, RawResponse, RequestContext}; + use litellm_host::machine::{CallMachine, MachineFault}; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; @@ -573,22 +603,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - struct Synthetic; - - impl Route for Synthetic { - type Response = String; - type Error = Error; - type Op = &'static str; - type OpResult = String; - type Chunk = std::convert::Infallible; - type StreamHead = std::convert::Infallible; + impl From for Error { + fn from(fault: MachineFault) -> Self { + Self(format!("{fault:?}")) + } } - /// Yields the scripted ops in order, then completes or fails as scripted. - struct ScriptedMachine { - ops: Vec>, - outcome: Option>, - answers: Vec, + struct Synthetic; + + impl Protocol for Synthetic { + type Response = String; + type Error = Error; + type Projection = String; + type Op = (&'static str, Reply); + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } fn wire() -> WireRequest { @@ -609,37 +638,6 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl Machine for ScriptedMachine { - type Route = Synthetic; - type Complete = String; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(async move { - if let Some(result) = result { - self.answers.push(match result { - HostResult::Route(value) => value, - HostResult::BeforeSend(wire) => wire.url, - HostResult::Emitted => "emitted".into(), - HostResult::Demand(demand) => format!("{demand:?}"), - }); - } - if !self.ops.is_empty() { - return Ok(MachineStep::Host(self.ops.remove(0))); - } - self.outcome - .take() - .ok_or_else(|| Error("resumed after completion".into()))? - .map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.ops.clear(); - self.outcome = None; - Box::pin(async move { Err(failure.into_error()) }) - } - } - #[derive(Default)] struct Log(Arc>>); @@ -677,22 +675,41 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl RouteHost for SyntheticHost { - type Route = Synthetic; + impl SyntheticHost { + fn answer(&self, value: impl FnOnce() -> String) -> Result> { + match self.op { + OpScript::Answer => Ok(value()), + OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), + OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), + } + } + } + + impl ProtocolHost for SyntheticHost { + type Protocol = Synthetic; type Failure = Classified; + fn project( + &mut self, + _: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result> { + self.log.push("project"); + self.answer(|| format!("project:{}", arguments.len())) + } + fn invoke( &mut self, _: Python<'_>, - arguments: &Bound<'_, PyDict>, - op: &'static str, - ) -> Result> { - self.log.push(format!("route:{op}")); - match self.op { - OpScript::Answer => Ok(format!("{op}:{}", arguments.len())), - OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), - OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), - } + (op, reply): (&'static str, Reply), + ) -> Result<(), InvokeError> { + self.log.push(format!("op:{op}")); + self.answer(|| op.to_string()) + .map(|answer| reply.send(answer)) + } + + fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult> { + match head {} } fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { @@ -719,7 +736,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } fn close(&mut self, _: Python<'_>) { - self.log.push("route.close"); + self.log.push("host.close"); } fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -828,7 +845,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_scripted( py: Python<'_>, - machine: ScriptedMachine, + machine: CallMachine, op: OpScript, script: AdapterScript, asynchronous: bool, @@ -848,12 +865,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_hosted( py: Python<'_>, - machine: ScriptedMachine, - route: SyntheticHost, + machine: CallMachine, + host: SyntheticHost, script: AdapterScript, asynchronous: bool, ) -> (PyResult>, Vec) { - let log = Log(route.log.0.clone()); + let log = Log(host.log.0.clone()); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script, @@ -863,7 +880,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let result = run_call( py, machine, - route, + host, Box::new(adapter), arguments.unbind(), asynchronous, @@ -884,21 +901,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri (result, log.entries()) } - fn success_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![ - HostOp::Route("project"), - HostOp::BeforeSend { - wire: Box::new(wire()), - context: Box::new(context()), - }, - HostOp::Emit(MachineEvent::ResponseReceived { - raw: litellm_host::event::RawResponse { body: "raw".into() }, - }), - ], - outcome: Some(Ok("done".into())), - answers: Vec::new(), - } + /// Answers to projection, to the route op and to `before_send` all reach the + /// response, so a driver that misroutes a reply changes what the call returns. + fn success_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + let projected = host.project().await?; + let signed = host.custom_op(|reply| ("sign", reply)).await?; + let wire = host.before_send(wire(), context()).await?; + host.emit(MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }) + .await?; + Ok(format!("{projected}|{signed}|{}", wire.url)) + }) + }) } #[test] @@ -917,32 +934,194 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri AdapterScript::Plain, asynchronous, ); - assert_eq!(result.unwrap().extract::(py).unwrap(), "done"); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:1|sign|rewritten" + ); assert_eq!( log, [ "started", "begin", - "route:project", + "project", + "op:sign", "before_send", "response:raw", "complete", "after_success", - "succeeded:done", + "succeeded:project:1|sign|rewritten", "adapter.close", - "route.close", + "host.close", ] ); } }); } - fn failing_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![HostOp::Route("project")], - outcome: Some(Err(Error("provider exploded".into()))), - answers: Vec::new(), + struct Streaming; + + impl Protocol for Streaming { + type Response = (); + type Error = Error; + type Projection = (); + type Op = std::convert::Infallible; + type Chunk = &'static str; + type StreamHead = Vec<(&'static str, &'static str)>; + } + + struct StreamingHost; + + impl ProtocolHost for StreamingHost { + type Protocol = Streaming; + type Failure = Classified; + + fn project( + &mut self, + _: Python<'_>, + _: &Bound<'_, PyDict>, + ) -> Result<(), InvokeError> { + Ok(()) } + + fn invoke( + &mut self, + _: Python<'_>, + op: std::convert::Infallible, + ) -> Result<(), InvokeError> { + match op {} + } + + fn head( + &mut self, + py: Python<'_>, + head: Vec<(&'static str, &'static str)>, + ) -> PyResult> { + let headers = PyDict::new(py); + for (name, value) in head { + headers.set_item(name, value)?; + } + let hidden = PyDict::new(py); + hidden.set_item("additional_headers", headers)?; + Ok(hidden.into_any().unbind()) + } + + fn chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult> { + Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind()) + } + + fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult> { + Ok(py.None()) + } + + fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + Ok(Classified(error.0)) + } + + fn host_error(error: &PyErr) -> Error { + Error(error.to_string()) + } + + fn close(&mut self, _: Python<'_>) {} + + fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { + Ok(()) + } + } + + fn streaming_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + host.project().await?; + if host.open(vec![("request-id", "req_1")]).await? == Demand::Detached { + return Ok(()); + } + for chunk in ["first", "second"] { + if host.deliver(chunk).await? == Demand::Detached { + break; + } + } + Ok(()) + }) + }) + } + + /// Drives a `Stream` (async) or `SyncStream` to completion from a sync test. + fn read_all(py: Python<'_>, stream: &Bound<'_, PyAny>, asynchronous: bool) -> Vec { + if !asynchronous { + return stream + .try_iter() + .unwrap() + .map(|chunk| chunk.unwrap().extract().unwrap()) + .collect(); + } + std::iter::from_fn(|| { + let stop = stream + .call_method0("__anext__") + .unwrap() + .call_method1("send", (py.None(),)) + .unwrap_err(); + if stop.is_instance_of::(py) { + return None; + } + assert!(stop.is_instance_of::(py)); + Some(stop.value(py).getattr("value").unwrap().extract().unwrap()) + }) + .collect() + } + + #[test] + fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + for asynchronous in [false, true] { + let log = Log::default(); + let adapter = SyntheticAdapter { + log: Log(log.0.clone()), + script: AdapterScript::Plain, + }; + let handed = run_call( + py, + streaming_machine(), + StreamingHost, + Box::new(adapter), + PyDict::new(py).unbind(), + asynchronous, + ) + .unwrap(); + let stream = if asynchronous { + let stop = handed.call_method1(py, "send", (py.None(),)).unwrap_err(); + stop.value(py).getattr("value").unwrap() + } else { + handed.into_bound(py) + }; + let hidden: std::collections::HashMap< + String, + std::collections::HashMap, + > = stream.getattr("_hidden_params").unwrap().extract().unwrap(); + assert_eq!( + hidden["additional_headers"], + std::collections::HashMap::from([( + "request-id".to_string(), + "req_1".to_string() + )]) + ); + assert_eq!(log.entries(), ["started", "begin", "opened"]); + assert_eq!(read_all(py, &stream, asynchronous), ["first", "second"]); + } + }); + } + + fn failing_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + host.project().await?; + Err(Error("provider exploded".into())) + }) + }) } #[test] @@ -969,11 +1148,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classified: provider exploded", "adapter.close", - "route.close", + "host.close", ] ); } @@ -1003,11 +1182,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:op rejected", "failed:Call:classified: op rejected", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1035,10 +1214,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "failed:Call:op failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1073,11 +1252,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classifier failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1106,7 +1285,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "begin", "failed:Host:begin failed", "adapter.close", - "route.close" + "host.close" ] ); }); @@ -1130,7 +1309,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ); assert_eq!(result.unwrap().extract::(py).unwrap(), "replaced"); assert!(log.contains(&"succeeded:replaced".to_string())); - assert!(!log.contains(&"succeeded:done".to_string())); + assert!(!log.contains(&"succeeded:project:1|rewritten".to_string())); } }); } @@ -1159,7 +1338,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "after_success", "failed:Host:after_success failed", "adapter.close", - "route.close" + "host.close" ] ); assert!(!log.iter().any(|entry| entry.starts_with("succeeded"))); @@ -1175,18 +1354,31 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri crate::initialize_python(); Python::attach(|py| { struct Cancelling(Log); - impl RouteHost for Cancelling { - type Route = Synthetic; + impl ProtocolHost for Cancelling { + type Protocol = Synthetic; type Failure = Classified; - fn invoke( + fn project( &mut self, _: Python<'_>, _: &Bound<'_, PyDict>, - _: &'static str, ) -> Result> { - self.0.push("route"); + self.0.push("project"); Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into()) } + fn invoke( + &mut self, + _: Python<'_>, + _: (&'static str, Reply), + ) -> Result<(), InvokeError> { + Err(missing_state().into()) + } + fn head( + &mut self, + _: Python<'_>, + head: std::convert::Infallible, + ) -> PyResult> { + match head {} + } fn chunk( &mut self, _: Python<'_>, @@ -1210,7 +1402,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } let log = Log::default(); - let route = Cancelling(Log(log.0.clone())); + let host = Cancelling(Log(log.0.clone())); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script: AdapterScript::Plain, @@ -1218,7 +1410,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let error = run_call( py, success_machine(), - route, + host, Box::new(adapter), PyDict::new(py).unbind(), false, @@ -1227,7 +1419,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri assert!(!error.is_instance_of::(py)); assert_eq!( log.entries(), - ["started", "begin", "route", "adapter.close"] + ["started", "begin", "project", "adapter.close"] ); }); } diff --git a/litellm-rust/crates/host-python/src/file_reader.rs b/litellm-rust/crates/host-python/src/file_reader.rs new file mode 100644 index 00000000000..bbc7a233b28 --- /dev/null +++ b/litellm-rust/crates/host-python/src/file_reader.rs @@ -0,0 +1,241 @@ +//! A caller's file-like object: anything with a callable `read`, kept as a handle and read +//! once, on the host's thread, into bytes Rust owns. + +use bytes::Bytes; +use pyo3::{ + exceptions::PyTypeError, + gc::{PyTraverseError, PyVisit}, + prelude::*, + pybacked::PyBackedBytes, + types::{PyBytes, PyString}, +}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct FileContent { + pub bytes: Bytes, + pub file_name: Option, +} + +#[derive(Debug)] +pub struct PythonFileReader { + reader: Py, + name: Option, +} + +impl PythonFileReader { + /// `None` when `file` has no callable `read`. The object's `name` is read now, its + /// contents only on [`read`](Self::read). + pub fn from_file_like(file: &Bound<'_, PyAny>) -> PyResult> { + let reader = file + .getattr_opt("read")? + .filter(|value| value.is_callable()); + let Some(reader) = reader else { + return Ok(None); + }; + let name = file + .getattr_opt("name")? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?; + Ok(Some(Self { + reader: reader.unbind(), + name, + })) + } + + pub fn read(&self, py: Python<'_>) -> PyResult { + let value = self.reader.bind(py).call0()?; + let bytes = if value.is_instance_of::() { + Bytes::from(value.extract::()?) + } else if value.is_instance_of::() { + py_bytes(&value)? + } else { + return Err(PyTypeError::new_err(format!( + "file read must return bytes or str, got {}", + value.get_type(), + ))); + }; + Ok(FileContent { + bytes, + file_name: self.name.clone(), + }) + } + + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.reader) + } +} + +/// An exact `bytes` object is shared without copying and keeps the Python object alive; +/// a `bytes` subclass is copied. +pub fn py_bytes(value: &Bound<'_, PyAny>) -> PyResult { + if value.is_exact_instance_of::() { + return Ok(Bytes::from_owner(value.extract::()?)); + } + Ok(Bytes::copy_from_slice( + value.extract::()?.as_ref(), + )) +} + +#[cfg(test)] +mod tests { + use pyo3::{exceptions::PyTypeError, types::PyDict}; + + use super::*; + + fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run(source, Some(&locals), Some(&locals)).unwrap(); + locals + } + + fn reader<'py>(locals: &Bound<'py, PyDict>, name: &str) -> PythonFileReader { + PythonFileReader::from_file_like(&locals.get_item(name).unwrap().unwrap()) + .unwrap() + .unwrap() + } + + #[test] + fn objects_without_a_callable_read_are_not_readers() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Attribute: + read = 'not callable' +plain = object() +attribute = Attribute() +", + ); + for name in ["plain", "attribute"] { + let file = locals.get_item(name).unwrap().unwrap(); + assert!(PythonFileReader::from_file_like(&file).unwrap().is_none()); + } + }); + } + + #[test] + fn the_name_is_taken_up_front_and_the_contents_only_on_read() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Reader: + name = 'scan.png' + def __init__(self): + self.reads = 0 + def read(self): + self.reads += 1 + return b'abc' +file = Reader() +", + ); + let reads = || { + locals + .get_item("file") + .unwrap() + .unwrap() + .getattr("reads") + .unwrap() + .extract::() + .unwrap() + }; + let file = reader(&locals, "file"); + assert_eq!(reads(), 0); + let content = file.read(py).unwrap(); + assert_eq!(reads(), 1); + assert_eq!( + content, + FileContent { + bytes: b"abc".as_slice().into(), + file_name: Some("scan.png".into()), + } + ); + }); + } + + #[test] + fn read_results_are_normalized_and_exceptions_keep_their_identity() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = KeyError('reader failed') +class Raising: + def read(self): + raise failure +class Text: + def read(self): + return 'héllo' +class Wrong: + def read(self): + return 7 +raising = Raising() +text = Text() +wrong = Wrong() +", + ); + let error = reader(&locals, "raising").read(py).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + assert_eq!( + reader(&locals, "text").read(py).unwrap().bytes.as_ref(), + "héllo".as_bytes() + ); + let error = reader(&locals, "wrong").read(py).unwrap_err(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("bytes or str")); + }); + } + + #[rstest::rstest] + #[case::read("read")] + #[case::name("name")] + fn attribute_failures_keep_their_identity(#[case] attribute: &str) { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = LookupError('file property failed') +class File: + def __getattribute__(self, name): + if name == attribute: + raise failure + return super().__getattribute__(name) + name = 'scan.pdf' + def read(self): + return b'abc' +file = File() +", + ); + locals.set_item("attribute", attribute).unwrap(); + let error = + PythonFileReader::from_file_like(&locals.get_item("file").unwrap().unwrap()) + .unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { + Python::initialize(); + let (bytes, pointer) = Python::attach(|py| { + let value = PyBytes::new(py, b"document bytes"); + let pointer = value.as_bytes().as_ptr() as usize; + (py_bytes(value.as_any()).unwrap(), pointer) + }); + assert_eq!(bytes.as_ptr() as usize, pointer); + assert_eq!(bytes.as_ref(), b"document bytes"); + } +} diff --git a/litellm-rust/crates/host-python/src/handle.rs b/litellm-rust/crates/host-python/src/handle.rs index 10abbadbda5..24adfd404d7 100644 --- a/litellm-rust/crates/host-python/src/handle.rs +++ b/litellm-rust/crates/host-python/src/handle.rs @@ -8,9 +8,9 @@ use pyo3::prelude::*; pub enum ExecutionStep { Return(Py), Await(Py), - /// The call streams: the caller gets a stream over this execution, which stays - /// suspended until the stream asks for a chunk. - Open, + /// The call streams: the caller gets a stream over this execution carrying this head, + /// and the execution stays suspended until the stream asks for a chunk. + Open(Py), Yield(Py), } @@ -75,7 +75,7 @@ impl Execution { let step = body.resume(result)?; let (tag, value, suspended) = match step { ExecutionStep::Await(value) => ("Await", value, true), - ExecutionStep::Open => ("Open", py.None(), true), + ExecutionStep::Open(head) => ("Open", head, true), ExecutionStep::Yield(value) => ("Yield", value, true), ExecutionStep::Return(value) => ("Complete", value, false), }; diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 55b27e34b46..7e17c4da51e 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,6 +1,6 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and //! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine) -//! against a Python route host and a Python lifecycle. Everything here is Python-specific by +//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. mod adapter; @@ -8,13 +8,14 @@ mod argument; mod callable; mod driver; mod execution; +mod file_reader; mod fork_gate; mod gil; mod handle; mod marshal; pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; @@ -24,6 +25,7 @@ pub use execution::{ reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value, runtime_started, }; +pub use file_reader::{FileContent, PythonFileReader, py_bytes}; pub use fork_gate::RuntimeAlreadyStarted; pub use gil::{PythonContext, attach_blocking, release_count, release_gil}; pub use handle::{Execution, ExecutionBody, ExecutionStep}; diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index 0c7c46192b5..bbbed68f345 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] litellm-auth.workspace = true +litellm-coroutine.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["sync"] } diff --git a/litellm-rust/crates/host/src/host.rs b/litellm-rust/crates/host/src/host.rs index aba35185a18..9714b9470a3 100644 --- a/litellm-rust/crates/host/src/host.rs +++ b/litellm-rust/crates/host/src/host.rs @@ -1,28 +1,27 @@ use std::future::Future; -use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; -use crate::route::Route; +pub use litellm_coroutine::{Abandoned, Answer, Reply, reply}; -/// One suspension point of a native call, performed by the host. -pub enum HostOp { - Route(R::Op), +use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; +use crate::protocol::Protocol; + +/// One suspension point of a native call, performed by the host and answered through the +/// [`Reply`] it carries. +pub enum HostOp { + /// The first op of every call: the caller's request as the host projects it. + Project(Reply), + Custom(R::Op), BeforeSend { wire: Box, context: Box, + reply: Reply, }, - Emit(MachineEvent), + Emit(MachineEvent, Reply<()>), /// The response streams: the host hands the caller a stream and answers once the /// caller asks for the first chunk or goes away. - Open(R::StreamHead), + Open(R::StreamHead, Reply), /// The next chunk of an open stream, answered once the caller asks for the one after. - Deliver(R::Chunk), -} - -pub enum HostResult { - Route(R::OpResult), - BeforeSend(Box), - Emitted, - Demand(Demand), + Deliver(R::Chunk, Reply), } /// Whether the caller of a streamed call still reads it. @@ -39,10 +38,13 @@ pub enum HostStep { Suspend(S), } -/// An in-process host: answers route operations and observes the call without leaving +/// An in-process host: answers custom operations and observes the call without leaving /// the Rust runtime. Language hosts implement their own driver instead. -pub trait Host: Send + Sync { - fn route(&self, op: R::Op) -> impl Future> + Send; +pub trait Host: Send + Sync { + fn project(&self) -> impl Future> + Send; + + /// Answers `op` through its reply, or fails the call. + fn custom_op(&self, op: R::Op) -> impl Future> + Send; fn before_send( &self, diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index 65479c2380f..c6b9e59b65a 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -1,12 +1,13 @@ //! The contract between a native call and the host runtime that drives it. //! //! A host is whatever sits on the far side of the language boundary: CPython today, -//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns +//! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns //! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers -//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent. +//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and +//! may rewrite the wire request before it is sent. pub mod event; pub mod host; pub mod machine; -pub mod route; +pub mod protocol; pub mod run; diff --git a/litellm-rust/crates/host/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs index ba7e242e766..76e3504ca28 100644 --- a/litellm-rust/crates/host/src/machine/auth.rs +++ b/litellm-rust/crates/host/src/machine/auth.rs @@ -1,22 +1,21 @@ use std::sync::Arc; use super::{HostChannel, MachineFault}; -use crate::route::Route; +use crate::{host::Reply, protocol::Protocol}; use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -/// A route whose host can mint credentials on the call's behalf. -pub trait TokenRoute: Route { - fn acquire_token_op() -> Self::Op; - fn token_credential(result: Self::OpResult) -> Option; +/// A protocol whose host can mint credentials on the call's behalf. +pub trait TokenProtocol: Protocol { + fn acquire_token_op(reply: Reply) -> Self::Op; } /// A [`TokenProvider`] that asks the host for each credential through the call's own /// operation channel, so the host answers it on the caller's thread and context. -pub struct HostTokenProvider { +pub struct HostTokenProvider { channel: HostChannel, } -impl std::fmt::Debug for HostTokenProvider { +impl std::fmt::Debug for HostTokenProvider { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter.write_str("HostTokenProvider") } @@ -24,7 +23,7 @@ impl std::fmt::Debug for HostTokenProvider { impl HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { pub fn handle(channel: HostChannel) -> TokenProviderHandle { @@ -34,19 +33,15 @@ where impl TokenProvider for HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { fn acquire(&self) -> TokenFuture<'_> { Box::pin(async move { - let result = self - .channel - .route(R::acquire_token_op()) + self.channel + .custom_op(R::acquire_token_op) .await - .map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?; - R::token_credential(result).ok_or_else(|| { - Error::AzureTokenAcquisition("invalid token provider host result".into()) - }) + .map_err(|error| Error::AzureTokenAcquisition(error.to_string())) }) } } diff --git a/litellm-rust/crates/host/src/machine/call_machine.rs b/litellm-rust/crates/host/src/machine/call_machine.rs new file mode 100644 index 00000000000..af0bc50fbe6 --- /dev/null +++ b/litellm-rust/crates/host/src/machine/call_machine.rs @@ -0,0 +1,137 @@ +//! The one machine every route runs on: the route's provider future as a +//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No +//! task is spawned; dropping the machine drops the in-flight call. + +use std::{future::Future, pin::Pin}; + +use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError}; + +use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; +use crate::{ + event::{MachineEvent, RequestContext, WireRequest}, + host::{Demand, HostOp, Reply}, + protocol::Protocol, +}; + +/// The machine's own failures, distinct from anything the provider call reports. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum MachineFault { + /// The host dropped an op's reply unanswered, or went away while the call waited. + Abandoned, + /// The host resumed the call out of turn. + Protocol(ResumeError), +} + +pub type ExecuteFuture = + Pin::Response, ::Error>> + Send>>; + +/// The provider side of the machine: how the in-flight call reaches its host. +pub struct HostChannel { + co: Co>, +} + +impl Clone for HostChannel { + fn clone(&self) -> Self { + Self { + co: self.co.clone(), + } + } +} + +impl HostChannel +where + R::Error: From, +{ + async fn yield_( + &self, + ask: impl FnOnce(Reply) -> HostOp + Send, + ) -> Result { + self.co + .yield_(ask) + .await + .map_err(|_| MachineFault::Abandoned.into()) + } + + pub async fn project(&self) -> Result { + self.yield_(HostOp::Project).await + } + + /// Asks the host to perform the custom operation `ask` builds around its reply, as in + /// `host.custom_op(OcrOp::AcquireAzureAdToken)`. + pub async fn custom_op( + &self, + ask: impl FnOnce(Reply) -> R::Op + Send, + ) -> Result { + self.yield_(|reply| HostOp::Custom(ask(reply))).await + } + + pub async fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + self.yield_(|reply| HostOp::BeforeSend { + wire: Box::new(wire), + context: Box::new(context), + reply, + }) + .await + } + + pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { + self.yield_(|reply| HostOp::Emit(event, reply)).await + } + + pub async fn open(&self, head: R::StreamHead) -> Result { + self.yield_(|reply| HostOp::Open(head, reply)).await + } + + pub async fn deliver(&self, chunk: R::Chunk) -> Result { + self.yield_(|reply| HostOp::Deliver(chunk, reply)).await + } +} + +type CallCoroutine = + Coroutine, Result<::Response, ::Error>>; + +pub struct CallMachine { + coroutine: CallCoroutine, +} + +impl CallMachine +where + R::Error: From, +{ + pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { + Self { + coroutine: Coroutine::new(|co| execute(HostChannel { co })), + } + } +} + +impl Machine for CallMachine +where + R::Error: From, +{ + type Protocol = R; + type Complete = R::Response; + + fn resume(&mut self) -> Step<'_, Self> { + Box::pin(async move { + match self + .coroutine + .resume() + .await + .map_err(MachineFault::Protocol)? + { + CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)), + CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete), + } + }) + } + + fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { + self.coroutine.cancel(); + Box::pin(async move { Err(failure.into_error()) }) + } +} diff --git a/litellm-rust/crates/host/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/mod.rs index 2c26db61582..0c7501633fa 100644 --- a/litellm-rust/crates/host/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -1,16 +1,16 @@ mod auth; -mod route_machine; +mod call_machine; use std::future::Future; use std::pin::Pin; -pub use auth::{HostTokenProvider, TokenRoute}; -pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine}; +pub use auth::{HostTokenProvider, TokenProtocol}; +pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault}; -use crate::host::{HostOp, HostResult}; -use crate::route::Route; +use crate::host::HostOp; +use crate::protocol::Protocol; -pub enum MachineStep { +pub enum MachineStep { Host(HostOp), Complete(C), } @@ -19,8 +19,8 @@ pub type Step<'a, M> = Pin< Box< dyn Future< Output = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >, > + Send + 'a, @@ -30,7 +30,10 @@ pub type Step<'a, M> = Pin< pub type Interrupted<'a, M> = Pin< Box< dyn Future< - Output = Result<::Complete, <::Route as Route>::Error>, + Output = Result< + ::Complete, + <::Protocol as Protocol>::Error, + >, > + Send + 'a, >, @@ -51,19 +54,18 @@ impl HostFailure { } /// A resumable call. Core implements it per route; a host drives it. Every suspension -/// point is an op the host performs and answers with a result. +/// point is an op the host performs and answers through the op's own reply before it +/// resumes the call again. pub trait Machine: Send { - type Route: Route; + type Protocol: Protocol; type Complete: Send + 'static; - /// `None` on the first call and whenever the previous step completed without - /// yielding an op; otherwise the result of the op last yielded. - fn resume(&mut self, result: Option>) -> Step<'_, Self>; + fn resume(&mut self) -> Step<'_, Self>; /// The host failed to perform the pending op, or the caller cancelled. The call /// yields no further ops. fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self>; } diff --git a/litellm-rust/crates/host/src/machine/route_machine.rs b/litellm-rust/crates/host/src/machine/route_machine.rs deleted file mode 100644 index 38a0b8bc16a..00000000000 --- a/litellm-rust/crates/host/src/machine/route_machine.rs +++ /dev/null @@ -1,199 +0,0 @@ -//! The one machine every route runs on: it owns the route's provider future, polls it in -//! place, and turns the host operations that future requests into [`Machine`] steps. No -//! task is spawned; dropping the machine drops the in-flight call. - -use std::{future::Future, pin::Pin}; - -use tokio::sync::{mpsc, oneshot}; - -use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; -use crate::{ - event::{MachineEvent, RequestContext, WireRequest}, - host::{Demand, HostOp, HostResult}, - route::Route, -}; - -/// The machine's own failures, distinct from anything the provider call reports. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MachineFault { - /// The host driver went away while the call was waiting on it. - Abandoned, - /// The host answered out of turn: a result with nothing pending, or nothing when a - /// result was pending. - Protocol(&'static str), - /// The host answered a route operation with the wrong result variant. - Mismatch, -} - -pub type ExecuteFuture = - Pin::Response, ::Error>> + Send>>; - -struct PendingOp { - op: HostOp, - reply: oneshot::Sender>, -} - -/// The provider side of the machine: how the in-flight call reaches its host. -pub struct HostChannel { - ops: mpsc::UnboundedSender>, -} - -impl Clone for HostChannel { - fn clone(&self) -> Self { - Self { - ops: self.ops.clone(), - } - } -} - -impl HostChannel -where - R::Error: From, -{ - async fn invoke(&self, op: HostOp) -> Result, R::Error> { - let (reply, answer) = oneshot::channel(); - self.ops - .send(PendingOp { op, reply }) - .map_err(|_| MachineFault::Abandoned)?; - answer.await.map_err(|_| MachineFault::Abandoned.into()) - } - - pub async fn route(&self, op: R::Op) -> Result { - match self.invoke(HostOp::Route(op)).await? { - HostResult::Route(result) => Ok(result), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn before_send( - &self, - wire: WireRequest, - context: RequestContext, - ) -> Result { - let op = HostOp::BeforeSend { - wire: Box::new(wire), - context: Box::new(context), - }; - match self.invoke(op).await? { - HostResult::BeforeSend(wire) => Ok(*wire), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { - match self.invoke(HostOp::Emit(event)).await? { - HostResult::Emitted => Ok(()), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn open(&self, head: R::StreamHead) -> Result { - self.demand(HostOp::Open(head)).await - } - - pub async fn deliver(&self, chunk: R::Chunk) -> Result { - self.demand(HostOp::Deliver(chunk)).await - } - - async fn demand(&self, op: HostOp) -> Result { - match self.invoke(op).await? { - HostResult::Demand(demand) => Ok(demand), - _ => Err(MachineFault::Mismatch.into()), - } - } -} - -enum Execution { - Unstarted(Box) -> ExecuteFuture + Send>), - Running(ExecuteFuture), - Done, -} - -pub struct RouteMachine { - execution: Execution, - ops: mpsc::UnboundedReceiver>, - channel: HostChannel, - reply: Option>>, -} - -impl RouteMachine -where - R::Error: From, -{ - pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { - let (ops_tx, ops) = mpsc::unbounded_channel(); - Self { - execution: Execution::Unstarted(Box::new(execute)), - ops, - channel: HostChannel { ops: ops_tx }, - reply: None, - } - } - - async fn step( - &mut self, - result: Option>, - ) -> Result, R::Error> { - match (self.reply.take(), result) { - (Some(reply), Some(result)) => { - reply - .send(result) - .map_err(|_| MachineFault::Protocol("the call stopped waiting on the host"))?; - } - (None, None) if matches!(self.execution, Execution::Unstarted(_)) => {} - (Some(reply), None) => { - self.reply = Some(reply); - return Err(MachineFault::Protocol("host operation result is required").into()); - } - (None, Some(_)) => { - return Err(MachineFault::Protocol("unexpected host operation result").into()); - } - (None, None) => { - return Err( - MachineFault::Protocol("call cannot be resumed after completion").into(), - ); - } - } - if let Execution::Unstarted(_) = self.execution { - let Execution::Unstarted(start) = - std::mem::replace(&mut self.execution, Execution::Done) - else { - unreachable!() - }; - self.execution = Execution::Running(start(self.channel.clone())); - } - let Execution::Running(future) = &mut self.execution else { - return Err(MachineFault::Protocol("call cannot be resumed after completion").into()); - }; - tokio::select! { - biased; - pending = self.ops.recv() => { - let pending = pending.ok_or(MachineFault::Abandoned)?; - self.reply = Some(pending.reply); - Ok(MachineStep::Host(pending.op)) - } - outcome = future => { - self.execution = Execution::Done; - outcome.map(MachineStep::Complete) - } - } - } -} - -impl Machine for RouteMachine -where - R::Error: From, -{ - type Route = R; - type Complete = R::Response; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(self.step(result)) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.reply = None; - self.execution = Execution::Done; - Box::pin(async move { Err(failure.into_error()) }) - } -} diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs new file mode 100644 index 00000000000..a7c0f3470b2 --- /dev/null +++ b/litellm-rust/crates/host/src/protocol.rs @@ -0,0 +1,17 @@ +/// One public call surface: what a completed call produces, how it fails, what the host +/// projects the caller's request into, and the protocol-specific operations only its host +/// can perform mid-call (token acquisition, for one). +pub trait Protocol: Send + Sync + 'static { + type Response: Send + 'static; + type Error: Clone + Send + Sync + 'static; + /// The caller's request as the host projects it, answered once before anything else. + type Projection: Send + 'static; + /// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through. + /// A protocol with no operations of its own uses `Infallible`. + type Op: Send + 'static; + /// One piece of a streamed response, handed to the caller as it arrives. A protocol + /// that never streams uses `Infallible`. + type Chunk: Send + 'static; + /// What the call knows once a streamed response starts, before its first chunk. + type StreamHead: Send + 'static; +} diff --git a/litellm-rust/crates/host/src/route.rs b/litellm-rust/crates/host/src/route.rs deleted file mode 100644 index 8ab2b125760..00000000000 --- a/litellm-rust/crates/host/src/route.rs +++ /dev/null @@ -1,14 +0,0 @@ -/// One public call surface: what a completed call produces, how it fails, and the -/// route-specific operations only its host can perform (request projection, file reads, -/// token acquisition). -pub trait Route: Send + Sync + 'static { - type Response: Send + 'static; - type Error: Clone + Send + Sync + 'static; - type Op: Send + 'static; - type OpResult: Send + 'static; - /// One piece of a streamed response, handed to the caller as it arrives. A route - /// that never streams uses `Infallible`. - type Chunk: Send + 'static; - /// What the route knows once a streamed response starts, before its first chunk. - type StreamHead: Send + 'static; -} diff --git a/litellm-rust/crates/host/src/run.rs b/litellm-rust/crates/host/src/run.rs index 6a0c08fba68..baa3b58e058 100644 --- a/litellm-rust/crates/host/src/run.rs +++ b/litellm-rust/crates/host/src/run.rs @@ -1,40 +1,28 @@ use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; -use crate::host::{Host, HostOp, HostResult}; +use crate::host::{Host, HostOp}; use crate::machine::{HostFailure, Machine, MachineStep}; -use crate::route::Route; +use crate::protocol::Protocol; /// Drives a machine to completion against an in-process host and emits exactly one /// terminal event. -pub async fn run(mut machine: M, host: &H) -> Result::Error> +pub async fn run( + mut machine: M, + host: &H, +) -> Result::Error> where M: Machine, - H: Host, + H: Host, { let start_time = epoch_seconds(); let _ = host.emit(&CallEvent::Started { start_time }).await; - let mut result = None; let outcome = loop { - let step = match machine.resume(result.take()).await { + let op = match machine.resume().await { Ok(MachineStep::Complete(complete)) => break Ok(complete), Ok(MachineStep::Host(op)) => op, Err(error) => break Err(error), }; - let answer = match step { - HostOp::Route(op) => host.route(op).await.map(HostResult::Route), - HostOp::BeforeSend { wire, context } => host - .before_send(*wire, &context) - .await - .map(|wire| HostResult::BeforeSend(Box::new(wire))), - HostOp::Emit(event) => host - .emit(&CallEvent::Machine(event)) - .await - .map(|()| HostResult::Emitted), - HostOp::Open(head) => host.open(head).await.map(HostResult::Demand), - HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand), - }; - match answer { - Ok(answer) => result = Some(answer), - Err(error) => break machine.interrupt(HostFailure::Error(error)).await, + if let Err(error) = perform(host, op).await { + break machine.interrupt(HostFailure::Error(error)).await; } }; let timing = Timing { @@ -52,44 +40,52 @@ where outcome } +async fn perform>(host: &H, op: HostOp) -> Result<(), R::Error> { + match op { + HostOp::Project(reply) => host + .project() + .await + .map(|projection| reply.send(projection)), + HostOp::Custom(op) => host.custom_op(op).await, + HostOp::BeforeSend { + wire, + context, + reply, + } => host + .before_send(*wire, &context) + .await + .map(|wire| reply.send(wire)), + HostOp::Emit(event, reply) => host + .emit(&CallEvent::Machine(event)) + .await + .map(|()| reply.send(())), + HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)), + HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)), + } +} + #[cfg(test)] mod tests { use std::sync::Mutex; use super::*; - use crate::machine::{Interrupted, Step}; + use crate::host::Reply; + use crate::machine::{CallMachine, MachineFault}; struct Unit; - impl Route for Unit { + impl Protocol for Unit { type Response = (); type Error = &'static str; - type Op = &'static str; - type OpResult = (); + type Projection = (); + type Op = (&'static str, Reply<()>); type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } - struct Scripted { - ops: Vec<&'static str>, - outcome: Result<(), &'static str>, - } - - impl Machine for Scripted { - type Route = Unit; - type Complete = (); - - fn resume(&mut self, _: Option>) -> Step<'_, Self> { - Box::pin(async move { - if !self.ops.is_empty() { - return Ok(MachineStep::Host(HostOp::Route(self.ops.remove(0)))); - } - self.outcome.map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure<&'static str>) -> Interrupted<'_, Self> { - Box::pin(async move { Err(failure.into_error()) }) + impl From for &'static str { + fn from(_: MachineFault) -> Self { + "machine fault" } } @@ -100,12 +96,21 @@ mod tests { } impl Host for Recording { - async fn route(&self, op: &'static str) -> Result<(), &'static str> { - self.seen.lock().unwrap().push(format!("route:{op}")); - match self.fail { - Some(failing) if failing == op => Err("host failed"), - _ => Ok(()), + async fn project(&self) -> Result<(), &'static str> { + self.seen.lock().unwrap().push("project".into()); + Ok(()) + } + + async fn custom_op( + &self, + (op, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + self.seen.lock().unwrap().push(format!("op:{op}")); + if self.fail == Some(op) { + return Err("host failed"); } + reply.send(()); + Ok(()) } async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { @@ -119,21 +124,29 @@ mod tests { } } - fn scripted(ops: &[&'static str], outcome: Result<(), &'static str>) -> Scripted { - Scripted { - ops: ops.to_vec(), - outcome, - } + fn scripted( + ops: &'static [&'static str], + outcome: Result<(), &'static str>, + ) -> CallMachine { + CallMachine::new(move |host| { + Box::pin(async move { + host.project().await?; + for op in ops { + host.custom_op(|reply| (*op, reply)).await?; + } + outcome + }) + }) } #[tokio::test] async fn forwards_every_op_then_emits_one_succeeded() { let host = Recording::default(); - let outcome = run(scripted(&["project", "send"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await; assert_eq!(outcome, Ok(())); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "succeeded"] + ["started", "project", "op:sign", "op:send", "succeeded"] ); } @@ -142,24 +155,32 @@ mod tests { let host = Recording::default(); let outcome = run(scripted(&[], Err("boom")), &host).await; assert_eq!(outcome, Err("boom")); - assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]); + assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]); let host = Recording { fail: Some("send"), ..Recording::default() }; - let outcome = run(scripted(&["project", "send", "never"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await; assert_eq!(outcome, Err("host failed")); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "failed"] + ["started", "project", "op:sign", "op:send", "failed"] ); } struct StartTimes(Mutex>); impl Host for StartTimes { - async fn route(&self, _: &'static str) -> Result<(), &'static str> { + async fn project(&self) -> Result<(), &'static str> { + Ok(()) + } + + async fn custom_op( + &self, + (_, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + reply.send(()); Ok(()) } @@ -178,7 +199,7 @@ mod tests { #[tokio::test] async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() { let host = StartTimes(Mutex::default()); - assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(())); + assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(())); let times = host.0.lock().unwrap(); assert_eq!(times.len(), 2); assert_eq!(times[0], times[1]); diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs index b3df8fc18c8..5d7be8b10df 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs @@ -118,7 +118,6 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "OCR host driver was abandoned".into(), MachineFault::Protocol(message) => format!("OCR {message}"), - MachineFault::Mismatch => "invalid OCR host operation result".into(), }) } } diff --git a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs index f00259984ba..c3377536545 100644 --- a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs @@ -554,6 +554,50 @@ async fn upload_bytes_async( mod tests { use super::*; + #[tokio::test] + async fn v3_body_keeps_explicit_null_options_and_drops_unknown_ones() { + use crate::base_llm::ocr::{handler::OcrClient, transformation::OcrRequestContext}; + + let overrides = + serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true})) + .unwrap(); + let params = ReductoParseV3Config + .map_ocr_params(&overrides, "parse-v3") + .unwrap(); + let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()); + let connection = OcrConnection::default(); + let document = serde_json::from_value( + json!({"type":"document_url","document_url":"reducto://ready.pdf"}), + ) + .unwrap(); + + let body = ReductoParseV3Config + .async_transform_ocr_request( + "parse-v3", + document, + ¶ms, + &[], + OcrRequestContext { + client: &client, + connection: &connection, + }, + ) + .await + .unwrap(); + + assert_eq!( + serde_json::to_value(body).unwrap(), + json!({"input":"reducto://ready.pdf", "formatting":null, "settings":{}}) + ); + let absent = ReductoParseV3Config + .map_ocr_params( + &litellm_core_utils::call_arguments::CallArguments::default(), + "parse-v3", + ) + .unwrap(); + assert_eq!(serde_json::to_value(absent).unwrap(), json!({})); + } + #[test] fn options_preserve_null_and_select_the_provider_fields() { let overrides = serde_json::from_value(json!({ diff --git a/litellm-rust/crates/llms/tests/ocr_handler.rs b/litellm-rust/crates/llms/tests/ocr_handler.rs new file mode 100644 index 00000000000..6e46e6f76d4 --- /dev/null +++ b/litellm-rust/crates/llms/tests/ocr_handler.rs @@ -0,0 +1,79 @@ +use std::time::Duration; + +use litellm_llms::base_llm::ocr::{error::Error, handler::read_response_bytes}; +use rstest::rstest; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, +}; + +/// Answers one request with raw `response` bytes and then holds the connection open, so a +/// read that waits for the rest of an oversized body hangs instead of passing. +async fn read_bounded(response: String, limit: usize) -> Result { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = [0; 4096]; + assert!(socket.read(&mut request).await.unwrap() > 0); + socket.write_all(response.as_bytes()).await.unwrap(); + std::future::pending::<()>().await; + }); + let response = reqwest::Client::new() + .get(format!("http://{address}")) + .send() + .await + .unwrap(); + let result = + tokio::time::timeout(Duration::from_secs(2), read_response_bytes(response, limit)).await; + server.abort(); + result.expect("bounded reads must finish without waiting for the rest of an oversized body") +} + +#[rstest] +#[case::declared("HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh")] +#[case::chunked( + "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n" +)] +#[tokio::test] +async fn a_body_of_exactly_the_limit_is_read(#[case] response: &str) { + assert_eq!(read_bounded(response.into(), 8).await.unwrap(), "abcdefgh"); +} + +#[rstest] +#[case::declared("HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n")] +#[case::chunked("HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n")] +#[tokio::test] +async fn a_body_over_the_limit_is_rejected(#[case] response: &str) { + assert!(matches!( + read_bounded(response.into(), 8).await, + Err(Error::TooLarge { limit: 8 }) + )); +} + +#[rstest] +#[case::declared("Content-Length: 1000000")] +#[case::chunked("Transfer-Encoding: chunked")] +#[tokio::test] +async fn an_oversized_error_keeps_its_status_and_a_bounded_body_without_draining( + #[case] headers: &str, +) { + let prefix = "x".repeat(4096); + let body = match headers.starts_with("Transfer") { + true => format!("{:x}\r\n{prefix}\r\n", prefix.len()), + false => prefix.clone(), + }; + + let error = read_bounded( + format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}"), + prefix.len(), + ) + .await + .unwrap_err(); + + let Error::Transport(litellm_http::transport::Error::Http { status, body }) = error else { + panic!("unexpected error: {error}"); + }; + assert_eq!(status, 429); + assert_eq!(body, prefix); +} diff --git a/litellm-rust/crates/model-catalog/AGENTS.md b/litellm-rust/crates/model-catalog/AGENTS.md new file mode 100644 index 00000000000..9fcbcd57a76 --- /dev/null +++ b/litellm-rust/crates/model-catalog/AGENTS.md @@ -0,0 +1,6 @@ +## Validation + +For `model_prices_and_context_window.json` validation, we should eventually: + +- Remove any schema file like `model_prices_and_context_window.schema.json` +- Stop skipping this crate's tests diff --git a/litellm-rust/crates/model-catalog/Cargo.toml b/litellm-rust/crates/model-catalog/Cargo.toml index ea75c6386d8..0b26e398ac8 100644 --- a/litellm-rust/crates/model-catalog/Cargo.toml +++ b/litellm-rust/crates/model-catalog/Cargo.toml @@ -14,12 +14,8 @@ schemars = { version = "1.0", optional = true } serde.workspace = true serde_json.workspace = true thiserror.workspace = true +time.workspace = true [dev-dependencies] -criterion.workspace = true +jsonschema = { version = "0.55.1", default-features = false } rstest.workspace = true -litellm-model-catalog = { path = ".", features = ["schema"] } - -[[bench]] -name = "catalog" -harness = false diff --git a/litellm-rust/crates/model-catalog/README.md b/litellm-rust/crates/model-catalog/README.md deleted file mode 100644 index 973f7190614..00000000000 --- a/litellm-rust/crates/model-catalog/README.md +++ /dev/null @@ -1,25 +0,0 @@ -# Model catalog - -`litellm-model-catalog` builds an immutable snapshot from caller supplied JSON bytes. It has no network, Python, registration, or refresh behavior. The caller supplies optional source, revision, and ETag provenance. Parse and validation are separate so small synthetic catalogs can use explicit integrity limits - -The parser treats `sample_spec` and `fallback_generalizations` as reserved top level metadata. `fallback_rules()` exposes the typed rule array when present; this crate does not execute regex generalizations. Model entries retain all JSON fields except `aliases`, including unknown fields. `field()` returns `None` for an absent key and a JSON null, false, or zero value for a present key. The returned values are borrowed, so callers cannot mutate the snapshot - -Each entry also deserializes into `ModelInfo`, a typed mirror of `model_prices_and_context_window.schema.json`'s `modelEntry` definition, reachable via `ModelEntry::info()`. All schema fields are optional on `ModelInfo`, including `litellm_provider` which the schema marks required, so small synthetic catalogs still parse. Unknown fields are not part of `ModelInfo`; they remain on `fields()`. Building with the `schema` feature adds `schemars` derives and exposes `model_entry_json_schema()` for emitting the entry's JSON Schema. Parse and validation failures are reported by the `Error` enum in `error.rs`, while catalog logic lives in `catalog.rs` - -The integration tests read the repository's catalog and schema files at test time, assert every entry round-trips through `ModelInfo`, and verify that the generated schema's properties match the repository schema - -Aliases point to their canonical entries. An alias that exactly matches any canonical key is skipped; the first canonical entry claiming an alias wins. Invalid alias lists and nonstring names are skipped and reported by `alias_issues()`. Exact lookup wins. For a case insensitive miss, the last key with the same lowercase spelling wins, following Python's lowercase map built after aliases are appended. This uses Rust Unicode lowercasing, which can differ from Python for unusual Unicode model IDs - -`validate()` counts canonical entries before alias expansion and excludes both reserved keys. It enforces an explicit minimum and backup shrink ratio, with Python defaults of 50 models and 0.5. Parsing rejects nonobject model entries and known fields with the wrong JSON type, but ignores unknown fields. It does not enforce every constraint in the JSON schema, calculate prices, resolve providers, or check provenance authenticity. The caller decides how to handle validation failures - -This snapshot does not represent Python's live mutable `litellm.model_cost`, nested dict and list mutation, or mutation of dicts previously returned by Python APIs. It has no bridge or runtime integration - -## Benchmarks - -`cargo bench -p litellm-model-catalog --bench catalog` measures parsing plus alias indexing and exact lookup. For a local Python baseline on the same fixture, use: - -```sh -python3 -m timeit -s 'import json, pathlib; body = pathlib.Path("../model_prices_and_context_window.json").read_bytes()' 'json.loads(body)' -``` - -Run these commands from `litellm-rust`. Python's command measures JSON loading only, without alias expansion or snapshot construction. The Rust benchmark does not include future Python object materialization, so these numbers are not an end to end runtime comparison diff --git a/litellm-rust/crates/model-catalog/benches/catalog.rs b/litellm-rust/crates/model-catalog/benches/catalog.rs deleted file mode 100644 index d1f51507c2b..00000000000 --- a/litellm-rust/crates/model-catalog/benches/catalog.rs +++ /dev/null @@ -1,21 +0,0 @@ -use criterion::{Criterion, criterion_group, criterion_main}; -use litellm_model_catalog::{Catalog, Provenance}; -use std::hint::black_box; - -fn benchmarks(c: &mut Criterion) { - let body = include_bytes!("../../../../model_prices_and_context_window.json"); - c.bench_function("parse_current_catalog", |b| { - b.iter(|| Catalog::parse(black_box(body), Provenance::default()).unwrap()) - }); - let catalog = Catalog::parse(body, Provenance::default()).unwrap(); - let key = catalog - .model_names() - .next() - .expect("catalog must have a benchmark key"); - c.bench_function("lookup_catalog_key", |b| { - b.iter(|| black_box(&catalog).lookup(black_box(key))) - }); -} - -criterion_group!(benches, benchmarks); -criterion_main!(benches); diff --git a/litellm-rust/crates/model-catalog/src/capabilities.rs b/litellm-rust/crates/model-catalog/src/capabilities.rs new file mode 100644 index 00000000000..66b5f1c5d2e --- /dev/null +++ b/litellm-rust/crates/model-catalog/src/capabilities.rs @@ -0,0 +1,80 @@ +use serde::{Deserialize, Serialize}; + +/// Primary API surface / task type of the model. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum Mode { + AudioSpeech, + AudioTranscription, + Chat, + Completion, + Embedding, + Evaluation, + Guardrail, + ImageEdit, + ImageGeneration, + Moderation, + Ocr, + Realtime, + Rerank, + Responses, + Search, + VectorStore, + VideoGeneration, +} + +/// Reasoning effort level accepted or applied by the model. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum ReasoningEffort { + None, + Minimal, + Low, + Medium, + High, + Xhigh, + Max, +} + +/// Gemini audio generation API the model is served through. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum VertexAiAudioApi { + LyriaPredict, + LyriaInteractions, +} + +/// Audio container format the model can return. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum AudioFormat { + Mp3, + Wav, +} + +/// Input modality the model accepts. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum InputModality { + Text, + Image, + Audio, + Video, +} + +/// Output modality the model can produce. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum OutputModality { + Text, + Image, + Audio, + Video, + Code, +} diff --git a/litellm-rust/crates/model-catalog/src/catalog.rs b/litellm-rust/crates/model-catalog/src/catalog.rs index dc7564f9bee..0b113436a5b 100644 --- a/litellm-rust/crates/model-catalog/src/catalog.rs +++ b/litellm-rust/crates/model-catalog/src/catalog.rs @@ -1,5 +1,6 @@ use crate::error::Error; -use crate::model_info::{FallbackGeneralizations, FallbackRule, ModelInfo}; +use crate::fallback::{FallbackGeneralizations, FallbackRule}; +use crate::model_info::ModelInfo; use indexmap::IndexMap; use serde::Deserialize; use serde_json::{Map, Value}; @@ -14,19 +15,9 @@ pub struct Provenance { #[derive(Clone, Copy, Debug, PartialEq)] pub struct IntegrityLimits { - pub backup_model_count: usize, + pub reference_model_count: usize, pub min_model_count: usize, - pub min_backup_ratio: f64, -} - -impl IntegrityLimits { - pub fn python_defaults(backup_model_count: usize) -> Self { - Self { - backup_model_count, - min_model_count: 50, - min_backup_ratio: 0.5, - } - } + pub min_reference_ratio: f64, } #[derive(Clone, Debug, PartialEq, Eq)] @@ -99,16 +90,11 @@ impl Catalog { } _ => {} } - let Value::Object(ref object) = value else { + let Value::Object(mut fields) = value else { return Err(Error::EntryNotObject { model: name }); }; - let info = ModelInfo::deserialize(object)?; - let Value::Object(mut fields) = value else { - unreachable!("value checked is_object above") - }; - if let Some(aliases) = fields.remove("aliases") - && !aliases.is_null() - { + let info = ModelInfo::deserialize(&fields)?; + if let Some(aliases) = fields.remove("aliases") { match aliases { Value::Array(names) => alias_lists.push((name.clone(), names)), _ => alias_issues.push(AliasIssue::InvalidList { @@ -161,7 +147,9 @@ impl Catalog { } pub fn validate(&self, limits: IntegrityLimits) -> Result<(), Error> { - if !limits.min_backup_ratio.is_finite() || !(0.0..=1.0).contains(&limits.min_backup_ratio) { + if !limits.min_reference_ratio.is_finite() + || !(0.0..=1.0).contains(&limits.min_reference_ratio) + { return Err(Error::InvalidRatio); } let actual = self.entries.len(); @@ -171,13 +159,13 @@ impl Catalog { minimum: limits.min_model_count, }); } - if limits.backup_model_count > 0 - && (actual as f64) < (limits.backup_model_count as f64) * limits.min_backup_ratio + if limits.reference_model_count > 0 + && (actual as f64) < (limits.reference_model_count as f64) * limits.min_reference_ratio { return Err(Error::Shrunk { actual, - backup: limits.backup_model_count, - ratio: limits.min_backup_ratio, + reference: limits.reference_model_count, + ratio: limits.min_reference_ratio, }); } Ok(()) diff --git a/litellm-rust/crates/model-catalog/src/error.rs b/litellm-rust/crates/model-catalog/src/error.rs index 83617312fff..edb6ba1eb11 100644 --- a/litellm-rust/crates/model-catalog/src/error.rs +++ b/litellm-rust/crates/model-catalog/src/error.rs @@ -1,6 +1,5 @@ use thiserror::Error; -/// Failures from parsing or validating a catalog snapshot. #[derive(Debug, Error)] pub enum Error { /// The body is not valid JSON, or a model entry fails typed deserialization. @@ -15,14 +14,14 @@ pub enum Error { /// Canonical entry count is under the configured minimum. #[error("catalog has {actual} models, below minimum {minimum}")] BelowMinimum { actual: usize, minimum: usize }, - /// Canonical entry count is under the configured backup shrink ratio. - #[error("catalog has {actual} models, below {ratio} of backup count {backup}")] + /// Canonical entry count is under the configured reference ratio. + #[error("catalog has {actual} models, below {ratio} of reference count {reference}")] Shrunk { actual: usize, - backup: usize, + reference: usize, ratio: f64, }, - /// The configured minimum backup ratio is not finite or outside `[0, 1]`. - #[error("minimum backup ratio must be finite and between zero and one")] + /// The configured minimum reference ratio is not finite or outside `[0, 1]`. + #[error("minimum reference ratio must be finite and between zero and one")] InvalidRatio, } diff --git a/litellm-rust/crates/model-catalog/src/fallback.rs b/litellm-rust/crates/model-catalog/src/fallback.rs new file mode 100644 index 00000000000..62291a84929 --- /dev/null +++ b/litellm-rust/crates/model-catalog/src/fallback.rs @@ -0,0 +1,23 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::BTreeMap; + +/// One regex rule generalizing unknown model ids to known families. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +pub struct FallbackRule { + pub name: String, + pub pattern: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(flatten)] + pub extra: BTreeMap, +} + +/// Regex rules that generalize unknown model ids to known families; not a model entry. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(deny_unknown_fields)] +pub struct FallbackGeneralizations { + pub rules: Vec, +} diff --git a/litellm-rust/crates/model-catalog/src/lib.rs b/litellm-rust/crates/model-catalog/src/lib.rs index 9c942a5521c..066c2c83c6b 100644 --- a/litellm-rust/crates/model-catalog/src/lib.rs +++ b/litellm-rust/crates/model-catalog/src/lib.rs @@ -1,16 +1,20 @@ +mod capabilities; mod catalog; mod error; +mod fallback; mod model_info; +mod pricing; +mod validation; + +pub use capabilities::*; +pub use catalog::*; +pub use error::*; +pub use fallback::*; +pub use model_info::*; +pub use pricing::*; +pub use validation::*; + #[cfg(feature = "schema")] mod schema; - -pub use catalog::{AliasIssue, Catalog, IntegrityLimits, ModelEntry, ModelMatch, Provenance}; -pub use error::Error; -pub use model_info::{ - AudioFormat, FallbackGeneralizations, FallbackRule, InputModality, Mode, ModelInfo, - OffPeakPricing, OffPeakWindow, OutputModality, ReasoningEffort, SearchContextCostPerQuery, - TieredRate, UtcHours, VertexAiAudioApi, WebSearchBillingUnit, Weekday, -}; - #[cfg(feature = "schema")] -pub use schema::model_entry_json_schema; +pub use schema::*; diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 4a56e1112d1..361cb56e9b1 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -1,673 +1,482 @@ +use crate::capabilities::{ + AudioFormat, InputModality, Mode, OutputModality, ReasoningEffort, VertexAiAudioApi, +}; +use crate::pricing::{OffPeakPricing, SearchContextCostPerQuery, TieredRate, WebSearchBillingUnit}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::BTreeMap; -/// Primary API surface / task type of the model. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum Mode { - AudioSpeech, - AudioTranscription, - Chat, - Completion, - Embedding, - Evaluation, - Guardrail, - ImageEdit, - ImageGeneration, - Moderation, - Ocr, - Realtime, - Rerank, - Responses, - Search, - VectorStore, - VideoGeneration, -} - -/// Reasoning effort level accepted or applied by the model. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum ReasoningEffort { - None, - Minimal, - Low, - Medium, - High, - Xhigh, - Max, -} - -/// Gemini audio generation API the model is served through. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum VertexAiAudioApi { - LyriaPredict, - LyriaInteractions, -} - -/// Whether web search is billed per query or per prompt. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum WebSearchBillingUnit { - PerQuery, - PerPrompt, -} - -/// Audio container format the model can return. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum AudioFormat { - Mp3, - Wav, -} - -/// Input modality the model accepts. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum InputModality { - Text, - Image, - Audio, - Video, -} - -/// Output modality the model can produce. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -pub enum OutputModality { - Text, - Image, - Audio, - Video, - Code, -} - -/// UTC "HH:MM-HH:MM" window, or a list of them; a window may wrap past midnight. -#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(untagged)] -pub enum UtcHours { - Single(String), - Multiple(Vec), -} - -/// ISO-8601 weekday number (1 = Monday .. 7 = Sunday) or English day name. -#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(untagged)] -pub enum Weekday { - Number(u8), - Name(String), -} - -/// One off-peak window entry inside `windows`. -#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(deny_unknown_fields)] -pub struct OffPeakWindow { - pub hours_utc: UtcHours, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub weekdays: Option>, -} - -/// Rates that replace the same-named base fields inside the stated UTC windows. -#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(deny_unknown_fields)] -pub struct OffPeakPricing { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub hours_utc: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub windows: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub weekday_timezone: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input_cost_per_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_reasoning_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_read_input_token_cost: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_creation_input_token_cost: Option, -} - -/// USD cost per web search query, keyed by search context size. -#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(deny_unknown_fields)] -pub struct SearchContextCostPerQuery { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub search_context_size_low: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub search_context_size_medium: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub search_context_size_high: Option, -} - -/// One tier of a context-length or result-count tiered rate. -#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(deny_unknown_fields)] -pub struct TieredRate { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub range: Option<[f64; 2]>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub max_results_range: Option<[f64; 2]>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input_cost_per_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_reasoning_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_read_input_token_cost: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_creation_input_token_cost: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input_cost_per_query: Option, -} - -/// One regex rule generalizing unknown model ids to known families. -#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -pub struct FallbackRule { - pub name: String, - pub pattern: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub description: Option, - #[serde(flatten)] - pub extra: BTreeMap, -} - -/// Regex rules that generalize unknown model ids to known families; not a model entry. -#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(deny_unknown_fields)] -pub struct FallbackGeneralizations { - pub rules: Vec, -} - /// Typed mirror of one catalog model entry. #[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] pub struct ModelInfo { - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub annotation_cost_per_page: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub annotation_cost_per_page_batches: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub audio_transcription_config: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub bedrock_converse_supports_strict_tools: Option, /// Highest reasoning effort the Bedrock output_config accepts for this model. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub bedrock_output_config_effort_ceiling: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_audio_token_cost: Option, /// USD per token written to the provider's prompt cache. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost_above_32k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_128k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_1hr: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_1hr_above_200k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_200k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_272k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_272k_tokens_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_272k_tokens_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_272k_tokens_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_creation_input_token_cost_above_32k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_audio_token_cost: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_image_token_cost: Option, /// USD per prompt token served from the provider's prompt cache. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost_above_32k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_128k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens_priority: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_272k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_272k_tokens_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_272k_tokens_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_272k_tokens_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_read_input_token_cost_above_32k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_512k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub citation_cost_per_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub code_interpreter_cost_per_session: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub comment: Option, /// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub default_reasoning_effort: Option, /// Date the provider deprecates the model, YYYY-MM-DD. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub deprecation_date: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub gemini_audio_only_live: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub gemini_native_audio: Option, /// USD per Grounding with Google Maps request; billed per query or per prompt per web_search_billing_unit. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub google_maps_grounding_cost_per_query: Option, /// USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub guardrail_cost_per_unit: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_audio_per_second: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_audio_per_second_above_128k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_audio_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_audio_token_batches: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_audio_token_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_character: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_character_above_128k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_image: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_image_above_128k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_image_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_image_token_batches: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_pixel: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_query: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_request: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_second: Option, /// USD per prompt token. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token_above_32k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_128k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens_priority: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_272k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_272k_tokens_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_272k_tokens_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_272k_tokens_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input_cost_per_token_above_32k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_512k_tokens: Option, /// USD per prompt token via the provider's batch API. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_batches: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_cache_hit: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_per_second: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_per_second_above_128k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_per_second_above_15s_interval: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_per_second_above_8s_interval: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_token_batches: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub input_dbu_cost_per_token: Option, /// LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub litellm_provider: Option, /// Maximum prompt/context tokens the model accepts. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub max_input_tokens: Option, /// Maximum tokens the model can generate in one response. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub max_output_tokens: Option, /// Legacy field: max output tokens if the provider specifies it, else max input tokens. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub max_tokens: Option, /// Free-form notes about the entry (e.g. pricing derivation). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub metadata: Option>, /// Primary API surface / task type of the model. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub mode: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub ocr_cost_per_credit: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub ocr_cost_per_page: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub ocr_cost_per_page_batches: Option, /// Rates that replace the same-named base fields while the request falls inside the stated UTC windows. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub off_peak_pricing: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_audio_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_character: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_character_above_128k_tokens: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_1024: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_1536: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_512: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_pixel: Option, /// USD per reasoning/thinking token, when billed separately. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_reasoning_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_1080p: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_2k: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_480p: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_4k: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_720p: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_second_768p: Option, /// USD per generated token. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token_above_32k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_128k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens_priority: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_272k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_272k_tokens_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_272k_tokens_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_272k_tokens_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_token_above_32k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_512k_tokens: Option, /// USD per generated token via the provider's batch API. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_batches: Option, /// Flex service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_flex: Option, /// Priority service-tier rate for the same-named base field. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_priority: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_video_per_second: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_video_token: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_dbu_cost_per_token: Option, /// Embedding dimension for embedding models. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub output_vector_size: Option, /// Smallest prefix the provider will actually cache; absent means the provider default applies. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub prompt_cache_min_tokens: Option, /// Provider-internal routing hints (e.g. bedrock_invocation_schema). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub provider_specific_entry: Option>, /// Exact reasoning_effort levels this deployment accepts; wins over supports_* flags. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub reasoning_effort_levels: Option>, /// Multiplier applied to all token costs when served from a non-global Vertex AI endpoint (e.g. 1.10 = +10%). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub regional_endpoint_uplift_multiplier: Option, /// Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub regional_processing_uplift_multiplier_eu: Option, /// Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub regional_processing_uplift_multiplier_us: Option, /// Provider default requests-per-minute limit. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub rpm: Option, /// USD cost per web search query, keyed by search context size. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub search_context_cost_per_query: Option, /// URL of the provider pricing/model page this entry was taken from. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub source: Option, /// Audio container formats the model can return. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supported_audio_formats: Option>, /// OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supported_endpoints: Option>, /// Input modalities the model accepts. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supported_modalities: Option>, /// Output modalities the model can produce. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supported_output_modalities: Option>, /// Cloud regions the model is available in ('global' or region ids). - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supported_regions: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_adaptive_thinking: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_anthropic_compaction: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_anthropic_thinking_payload: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_assistant_prefill: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_audio_input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_audio_output: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_computer_use: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_embedding_image_input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_fast_mode: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_forced_tool_use: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_function_calling: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_image_input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_image_size: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_legacy_thinking: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_low_reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_max_reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_mid_conversation_system: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_minimal_reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_multimodal: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_native_streaming: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_native_structured_output: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_none_reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_nova_canvas_image_edit: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_output_config: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_parallel_function_calling: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_parallel_tool_use_config: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_pdf_input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_prompt_cache_breakpoint: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_prompt_caching: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_reasoning: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_response_schema: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_sampling_params: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_speed: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_system_messages: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_thinking_cache_preservation: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_tool_choice: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_tool_search: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_url_context: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_video_input: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_vision: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_web_search: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub supports_xhigh_reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub thinking_always_on: Option, /// Context-length or result-count tiered rates; each tier's costs apply within its range. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub tiered_pricing: Option>, /// Provider default tokens-per-minute limit. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub tpm: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub use_openai_responses_path: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub uses_embed_content: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub vertex_ai_audio_api: Option, /// Whether web search is billed per query or per prompt. - #[serde(default, skip_serializing_if = "Option::is_none")] + #[serde(skip_serializing_if = "Option::is_none")] pub web_search_billing_unit: Option, } diff --git a/litellm-rust/crates/model-catalog/src/pricing.rs b/litellm-rust/crates/model-catalog/src/pricing.rs new file mode 100644 index 00000000000..b8ed652451c --- /dev/null +++ b/litellm-rust/crates/model-catalog/src/pricing.rs @@ -0,0 +1,97 @@ +use serde::{Deserialize, Serialize}; + +/// Whether web search is billed per query or per prompt. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(rename_all = "snake_case")] +pub enum WebSearchBillingUnit { + PerQuery, + PerPrompt, +} + +/// UTC "HH:MM-HH:MM" window, or a list of them; a window may wrap past midnight. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(untagged)] +pub enum UtcHours { + Single(String), + Multiple(Vec), +} + +/// ISO-8601 weekday number (1 = Monday .. 7 = Sunday) or English day name. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(untagged)] +pub enum Weekday { + Number(u8), + Name(String), +} + +/// One off-peak window entry inside `windows`. +#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(deny_unknown_fields)] +pub struct OffPeakWindow { + pub hours_utc: UtcHours, + #[serde(skip_serializing_if = "Option::is_none")] + pub weekdays: Option>, +} + +/// Rates that replace the same-named base fields inside the stated UTC windows. +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(deny_unknown_fields)] +pub struct OffPeakPricing { + #[serde(skip_serializing_if = "Option::is_none")] + pub hours_utc: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub windows: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub weekday_timezone: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_reasoning_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost: Option, +} + +/// USD cost per web search query, keyed by search context size. +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(deny_unknown_fields)] +pub struct SearchContextCostPerQuery { + #[serde(skip_serializing_if = "Option::is_none")] + pub search_context_size_low: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub search_context_size_medium: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub search_context_size_high: Option, +} + +/// One tier of a context-length or result-count tiered rate. +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[serde(deny_unknown_fields)] +pub struct TieredRate { + #[serde(skip_serializing_if = "Option::is_none")] + pub range: Option<[f64; 2]>, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_results_range: Option<[f64; 2]>, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_reasoning_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_query: Option, +} diff --git a/litellm-rust/crates/model-catalog/src/schema.rs b/litellm-rust/crates/model-catalog/src/schema.rs index 82cfd6c0352..7de988b3ff3 100644 --- a/litellm-rust/crates/model-catalog/src/schema.rs +++ b/litellm-rust/crates/model-catalog/src/schema.rs @@ -1,7 +1,137 @@ -use crate::model_info::ModelInfo; +use schemars::Schema; +use serde_json::{Map, Value, json}; -/// JSON Schema for one catalog model entry, mirroring -/// `model_prices_and_context_window.schema.json`'s `modelEntry` definition. -pub fn model_entry_json_schema() -> schemars::Schema { - schemars::schema_for!(ModelInfo) +/// JSON Schema for one model entry, including registry validation constraints. +pub fn model_entry_json_schema() -> Schema { + let mut schema = serde_json::to_value(schemars::schema_for!(crate::ModelInfo)) + .expect("derived model schema serializes"); + remove_nullable_optional_fields(&mut schema); + decorate_model_entry(&mut schema); + Schema::from( + schema + .as_object() + .expect("derived schema is an object") + .clone(), + ) +} + +/// JSON Schema for the complete model prices registry document. +pub fn registry_json_schema() -> Schema { + let mut entry = model_entry_json_schema().as_value().clone(); + let mut definitions = take_definitions(&mut entry); + entry.as_object_mut().unwrap().remove("$schema"); + definitions.insert("modelEntry".into(), entry); + + let mut fallback = serde_json::to_value(schemars::schema_for!(crate::FallbackGeneralizations)) + .expect("derived fallback schema serializes"); + remove_nullable_optional_fields(&mut fallback); + definitions.extend(take_definitions(&mut fallback)); + fallback.as_object_mut().unwrap().remove("$schema"); + + let root = json!({ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "title": "LiteLLM model prices and context window registry", + "type": "object", + "properties": { + "sample_spec": {"type": "object"}, + "fallback_generalizations": fallback + }, + "additionalProperties": {"$ref": "#/$defs/modelEntry"}, + "$defs": definitions + }); + Schema::from(root.as_object().unwrap().clone()) +} + +fn take_definitions(schema: &mut Value) -> Map { + schema + .as_object_mut() + .unwrap() + .remove("$defs") + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default() +} + +fn remove_nullable_optional_fields(value: &mut Value) { + match value { + Value::Array(values) => values.iter_mut().for_each(remove_nullable_optional_fields), + Value::Object(map) => { + map.values_mut().for_each(remove_nullable_optional_fields); + if let Some(Value::Array(types)) = map.get_mut("type") { + types.retain(|value| value != "null"); + if types.len() == 1 { + let only = types[0].clone(); + map.insert("type".into(), only); + } + } + if let Some(Value::Array(branches)) = map.get_mut("anyOf") { + branches.retain(|branch| branch.get("type") != Some(&Value::String("null".into()))); + if branches.len() == 1 { + let only = branches[0] + .as_object() + .expect("schema branch is an object") + .clone(); + map.remove("anyOf"); + map.extend(only); + } + } + } + _ => {} + } +} + +fn decorate_model_entry(schema: &mut Value) { + let object = schema.as_object_mut().unwrap(); + object.insert("required".into(), json!(["litellm_provider"])); + object.insert("additionalProperties".into(), Value::Bool(true)); + let properties = object + .get_mut("properties") + .unwrap() + .as_object_mut() + .unwrap(); + properties.insert( + "aliases".into(), + json!({"type": "array", "items": {"type": "string"}}), + ); + properties.get_mut("deprecation_date").unwrap()["format"] = json!("date"); + properties.get_mut("deprecation_date").unwrap()["pattern"] = + json!(r"^\d{4}-(0[1-9]|1[0-2])-(0[1-9]|[12]\d|3[01])$"); + + properties.iter_mut().for_each(|(name, property)| { + if name.contains("cost") { + property["minimum"] = json!(0); + } else if name.contains("uplift_multiplier") { + property["minimum"] = json!(1); + } + }); + properties.get_mut("guardrail_cost_per_unit").unwrap()["additionalProperties"]["minimum"] = + json!(0); + + let definitions = object.get_mut("$defs").unwrap().as_object_mut().unwrap(); + for definition in ["OffPeakPricing", "TieredRate", "SearchContextCostPerQuery"] { + let properties = definitions[definition]["properties"] + .as_object_mut() + .unwrap(); + properties.iter_mut().for_each(|(name, property)| { + if name.contains("cost") || definition == "SearchContextCostPerQuery" { + property["minimum"] = json!(0); + } + }); + } + definitions["OffPeakPricing"]["anyOf"] = json!([ + {"required": ["hours_utc"]}, + {"required": ["windows"]} + ]); + definitions["OffPeakPricing"]["properties"]["windows"]["minItems"] = json!(1); + definitions["OffPeakWindow"]["properties"]["weekdays"]["minItems"] = json!(1); + definitions["TieredRate"]["properties"]["range"]["items"]["minimum"] = json!(0); + definitions["TieredRate"]["properties"]["max_results_range"]["items"]["minimum"] = json!(0); + definitions["Weekday"]["anyOf"][0]["minimum"] = json!(1); + definitions["Weekday"]["anyOf"][0]["maximum"] = json!(7); + definitions["Weekday"]["anyOf"][1]["pattern"] = json!( + r"(?i)^(mon|monday|tue|tues|tuesday|wed|wednesday|thu|thur|thurs|thursday|fri|friday|sat|saturday|sun|sunday)$" + ); + let window_pattern = json!(r"^([01]\d|2[0-3]):[0-5]\d-([01]\d|2[0-3]):[0-5]\d$"); + definitions["UtcHours"]["anyOf"][0]["pattern"] = window_pattern.clone(); + definitions["UtcHours"]["anyOf"][1]["items"]["pattern"] = window_pattern; + definitions["UtcHours"]["anyOf"][1]["minItems"] = json!(1); } diff --git a/litellm-rust/crates/model-catalog/src/validation.rs b/litellm-rust/crates/model-catalog/src/validation.rs new file mode 100644 index 00000000000..73f08a865c1 --- /dev/null +++ b/litellm-rust/crates/model-catalog/src/validation.rs @@ -0,0 +1,218 @@ +use std::collections::BTreeSet; + +use serde_json::{Map, Value}; +use thiserror::Error; + +use crate::{AliasIssue, Catalog, ModelInfo, UtcHours, Weekday}; + +/// A registry entry violates the checked-in catalog contract. +#[derive(Debug, Error)] +pub enum RegistryValidationError { + #[error("{reason}")] + Entry { model: String, reason: String }, + #[error("alias issue: {0:?}")] + Alias(AliasIssue), +} + +/// Validate one registry entry without restricting the tolerant catalog reader. +pub fn validate_model_entry(model: &str, value: &Value) -> Result<(), RegistryValidationError> { + validate_entry_inner(model, value).map_err(|reason| RegistryValidationError::Entry { + model: model.to_owned(), + reason, + }) +} + +/// Check every model and alias in a parsed catalog against registry rules. +pub fn validate_registry(catalog: &Catalog) -> Result<(), RegistryValidationError> { + if let Some(issue) = catalog.alias_issues().first() { + return Err(RegistryValidationError::Alias(issue.clone())); + } + catalog.model_names().try_for_each(|name| { + let entry = catalog.lookup(name).expect("catalog name must resolve"); + validate_model_entry(name, &Value::Object(entry.entry.fields().clone())) + }) +} + +fn json_eq(left: &Value, right: &Value) -> bool { + match (left, right) { + (Value::Number(left), Value::Number(right)) => left.as_f64() == right.as_f64(), + (Value::Array(left), Value::Array(right)) => { + left.len() == right.len() && left.iter().zip(right).all(|(a, b)| json_eq(a, b)) + } + (Value::Object(left), Value::Object(right)) => { + left.len() == right.len() + && left + .iter() + .all(|(key, value)| right.get(key).is_some_and(|other| json_eq(value, other))) + } + _ => left == right, + } +} + +fn keys(value: &Map) -> BTreeSet { + value.keys().cloned().collect() +} + +fn symmetric_difference(left: &BTreeSet, right: &BTreeSet) -> BTreeSet { + left.symmetric_difference(right).cloned().collect() +} + +fn validate_entry_inner(model_name: &str, value: &Value) -> Result<(), String> { + let object = value + .as_object() + .ok_or_else(|| format!("{model_name} must be an object"))?; + if let Some(aliases) = object.get("aliases") { + let names = aliases + .as_array() + .ok_or_else(|| format!("{model_name}.aliases must be an array"))?; + if names.iter().any(|name| !name.is_string()) { + return Err(format!("{model_name}.aliases must contain strings")); + } + } + let info: ModelInfo = + serde_json::from_value(value.clone()).map_err(|error| format!("{model_name}: {error}"))?; + if info.litellm_provider.is_none() { + return Err(format!("{model_name}.litellm_provider is required")); + } + validate_dates_and_windows(model_name, &info)?; + let serialized = serde_json::to_value(info).map_err(|error| error.to_string())?; + let mut expected = object.clone(); + expected.remove("aliases"); + if !json_eq(&Value::Object(expected.clone()), &serialized) { + let actual = serialized + .as_object() + .expect("ModelInfo serializes as an object"); + return Err(format!( + "{model_name} has an unknown field, null, or changed value: {:?}", + symmetric_difference(&keys(&expected), &keys(actual)) + )); + } + check_prices(model_name, value) +} + +fn validate_dates_and_windows(model_name: &str, info: &ModelInfo) -> Result<(), String> { + if let Some(date) = &info.deprecation_date { + let format = time::format_description::parse_borrowed::<2>("[year]-[month]-[day]").unwrap(); + time::Date::parse(date, &format) + .map_err(|error| format!("{model_name}.deprecation_date: {error}"))?; + } + let Some(pricing) = &info.off_peak_pricing else { + return Ok(()); + }; + if pricing.hours_utc.is_none() && pricing.windows.is_none() { + return Err(format!( + "{model_name}.off_peak_pricing needs hours or windows" + )); + } + if let Some(hours) = &pricing.hours_utc { + validate_hours(hours)?; + } + if let Some(windows) = &pricing.windows { + if windows.is_empty() { + return Err(format!("{model_name}.off_peak_pricing.windows is empty")); + } + windows.iter().try_for_each(|window| { + validate_hours(&window.hours_utc)?; + if let Some(days) = &window.weekdays + && (days.is_empty() || days.iter().any(|day| !valid_weekday(day))) + { + return Err(format!("{model_name}.off_peak_pricing.weekdays is invalid")); + } + Ok(()) + })?; + } + Ok(()) +} + +fn validate_hours(hours: &UtcHours) -> Result<(), String> { + let values = match hours { + UtcHours::Single(value) => std::slice::from_ref(value), + UtcHours::Multiple(values) => values.as_slice(), + }; + if values.is_empty() || values.iter().any(|value| !valid_utc_window(value)) { + return Err("off_peak_pricing.hours_utc is invalid".into()); + } + Ok(()) +} + +fn valid_utc_window(value: &str) -> bool { + let Some((start, end)) = value.split_once('-') else { + return false; + }; + [start, end].into_iter().all(|clock| { + let Some((hour, minute)) = clock.split_once(':') else { + return false; + }; + hour.len() == 2 + && minute.len() == 2 + && hour.parse::().is_ok_and(|hour| hour < 24) + && minute.parse::().is_ok_and(|minute| minute < 60) + }) +} + +fn valid_weekday(day: &Weekday) -> bool { + match day { + Weekday::Number(number) => (1..=7).contains(number), + Weekday::Name(name) => matches!( + name.to_ascii_lowercase().as_str(), + "mon" + | "monday" + | "tue" + | "tues" + | "tuesday" + | "wed" + | "wednesday" + | "thu" + | "thur" + | "thurs" + | "thursday" + | "fri" + | "friday" + | "sat" + | "saturday" + | "sun" + | "sunday" + ), + } +} + +fn check_prices(path: &str, value: &Value) -> Result<(), String> { + let Some(object) = value.as_object() else { + return Ok(()); + }; + object.iter().try_for_each(|(key, field)| { + let field_path = format!("{path}.{key}"); + if (key.contains("cost") + || path.ends_with(".guardrail_cost_per_unit") + || path.ends_with(".search_context_cost_per_query")) + && let Some(number) = field.as_f64() + && number < 0.0 + { + return Err(format!("{field_path} must be nonnegative")); + } + if key.contains("uplift_multiplier") + && let Some(number) = field.as_f64() + && number < 1.0 + { + return Err(format!("{field_path} must be at least one")); + } + if matches!(key.as_str(), "range" | "max_results_range") + && field.as_array().is_some_and(|values| { + values + .iter() + .any(|value| value.as_f64().is_some_and(|n| n < 0.0)) + }) + { + return Err(format!("{field_path} must be nonnegative")); + } + if matches!(key.as_str(), "metadata" | "provider_specific_entry") { + return Ok(()); + } + match field.as_array() { + Some(items) => items.iter().enumerate().try_for_each(|(index, item)| { + check_prices(&format!("{field_path}[{index}]"), item) + }), + None => check_prices(&field_path, field), + } + }) +} diff --git a/litellm-rust/crates/model-catalog/tests/catalog.rs b/litellm-rust/crates/model-catalog/tests/catalog.rs index bcadd38e908..d87de3c5b3b 100644 --- a/litellm-rust/crates/model-catalog/tests/catalog.rs +++ b/litellm-rust/crates/model-catalog/tests/catalog.rs @@ -43,6 +43,7 @@ fn fixture_catalog() -> Catalog { } #[rstest] +#[ignore] fn preserves_fields_and_metadata(fixture_catalog: Catalog) { let catalog = fixture_catalog; let entry = catalog.lookup("SHORT").unwrap(); @@ -68,6 +69,7 @@ fn preserves_fields_and_metadata(fixture_catalog: Catalog) { } #[rstest] +#[ignore] fn snapshot_does_not_borrow_source() { let mut source = ALPHA_FIXTURE.to_vec(); let catalog = Catalog::parse(&source, Provenance::default()).unwrap(); @@ -84,7 +86,8 @@ fn snapshot_does_not_borrow_source() { #[case("shared", "Second")] #[case("FIRST", "First")] #[case("sHaReD", "Second")] -fn alias_collisions_and_case_fallback_follow_python_order( +#[ignore] +fn alias_collisions_and_case_fallback_follow_entry_order( #[case] lookup: &str, #[case] expected: &str, ) { @@ -117,6 +120,36 @@ fn alias_collisions_and_case_fallback_follow_python_order( ); } +#[test] +#[ignore] +fn json_entry_order_controls_alias_ownership_and_case_fallback() { + let forward = Catalog::parse( + br#"{ + "Alpha":{"aliases":["shared"]}, + "Beta":{"aliases":["shared"]}, + "Foo":{}, + "fOO":{} + }"#, + Provenance::default(), + ) + .unwrap(); + let reversed = Catalog::parse( + br#"{ + "fOO":{}, + "Foo":{}, + "Beta":{"aliases":["shared"]}, + "Alpha":{"aliases":["shared"]} + }"#, + Provenance::default(), + ) + .unwrap(); + + assert_eq!(forward.lookup("shared").unwrap().canonical_key, "Alpha"); + assert_eq!(reversed.lookup("shared").unwrap().canonical_key, "Beta"); + assert_eq!(forward.lookup("foo").unwrap().canonical_key, "fOO"); + assert_eq!(reversed.lookup("foo").unwrap().canonical_key, "Foo"); +} + #[derive(Debug)] enum ValidationOutcome { Ok, @@ -128,36 +161,37 @@ enum ValidationOutcome { #[rstest] #[case( IntegrityLimits { - backup_model_count: 2, + reference_model_count: 2, min_model_count: 1, - min_backup_ratio: 0.5, + min_reference_ratio: 0.5, }, ValidationOutcome::Ok )] #[case( IntegrityLimits { - backup_model_count: 3, + reference_model_count: 3, min_model_count: 1, - min_backup_ratio: 0.5, + min_reference_ratio: 0.5, }, ValidationOutcome::Shrunk )] #[case( IntegrityLimits { - backup_model_count: 0, + reference_model_count: 0, min_model_count: 2, - min_backup_ratio: 0.5, + min_reference_ratio: 0.5, }, ValidationOutcome::BelowMinimum )] #[case( IntegrityLimits { - backup_model_count: 0, + reference_model_count: 0, min_model_count: 0, - min_backup_ratio: f64::NAN, + min_reference_ratio: f64::NAN, }, ValidationOutcome::InvalidRatio )] +#[ignore] fn integrity_uses_canonical_count_and_strict_shrink_boundary( #[case] limits: IntegrityLimits, #[case] expected: ValidationOutcome, @@ -195,6 +229,7 @@ enum MalformedOutcome { br#"{"fallback_generalizations":{},"a":{}}"#, MalformedOutcome::Json )] +#[ignore] fn malformed_input_and_aliases_have_typed_outcomes( #[case] body: &[u8], #[case] expected: MalformedOutcome, @@ -210,6 +245,7 @@ fn malformed_input_and_aliases_have_typed_outcomes( } #[rstest] +#[ignore] fn invalid_aliases_are_reported_not_fatal() { let catalog = Catalog::parse( br#"{"a":{"aliases":"bad"},"b":{"aliases":[9,"ok"]}}"#, @@ -228,7 +264,8 @@ fn invalid_aliases_are_reported_not_fatal() { } #[rstest] -fn parses_current_and_packaged_catalogs_without_pinning_counts( +#[ignore] +fn parses_current_and_packaged_catalogs_against_independent_baseline( current_catalog: Catalog, backup_catalog: Catalog, ) { @@ -236,18 +273,17 @@ fn parses_current_and_packaged_catalogs_without_pinning_counts( assert!(backup_catalog.model_count() > 0); assert!(current_catalog.sample_spec().is_some()); assert!(backup_catalog.sample_spec().is_some()); - assert!( - current_catalog - .validate(IntegrityLimits::python_defaults( - backup_catalog.model_count() - )) - .is_ok() - ); - for name in current_catalog.model_names() { + // Snapshot from 2026-09-23; the backup file mirrors the current file and cannot detect shrinkage. + const REFERENCE_MODEL_COUNT: usize = 4303; + current_catalog + .validate(IntegrityLimits { + reference_model_count: REFERENCE_MODEL_COUNT, + min_model_count: 50, + min_reference_ratio: 0.9, + }) + .unwrap(); + assert!(current_catalog.model_names().all(|name| { let entry = current_catalog.lookup(name).unwrap().entry; - assert_eq!( - entry.info().litellm_provider.is_some(), - entry.field("litellm_provider").is_some() - ); - } + entry.info().litellm_provider.is_some() == entry.field("litellm_provider").is_some() + })); } diff --git a/litellm-rust/crates/model-catalog/tests/registry_validation.rs b/litellm-rust/crates/model-catalog/tests/registry_validation.rs new file mode 100644 index 00000000000..8f555e1f884 --- /dev/null +++ b/litellm-rust/crates/model-catalog/tests/registry_validation.rs @@ -0,0 +1,76 @@ +use std::path::{Path, PathBuf}; + +use litellm_model_catalog::{ + Catalog, FallbackGeneralizations, Provenance, validate_model_entry, validate_registry, +}; +use rstest::{fixture, rstest}; +use serde_json::{Map, Value}; + +#[fixture] +fn repo_root() -> PathBuf { + Path::new(env!("CARGO_MANIFEST_DIR")).join("../../..") +} + +#[rstest] +#[case("model_prices_and_context_window.json")] +#[case("litellm/model_prices_and_context_window_backup.json")] +#[ignore] +fn checked_in_registry_passes_strict_validation(repo_root: PathBuf, #[case] filename: &str) { + let body = std::fs::read(repo_root.join(filename)).unwrap(); + let catalog = Catalog::parse(&body, Provenance::default()).unwrap(); + validate_registry(&catalog).unwrap(); +} + +#[rstest] +#[ignore] +fn fallback_generalizations_are_typed(repo_root: PathBuf) { + let body = std::fs::read(repo_root.join("model_prices_and_context_window.json")).unwrap(); + let document: Map = serde_json::from_slice(&body).unwrap(); + let Some(raw_rules) = document.get("fallback_generalizations") else { + return; + }; + let _: FallbackGeneralizations = serde_json::from_value(raw_rules.clone()).unwrap(); + let catalog = Catalog::parse(&body, Provenance::default()).unwrap(); + assert!( + catalog + .fallback_rules() + .is_some_and(|rules| !rules.is_empty()) + ); +} + +#[rstest] +#[case::missing_provider(serde_json::json!({"mode": "chat"}), "litellm_provider")] +#[case::unknown_field(serde_json::json!({"litellm_provider": "test", "typo": true}), "unknown field")] +#[case::negative_price(serde_json::json!({"litellm_provider": "test", "input_cost_per_token": -1}), "nonnegative")] +#[case::negative_nested_price(serde_json::json!({"litellm_provider": "test", "guardrail_cost_per_unit": {"unit": -1}}), "nonnegative")] +#[case::invalid_mode(serde_json::json!({"litellm_provider": "test", "mode": "invalid"}), "unknown variant")] +#[case::invalid_date(serde_json::json!({"litellm_provider": "test", "deprecation_date": "2026-02-31"}), "deprecation_date")] +#[case::invalid_hours(serde_json::json!({"litellm_provider": "test", "off_peak_pricing": {"hours_utc": "25:00-01:00"}}), "hours_utc")] +#[case::empty_windows(serde_json::json!({"litellm_provider": "test", "off_peak_pricing": {"windows": []}}), "windows is empty")] +#[case::invalid_weekday(serde_json::json!({"litellm_provider": "test", "off_peak_pricing": {"windows": [{"hours_utc": "00:00-01:00", "weekdays": [0]}]}}), "weekdays is invalid")] +#[case::invalid_aliases(serde_json::json!({"litellm_provider": "test", "aliases": ["good", 7]}), "aliases must contain strings")] +#[case::null_aliases(serde_json::json!({"litellm_provider": "test", "aliases": null}), "aliases must be an array")] +#[ignore] +fn registry_validation_rejects_malformed_entries(#[case] entry: Value, #[case] expected: &str) { + assert!( + validate_model_entry("test", &entry) + .unwrap_err() + .to_string() + .contains(expected) + ); +} + +#[test] +#[ignore] +fn checked_in_catalog_and_backup_match() { + let root = Path::new(env!("CARGO_MANIFEST_DIR")).join("../../.."); + let current = std::fs::read(root.join("model_prices_and_context_window.json")).unwrap(); + let backup = + std::fs::read(root.join("litellm/model_prices_and_context_window_backup.json")).unwrap(); + assert_eq!(current, backup); + let catalog = Catalog::parse(¤t, Provenance::default()).unwrap(); + assert!( + catalog.alias_issues().is_empty(), + "invalid registry aliases" + ); +} diff --git a/litellm-rust/crates/model-catalog/tests/schema.rs b/litellm-rust/crates/model-catalog/tests/schema.rs new file mode 100644 index 00000000000..f9a028a78ab --- /dev/null +++ b/litellm-rust/crates/model-catalog/tests/schema.rs @@ -0,0 +1,121 @@ +#![cfg(feature = "schema")] + +use std::collections::BTreeSet; +use std::path::Path; + +use litellm_model_catalog::{model_entry_json_schema, registry_json_schema}; +use rstest::rstest; +use serde_json::{Value, json}; + +fn schema() -> Value { + serde_json::to_value(model_entry_json_schema()).expect("generated schema serializes") +} + +fn registry_validator() -> jsonschema::Validator { + let schema = serde_json::to_value(registry_json_schema()).unwrap(); + jsonschema::options() + .should_validate_formats(true) + .build(&schema) + .expect("generated registry schema is valid") +} + +#[rstest] +#[case("model_prices_and_context_window.json")] +#[case("litellm/model_prices_and_context_window_backup.json")] +#[ignore] +fn generated_registry_schema_validates_checked_in_catalog(#[case] path: &str) { + let root = Path::new(env!("CARGO_MANIFEST_DIR")).join("../../.."); + let catalog: Value = serde_json::from_slice(&std::fs::read(root.join(path)).unwrap()).unwrap(); + let validator = registry_validator(); + let errors: Vec<_> = validator + .iter_errors(&catalog) + .map(|error| error.to_string()) + .collect(); + assert!(errors.is_empty(), "{path}: {errors:?}"); +} + +#[rstest] +#[case(json!({"example": {"litellm_provider": "test"}}))] +#[case(json!({"example": {"litellm_provider": "test", "future_field": true}}))] +#[case(json!({"sample_spec": {"litellm_provider": "placeholder"}}))] +#[ignore] +fn generated_registry_schema_keeps_reader_compatibility(#[case] document: Value) { + assert!(registry_validator().is_valid(&document)); +} + +#[rstest] +#[case::missing_provider(json!({"mode": "chat"}))] +#[case::negative_cost(json!({"litellm_provider": "test", "input_cost_per_token": -1}))] +#[case::negative_guardrail_cost(json!({"litellm_provider": "test", "guardrail_cost_per_unit": {"unit": -1}}))] +#[case::negative_search_cost(json!({"litellm_provider": "test", "search_context_cost_per_query": {"search_context_size_low": -1}}))] +#[case::negative_tier_cost(json!({"litellm_provider": "test", "tiered_pricing": [{"input_cost_per_token": -1}]}))] +#[case::negative_tier_range(json!({"litellm_provider": "test", "tiered_pricing": [{"range": [-1, 2]}]}))] +#[case::low_uplift(json!({"litellm_provider": "test", "regional_endpoint_uplift_multiplier": 0.5}))] +#[case::nullable_cost(json!({"litellm_provider": "test", "input_cost_per_token": null}))] +#[case::invalid_mode(json!({"litellm_provider": "test", "mode": "telepathy"}))] +#[case::invalid_date(json!({"litellm_provider": "test", "deprecation_date": "2026-02-31"}))] +#[case::invalid_hours(json!({"litellm_provider": "test", "off_peak_pricing": {"hours_utc": "25:00-01:00"}}))] +#[case::empty_windows(json!({"litellm_provider": "test", "off_peak_pricing": {"windows": []}}))] +#[case::invalid_weekday(json!({"litellm_provider": "test", "off_peak_pricing": {"windows": [{"hours_utc": "00:00-01:00", "weekdays": [0]}]}}))] +#[case::invalid_aliases(json!({"litellm_provider": "test", "aliases": "wrong"}))] +#[case::non_object_model(json!(4))] +#[ignore] +fn generated_registry_schema_rejects_invalid_entries(#[case] entry: Value) { + assert!(!registry_validator().is_valid(&json!({"example": entry}))); +} + +#[rstest] +#[case("model_prices_and_context_window.json")] +#[case("litellm/model_prices_and_context_window_backup.json")] +#[ignore] +fn generated_schema_covers_catalog_fields(#[case] path: &str) { + let root = Path::new(env!("CARGO_MANIFEST_DIR")).join("../../.."); + let catalog: Value = serde_json::from_slice(&std::fs::read(root.join(path)).unwrap()).unwrap(); + let schema = schema(); + let properties = schema["properties"] + .as_object() + .expect("ModelInfo schema has properties"); + let fields: BTreeSet<&str> = catalog + .as_object() + .expect("catalog is an object") + .iter() + .filter(|(name, _)| *name != "sample_spec" && *name != "fallback_generalizations") + .flat_map(|(_, entry)| entry.as_object().expect("model entry is an object").keys()) + .map(String::as_str) + .filter(|name| *name != "aliases") + .collect(); + let missing: Vec<_> = fields + .into_iter() + .filter(|name| !properties.contains_key(*name)) + .collect(); + + assert!( + missing.is_empty(), + "{path}: fields missing from schema: {missing:?}" + ); +} + +#[rstest] +#[case("Mode", "chat")] +#[case("ReasoningEffort", "high")] +#[case("InputModality", "image")] +#[ignore] +fn generated_schema_includes_enum_values(#[case] definition: &str, #[case] value: &str) { + let schema = schema(); + let variants = schema["$defs"][definition]["enum"] + .as_array() + .expect("enum definition has variants"); + + assert!(variants.iter().any(|variant| variant == value)); +} + +#[test] +#[ignore] +fn generated_schema_includes_nested_pricing_types() { + let schema = schema(); + let definitions = schema["$defs"].as_object().expect("schema has definitions"); + + assert!(definitions.contains_key("OffPeakPricing")); + assert!(definitions.contains_key("TieredRate")); + assert!(definitions.contains_key("UtcHours")); +} diff --git a/litellm-rust/crates/model-catalog/tests/spec_parity.rs b/litellm-rust/crates/model-catalog/tests/spec_parity.rs deleted file mode 100644 index 7296d96f798..00000000000 --- a/litellm-rust/crates/model-catalog/tests/spec_parity.rs +++ /dev/null @@ -1,121 +0,0 @@ -use std::collections::{BTreeSet, HashSet}; -use std::path::{Path, PathBuf}; - -use indexmap::IndexMap; -use litellm_model_catalog::{ - Catalog, FallbackGeneralizations, ModelInfo, Provenance, model_entry_json_schema, -}; -use rstest::{fixture, rstest}; -use serde_json::{Map, Value}; - -#[fixture] -fn repo_root() -> PathBuf { - Path::new(env!("CARGO_MANIFEST_DIR")).join("../../..") -} - -fn json_eq(left: &Value, right: &Value) -> bool { - match (left, right) { - (Value::Number(left), Value::Number(right)) => left.as_f64() == right.as_f64(), - (Value::Array(left), Value::Array(right)) => { - left.len() == right.len() && left.iter().zip(right).all(|(a, b)| json_eq(a, b)) - } - (Value::Object(left), Value::Object(right)) => { - left.len() == right.len() - && left - .iter() - .all(|(key, value)| right.get(key).is_some_and(|other| json_eq(value, other))) - } - _ => left == right, - } -} - -fn keys(value: &Map) -> BTreeSet { - value.keys().cloned().collect() -} - -fn symmetric_difference(left: &BTreeSet, right: &BTreeSet) -> BTreeSet { - left.symmetric_difference(right).cloned().collect() -} - -#[rstest] -#[case("model_prices_and_context_window.json")] -#[case("litellm/model_prices_and_context_window_backup.json")] -fn every_entry_round_trips_through_model_info(repo_root: PathBuf, #[case] filename: &str) { - let body = std::fs::read(repo_root.join(filename)).unwrap(); - let document: IndexMap = serde_json::from_slice(&body).unwrap(); - for (model_name, value) in document { - if matches!( - model_name.as_str(), - "sample_spec" | "fallback_generalizations" - ) { - continue; - } - let object = value - .as_object() - .unwrap_or_else(|| panic!("{model_name} is not an object")); - let info: ModelInfo = serde_json::from_value(value.clone()) - .unwrap_or_else(|error| panic!("{model_name} does not deserialize: {error}")); - let serialized = serde_json::to_value(info).unwrap(); - let serialized_object = serialized - .as_object() - .unwrap_or_else(|| panic!("{model_name} did not serialize as an object")); - let mut expected = object.clone(); - expected.remove("aliases"); - let expected_keys = keys(&expected); - let serialized_keys = keys(serialized_object); - assert_eq!( - expected_keys, - serialized_keys, - "{model_name} key difference: {:?}", - symmetric_difference(&expected_keys, &serialized_keys) - ); - assert!( - json_eq(&Value::Object(expected), &serialized), - "{model_name} changed during ModelInfo round-trip" - ); - } -} - -#[rstest] -fn fallback_generalizations_are_typed(repo_root: PathBuf) { - let body = std::fs::read(repo_root.join("model_prices_and_context_window.json")).unwrap(); - let document: Map = serde_json::from_slice(&body).unwrap(); - let Some(raw_rules) = document.get("fallback_generalizations") else { - return; - }; - let _: FallbackGeneralizations = serde_json::from_value(raw_rules.clone()).unwrap(); - let catalog = Catalog::parse(&body, Provenance::default()).unwrap(); - assert!( - catalog - .fallback_rules() - .is_some_and(|rules| !rules.is_empty()) - ); -} - -#[rstest] -fn generated_schema_properties_match_repo_schema(repo_root: PathBuf) { - let body = - std::fs::read(repo_root.join("model_prices_and_context_window.schema.json")).unwrap(); - let document: Value = serde_json::from_slice(&body).unwrap(); - let repo_entry_properties = document["$defs"]["modelEntry"]["properties"] - .as_object() - .unwrap(); - let generated = serde_json::to_value(model_entry_json_schema()).unwrap(); - let generated_properties = generated["properties"].as_object().unwrap(); - let expected = keys(repo_entry_properties); - let actual = keys(generated_properties); - assert_eq!( - expected, - actual, - "modelEntry property difference: {:?}", - symmetric_difference(&expected, &actual) - ); - - let repo_root_properties = document["properties"].as_object().unwrap(); - let actual_root: HashSet = repo_root_properties.keys().cloned().collect(); - let expected_root: HashSet = ["sample_spec", "fallback_generalizations"] - .into_iter() - .map(str::to_owned) - .collect(); - assert_eq!(actual_root, expected_root); -} diff --git a/litellm-rust/crates/python-bridge/src/logger/machine.rs b/litellm-rust/crates/python-bridge/src/logger/machine.rs index 7234308e67e..54f8f3d4b3f 100644 --- a/litellm-rust/crates/python-bridge/src/logger/machine.rs +++ b/litellm-rust/crates/python-bridge/src/logger/machine.rs @@ -1,9 +1,8 @@ use std::sync::OnceLock; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, Step}, - route::Route, + protocol::Protocol, }; use litellm_tracing::Logger; use pyo3::Python; @@ -23,17 +22,17 @@ impl LoggedMachine { } impl Machine for LoggedMachine { - type Route = M::Route; + type Protocol = M::Protocol; type Complete = M::Complete; - fn resume(&mut self, result: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); - Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result)))) + Box::pin(logger.instrument(logger.scope(|| self.machine.resume()))) } fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure)))) diff --git a/litellm-rust/crates/python-bridge/src/logger/tests.rs b/litellm-rust/crates/python-bridge/src/logger/tests.rs index 9312d4c187c..b65e7d37023 100644 --- a/litellm-rust/crates/python-bridge/src/logger/tests.rs +++ b/litellm-rust/crates/python-bridge/src/logger/tests.rs @@ -1,29 +1,28 @@ use std::{process::Command, task::Poll}; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, - route::Route, + protocol::Protocol, }; use pyo3::{prelude::*, types::PyDict}; struct DiagnosticMachine; -impl Route for DiagnosticMachine { +impl Protocol for DiagnosticMachine { type Response = (); type Error = String; + type Projection = (); type Op = (); - type OpResult = (); type Chunk = (); type StreamHead = (); } impl Machine for DiagnosticMachine { - type Route = Self; + type Protocol = Self; type Complete = (); - fn resume(&mut self, _: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { litellm_tracing::warn!("machine started"); Box::pin(async { tokio::task::yield_now().await; @@ -45,7 +44,7 @@ fn machine_warning(py: Python<'_>) -> PyResult> { let mut machine = super::LoggedMachine::new(DiagnosticMachine); let mut future = Box::pin(async move { machine - .resume(None) + .resume() .await .map_err(pyo3::exceptions::PyValueError::new_err)?; machine diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 9d97094aeda..a253f4f5670 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -1,10 +1,12 @@ +use std::convert::Infallible; + use bytes::Bytes; use litellm_core::messages::{ Error, - route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput}, + route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead}, types::MessagesShaping, }; -use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ProviderSpecificHeaders; use pyo3::{ @@ -80,16 +82,16 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. -pub(super) struct MessagesRouteHost { +pub(super) struct MessagesPythonHost { request: Py, } -impl MessagesRouteHost { +impl MessagesPythonHost { pub(super) fn new(request: Py) -> Self { Self { request } } - fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { let request = self.request.bind(py); let argument = |name: &str| -> PyResult>> { Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) @@ -208,22 +210,21 @@ impl MessagesRouteHost { } } -impl RouteHost for MessagesRouteHost { - type Route = Messages; +impl ProtocolHost for MessagesPythonHost { + type Protocol = Messages; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: MessagesOp, - ) -> Result> { - match op { - MessagesOp::ProjectRequest => self - .project(py, arguments) - .map(|call| MessagesOpResult::Request(Box::new(call))) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))), - } + ) -> Result> { + self.projection(py, arguments) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))) + } + + fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { + match op {} } fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { @@ -237,6 +238,13 @@ impl RouteHost for MessagesRouteHost { } } + fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult> { + py.import(ROUTE_HOST_MODULE)? + .getattr("stream_hidden_params")? + .call1((to_py(py, &head.headers)?,)) + .map(Bound::unbind) + } + fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult> { Ok(PyBytes::new(py, &chunk).into_any().unbind()) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index fd474e6b2d4..dae8623979a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,17 +1,15 @@ mod host; -use host::MessagesRouteHost; +use host::MessagesPythonHost; use litellm_callbacks_legacy_python::{ LegacySurface, PassThroughStream, PublicCall, run_legacy_call, }; -use litellm_core::messages::route::{messages_machine, supports}; +use litellm_core::messages::route::messages_machine; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, }; -use crate::errors::RustBridgeDeclined; - const SURFACE: LegacySurface = LegacySurface { call_type: "anthropic_messages", input_description: "Messages", @@ -28,24 +26,13 @@ fn run_messages( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let model: String = request.getattr("model")?.extract()?; - let provider: Option = request.getattr("custom_llm_provider")?.extract()?; - let stream = request - .getattr("stream")? - .extract::>()? - .unwrap_or(false); - if !supports(&model, provider.as_deref(), stream) { - return Err(RustBridgeDeclined::new_err( - "the Rust Messages route does not serve this provider", - )); - } let secrets = crate::secrets::source(py)?; run_legacy_call( py, SURFACE, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(messages_machine(secrets)), - MessagesRouteHost::new(request.unbind()), + MessagesPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index a928e62d5b7..e6821241c89 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -1,58 +1,38 @@ use std::path::PathBuf; -use bytes::Bytes; -use litellm_core::ocr::types::{OcrDocumentInput, OcrFileContent}; +use litellm_core::ocr::types::OcrDocumentInput; +use litellm_host_python::{PythonFileReader, py_bytes}; use pyo3::{ - exceptions::{PyTypeError, PyValueError}, - gc::{PyTraverseError, PyVisit}, + exceptions::PyValueError, prelude::*, - pybacked::PyBackedBytes, sync::PyOnceLock, types::{PyBytes, PyString, PyType}, }; -#[derive(Debug)] -pub(super) struct PythonFileReader { - reader: Py, - name: Option, +/// A `type='file'` document as projected: paths and bytes are typed inputs already; a +/// file-like object is a reader the projection consumes once every other field is read. +pub(super) enum FileDocumentInput { + Ready(OcrDocumentInput), + Deferred { + reader: PythonFileReader, + mime_type: Option, + }, } -impl PythonFileReader { - pub(super) fn read(&self, py: Python<'_>) -> PyResult { - let value = self.reader.bind(py).call0()?; - let bytes = if value.is_instance_of::() { - Bytes::from(value.extract::()?) - } else if value.is_instance_of::() { - extract_bytes(&value)? - } else { - return Err(PyTypeError::new_err(format!( - "OCR file read must return bytes or str, got {}", - value.get_type(), - ))); - }; - Ok(OcrFileContent { - bytes, - file_name: self.name.clone(), - }) +impl FileDocumentInput { + pub(super) fn resolve(self, py: Python<'_>) -> PyResult { + match self { + Self::Ready(input) => Ok(input), + Self::Deferred { reader, mime_type } => { + let content = reader.read(py)?; + Ok(OcrDocumentInput::Bytes { + bytes: content.bytes, + file_name: content.file_name, + mime_type, + }) + } + } } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.reader) - } -} - -fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult { - if value.is_exact_instance_of::() { - return Ok(Bytes::from_owner(value.extract::()?)); - } - Ok(Bytes::copy_from_slice( - value.extract::()?.as_ref(), - )) -} - -pub(super) struct FileDocumentInput { - pub input: OcrDocumentInput, - pub reader: Option, } impl FromPyObject<'_, '_> for FileDocumentInput { @@ -87,51 +67,31 @@ impl FromPyObject<'_, '_> for FileDocumentInput { } static PATH_LIKE: PyOnceLock> = PyOnceLock::new(); if file.is_instance(PATH_LIKE.import(py, "os", "PathLike")?)? { - return Ok(Self { - input: OcrDocumentInput::Path { - path: file.extract::()?, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Path { + path: file.extract::()?, + mime_type, + })); } if file.is_instance_of::() { - return Ok(Self { - input: OcrDocumentInput::Bytes { - bytes: extract_bytes(&file)?, - file_name: None, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Bytes { + bytes: py_bytes(&file)?, + file_name: None, + mime_type, + })); } - let reader = file - .getattr_opt("read")? - .filter(|value| value.is_callable()); - let Some(reader) = reader else { - return Err(PyValueError::new_err(format!( + match PythonFileReader::from_file_like(&file)? { + Some(reader) => Ok(Self::Deferred { reader, mime_type }), + None => Err(PyValueError::new_err(format!( "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", file.get_type(), - ))); - }; - let name = file - .getattr_opt("name")? - .filter(|value| !value.is_none()) - .map(|value| value.extract::()) - .transpose()?; - Ok(Self { - input: OcrDocumentInput::HostReader { mime_type }, - reader: Some(PythonFileReader { - reader: reader.unbind(), - name, - }), - }) + ))), + } } } #[cfg(test)] mod tests { - use pyo3::types::PyDict; + use pyo3::{exceptions::PyTypeError, types::PyDict}; use super::*; @@ -141,6 +101,13 @@ mod tests { locals } + fn ready(input: FileDocumentInput) -> OcrDocumentInput { + match input { + FileDocumentInput::Ready(input) => input, + FileDocumentInput::Deferred { .. } => panic!("expected a ready document"), + } + } + #[test] fn extraction_validates_required_file_and_optional_mime_type() { Python::initialize(); @@ -167,13 +134,19 @@ mod tests { .unwrap(); assert!(error.is_instance_of::(py)); assert!(error.to_string().contains("bare str")); + let error = py + .eval(c"{'file': object()}", None, None) + .unwrap() + .extract::() + .err() + .unwrap(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("Unsupported file input type")); let document = py .eval(c"{'file': b'abc', 'mime_type': 'image/png'}", None, None) .unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: None, @@ -199,7 +172,7 @@ class Reader: return b'abc' reader = Reader() document = {'file': reader, 'mime_type': 7} -reader_document = {'file': reader} +reader_document = {'file': reader, 'mime_type': 'application/pdf'} path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_type': 'image/png'}", ); let document = locals.get_item("document").unwrap().unwrap(); @@ -208,10 +181,6 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ let document = locals.get_item("reader_document").unwrap().unwrap(); let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!( - input.input, - OcrDocumentInput::HostReader { mime_type: None } - ); let reads = || { locals .get_item("reader") @@ -223,21 +192,20 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ .unwrap() }; assert_eq!(reads(), 0); - let content = input.reader.unwrap().read(py).unwrap(); + let resolved = input.resolve(py).unwrap(); assert_eq!(reads(), 1); assert_eq!( - content, - OcrFileContent { + resolved, + OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: Some("scan.png".into()), + mime_type: Some("application/pdf".into()), } ); let document = locals.get_item("path_document").unwrap().unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Path { path: PathBuf::from("/nonexistent/ocr-projection-test.pdf"), mime_type: Some("image/png".into()), @@ -245,97 +213,4 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ ); }); } - - #[test] - fn reader_results_are_normalized_and_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = KeyError('reader failed') -class Raising: - def read(self): - raise failure -class Text: - def read(self): - return 'héllo' -class Wrong: - def read(self): - return 7 -raising = {'file': Raising()} -text = {'file': Text()} -wrong = {'file': Wrong()}", - ); - let reader = |name: &str| { - locals - .get_item(name) - .unwrap() - .unwrap() - .extract::() - .unwrap() - .reader - .unwrap() - }; - let error = reader("raising").read(py).unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - assert_eq!( - reader("text").read(py).unwrap().bytes.as_ref(), - "héllo".as_bytes() - ); - let error = reader("wrong").read(py).unwrap_err(); - assert!(error.is_instance_of::(py)); - assert!(error.to_string().contains("bytes or str")); - }); - } - - #[rstest::rstest] - #[case::read("read")] - #[case::name("name")] - fn reader_attribute_failures_keep_their_identity(#[case] attribute: &str) { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = LookupError('file property failed') -class File: - def __getattribute__(self, name): - if name == attribute: - raise failure - return super().__getattribute__(name) - name = 'scan.pdf' - def read(self): - return b'abc' -document = {'file': File()}", - ); - locals.set_item("attribute", attribute).unwrap(); - let error = locals - .get_item("document") - .unwrap() - .unwrap() - .extract::() - .err() - .unwrap(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { - Python::initialize(); - let (bytes, pointer) = Python::attach(|py| { - let value = PyBytes::new(py, b"document bytes"); - let pointer = value.as_bytes().as_ptr() as usize; - (extract_bytes(value.as_any()).unwrap(), pointer) - }); - assert_eq!(bytes.as_ptr() as usize, pointer); - assert_eq!(bytes.as_ref(), b"document bytes"); - } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 8bf99cd355f..dc01ced15a0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,6 +1,6 @@ use litellm_auth::ResolvedCredential; -use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult}; -use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py}; +use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection}; +use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -20,14 +20,15 @@ enum OcrHostData { Released, } -/// The Python side of the OCR route: projects the prepared arguments, reads file-like -/// documents, acquires Azure AD tokens, and builds the public response and exception. -pub(super) struct OcrRouteHost { +/// The Python side of the OCR route: projects the prepared arguments (reading a file-like +/// document as it goes), acquires Azure AD tokens, and builds the public response and +/// exception. +pub(super) struct OcrPythonHost { request: Py, data: OcrHostData, } -impl OcrRouteHost { +impl OcrPythonHost { pub(super) fn new(request: Py) -> Self { Self { request, @@ -42,14 +43,6 @@ impl OcrRouteHost { } } - fn read_document(&self, py: Python<'_>) -> PyResult { - self.handles()? - .reader - .as_ref() - .ok_or_else(missing_state)? - .read(py) - } - fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { self.handles()? .azure_ad_token_provider @@ -58,30 +51,21 @@ impl OcrRouteHost { .acquire(py) } - fn answer( + fn projection( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> PyResult { - match op { - OcrOp::ProjectRequest => { - let OcrHostData::Unprojected = self.data else { - return Err(missing_state()); - }; - let (request, handles) = project_request(self.request.bind(py), arguments)?; - let caller_token = handles.azure_ad_token_provider.is_some(); - self.data = OcrHostData::Projected(Box::new(handles)); - Ok(OcrOpResult::Request { - request: Box::new(request), - caller_token, - }) - } - OcrOp::ReadDocument => self.read_document(py).map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => self - .acquire_azure_ad_token(py) - .map(OcrOpResult::AzureAdToken), - } + ) -> PyResult { + let OcrHostData::Unprojected = self.data else { + return Err(missing_state()); + }; + let (request, handles) = project_request(self.request.bind(py), arguments)?; + let caller_token = handles.azure_ad_token_provider.is_some(); + self.data = OcrHostData::Projected(Box::new(handles)); + Ok(OcrProjection { + request, + caller_token, + }) } fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr { @@ -104,20 +88,28 @@ impl OcrRouteHost { } } -impl RouteHost for OcrRouteHost { - type Route = Ocr; +impl ProtocolHost for OcrPythonHost { + type Protocol = Ocr; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> Result> { - self.answer(py, arguments, op) + ) -> Result> { + self.projection(py, arguments) .map_err(|error| InvokeError::Python(self.map_failure(py, error))) } + fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError> { + match op { + OcrOp::AcquireAzureAdToken(reply) => self + .acquire_azure_ad_token(py) + .map(|token| reply.send(token)) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))), + } + } + fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { py.import("litellm.rust_bridge.ocr.route_host")? .getattr("response")? @@ -125,6 +117,10 @@ impl RouteHost for OcrRouteHost { .map(Bound::unbind) } + fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult> { + match head {} + } + fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { match chunk {} } @@ -148,13 +144,10 @@ impl RouteHost for OcrRouteHost { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.request)?; - if let OcrHostData::Projected(handles) = &self.data { - if let Some(reader) = &handles.reader { - reader.traverse(visit)?; - } - if let Some(provider) = &handles.azure_ad_token_provider { - provider.traverse(visit)?; - } + if let OcrHostData::Projected(handles) = &self.data + && let Some(provider) = &handles.azure_ad_token_provider + { + provider.traverse(visit)?; } Ok(()) } @@ -205,20 +198,13 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrRouteHost::new(py.None()); - let projected = host.invoke(py, &kwargs, OcrOp::ProjectRequest).unwrap(); - assert!(matches!( - projected, - OcrOpResult::Request { - caller_token: true, - .. - } - )); + let mut host = OcrPythonHost::new(py.None()); + assert!(host.project(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); + let (reply, _) = litellm_host::host::reply(); assert_eq!( - host.invoke(py, &PyDict::new(py), OcrOp::AcquireAzureAdToken) - .is_ok(), + host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(), succeeds ); let alive = || { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b54316b258b..a4f2bf851d7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -5,7 +5,7 @@ mod project; use std::sync::LazyLock; -use host::OcrRouteHost; +use host::OcrPythonHost; use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::{provider_config, route::ocr_machine}; @@ -69,7 +69,7 @@ fn run_ocr( if asynchronous { ASYNC_SURFACE } else { SURFACE }, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(ocr_machine(client)), - OcrRouteHost::new(request.unbind()), + OcrPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 697b935a1d4..be43a1b7711 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -8,19 +8,15 @@ use litellm_llms::base_llm::ocr::error::Error; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; use serde_json::{Map, Value}; -use super::{ - document::{FileDocumentInput, PythonFileReader}, - errors::to_pyerr as ocr_error_to_pyerr, -}; +use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr}; use crate::{ credentials::{self, CallerTokenProvider}, marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}, }; -/// What the host keeps after projection: the caller's callables that answer the document -/// read and token operations, and the provider name the failure mapping reports. +/// What the host keeps after projection: the caller's token callable that answers the +/// token operation, and the provider name the failure mapping reports. pub(super) struct OcrHostHandles { - pub reader: Option, pub azure_ad_token_provider: Option, pub provider: &'static str, } @@ -104,13 +100,11 @@ impl ProjectedDocument { Ok(Self::File(document.extract()?)) } - fn into_parts(self) -> PyResult<(OcrDocumentInput, Option)> { + /// Reads a file-like document now, so it runs after every other argument was read. + fn resolve(self, py: Python<'_>) -> PyResult { match self { - Self::File(FileDocumentInput { input, reader }) => Ok((input, reader)), - Self::Other(wire) => Ok(( - decode_document(wire).map_err(ocr_error_to_pyerr)?.into(), - None, - )), + Self::File(file) => file.resolve(py), + Self::Other(wire) => Ok(decode_document(wire).map_err(ocr_error_to_pyerr)?.into()), } } } @@ -136,24 +130,25 @@ pub(super) fn project_request( .chain(["api_key", "api_base", "extra_headers"]), )?; let azure_ad_token_provider = credentials::azure_ad_token_provider(kwargs)?; - let (document, reader) = document.into_parts()?; + let api_base = arguments.api_base()?; + let extra_headers = arguments.extra_headers()?; + let timeout_seconds = arguments.timeout_seconds()?; let wire = OcrWireRequest { model, - document, + document: document.resolve(request.py())?, api_key, - api_base: arguments.api_base()?, + api_base, custom_llm_provider, - extra_headers: arguments.extra_headers()?, + extra_headers, optional_params, input_sources, - timeout_seconds: arguments.timeout_seconds()?, + timeout_seconds, }; let request = decode_request_input(wire).map_err(ocr_error_to_pyerr)?; let provider = request.provider_name(); Ok(( request, OcrHostHandles { - reader, azure_ad_token_provider, provider, }, @@ -180,10 +175,8 @@ mod tests { OcrArguments { request, kwargs } } - fn project_document( - document: &Bound<'_, PyAny>, - ) -> PyResult<(OcrDocumentInput, Option)> { - ProjectedDocument::project(document)?.into_parts() + fn project_document(document: &Bound<'_, PyAny>) -> PyResult { + ProjectedDocument::project(document)?.resolve(document.py()) } fn url_document(url: &str) -> OcrDocumentInput { @@ -342,8 +335,11 @@ kwargs = {} }); } + /// A reader that rewrites the request while it runs shows which arguments projection + /// read before it and which after: every other argument is read first, and the read + /// happens exactly once. #[test] - fn document_readers_are_not_consumed_during_projection() { + fn document_readers_are_read_once_after_every_other_argument() { Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); @@ -351,17 +347,24 @@ kwargs = {} py, c" class Request: - api_base = 'original' + model = 'mistral/mistral-ocr-latest' + custom_llm_provider = None + api_key = None + api_base = 'https://original.example.com' + extra_headers = {'x-source': 'original'} timeout = 1 @property def document(self): return document class Reader: + reads = 0 def read(self): - Request.api_base = 'mutated' + Reader.reads += 1 + Request.api_base = 'https://mutated.example.com' + Request.extra_headers = {'x-source': 'mutated'} Request.timeout = 9 return b'abc' -document = {'type': 'file', 'file': Reader()} +document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'} request = Request() kwargs = {} ", @@ -373,15 +376,38 @@ kwargs = {} .unwrap() .cast_into::() .unwrap(); - let arguments = arguments(&request, &kwargs); - let document = arguments.document().unwrap(); - let (input, reader) = project_document(&document).unwrap(); - assert_eq!(input, OcrDocumentInput::HostReader { mime_type: None }); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("original")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(1.0)); - reader.unwrap().read(py).unwrap(); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0)); + let (projected, _) = project_request(&request, &kwargs).unwrap(); + assert_eq!( + py.eval(c"Reader.reads", Some(&locals), Some(&locals)) + .unwrap() + .extract::() + .unwrap(), + 1 + ); + assert_eq!( + projected.document, + OcrDocumentInput::Bytes { + bytes: b"abc".as_slice().into(), + file_name: None, + mime_type: Some("application/pdf".into()), + } + ); + assert_eq!( + projected + .credentials + .api_base + .as_ref() + .map(|base| base.value().as_str()), + Some("https://original.example.com") + ); + assert_eq!( + projected.transport.extra_headers, + [("x-source".to_string(), "original".to_string())] + ); + assert_eq!( + projected.transport.timeout, + Some(std::time::Duration::from_secs(1)) + ); }); } @@ -396,16 +422,14 @@ kwargs = {} None, ) .unwrap(); - let (input, reader) = project_document(&file).unwrap(); assert_eq!( - input, + project_document(&file).unwrap(), OcrDocumentInput::Bytes { bytes: b"%PDF-1.4".as_slice().into(), file_name: None, mime_type: Some("application/pdf".into()), } ); - assert!(reader.is_none()); let original = py .eval( @@ -414,8 +438,10 @@ kwargs = {} None, ) .unwrap(); - let (input, _) = project_document(&original).unwrap(); - assert_eq!(input, url_document("https://example.com/a.pdf")); + assert_eq!( + project_document(&original).unwrap(), + url_document("https://example.com/a.pdf") + ); }); } @@ -617,7 +643,7 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let (input, _) = project_document(&document).unwrap(); + let input = project_document(&document).unwrap(); assert!(matches!(input, OcrDocumentInput::Bytes { .. })); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); diff --git a/litellm-rust/crates/testkit/Cargo.toml b/litellm-rust/crates/testkit/Cargo.toml new file mode 100644 index 00000000000..98a36a1e87f --- /dev/null +++ b/litellm-rust/crates/testkit/Cargo.toml @@ -0,0 +1,32 @@ +[package] +name = "litellm-testkit" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +publish = false + +[dependencies] +flate2.workspace = true +reqwest.workspace = true +serde.workspace = true +semver.workspace = true +serde_json.workspace = true +sha2.workspace = true +tar.workspace = true +target-lexicon.workspace = true +thiserror.workspace = true +tokio = { workspace = true, features = ["fs", "process"] } +zip.workspace = true + +[dev-dependencies] +flate2.workspace = true +rstest.workspace = true +sha2.workspace = true +tar.workspace = true +target-lexicon.workspace = true +futures-util.workspace = true +tempfile.workspace = true +tokio.workspace = true +toml = "0.9" +zip.workspace = true diff --git a/litellm-rust/crates/testkit/src/agent/claude.rs b/litellm-rust/crates/testkit/src/agent/claude.rs new file mode 100644 index 00000000000..6870fff6bf8 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/claude.rs @@ -0,0 +1,181 @@ +use std::collections::BTreeMap; +use std::path::Path; + +use semver::Version; +use serde::Deserialize; + +use super::{ + Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines, + path_string, +}; +use crate::install::release::parse; +use crate::install::{Packaging, Release}; +use crate::{Error, Fetch, Target}; + +const RELEASES: &str = "https://downloads.claude.ai/claude-code-releases"; + +pub struct ClaudeCode; + +#[derive(Deserialize)] +struct Manifest { + platforms: BTreeMap, +} + +#[derive(Deserialize)] +struct Platform { + checksum: String, +} + +impl Install for ClaudeCode { + fn binary(&self) -> &'static str { + "claude" + } + + async fn release( + &self, + fetch: &impl Fetch, + version: &Version, + target: Target, + ) -> Result { + let manifest_url = format!("{RELEASES}/{version}/manifest.json"); + let manifest: Manifest = parse(&manifest_url, &fetch.get(&manifest_url).await?)?; + let key = format!( + "{}-{}{}", + target.os_name(), + target.arch_name(), + target.musl_suffix() + ); + let platform = manifest + .platforms + .get(&key) + .ok_or_else(|| Error::AssetNotFound(key.clone()))?; + Ok(Release { + url: format!("{RELEASES}/{version}/{key}/claude"), + asset: key, + sha256: platform.checksum.clone(), + packaging: Packaging::Bare, + }) + } +} + +impl Configure for ClaudeCode { + fn configure( + &self, + _version: &Version, + settings: &Settings, + home: &Path, + ) -> Result { + if settings.wire != Wire::Messages { + return Err(Error::UnsupportedWire { + agent: "claude", + wire: settings.wire, + }); + } + Ok(LaunchSpec { + env: env([ + ("HOME", path_string(home)), + ("CLAUDE_CONFIG_DIR", path_string(&home.join(".claude"))), + ("ANTHROPIC_BASE_URL", settings.base_url.clone()), + ("ANTHROPIC_AUTH_TOKEN", settings.api_key.clone()), + ("ANTHROPIC_MODEL", settings.model.clone()), + ("DISABLE_AUTOUPDATER", "1".to_owned()), + ]), + files: BTreeMap::new(), + }) + } +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum Event { + Assistant { + message: AssistantMessage, + }, + Result(Finished), + #[serde(other)] + Other, +} + +#[derive(Deserialize)] +struct AssistantMessage { + content: Vec, +} + +#[derive(Deserialize)] +struct Block { + #[serde(rename = "type")] + kind: String, + name: Option, +} + +#[derive(Deserialize)] +struct Finished { + is_error: bool, + result: Option, + usage: Option, +} + +#[derive(Deserialize)] +struct TokenUsage { + input_tokens: u64, + output_tokens: u64, +} + +impl Drive for ClaudeCode { + fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec { + let base = [ + "-p", + &prompt.text, + "--output-format", + "stream-json", + "--verbose", + "--model", + &settings.model, + ]; + let tools = ["--allowedTools", "Bash,Read,Write,Edit"]; + base.into_iter() + .chain(tools.into_iter().filter(|_| prompt.allow_tools)) + .map(str::to_owned) + .collect() + } + + fn parse(&self, _version: &Version, stdout: &str) -> Outcome { + let events: Vec = json_lines(stdout).collect(); + let tool_calls = events + .iter() + .filter_map(|event| match event { + Event::Assistant { message } => Some(&message.content), + _ => None, + }) + .flatten() + .filter(|block| block.kind == "tool_use") + .filter_map(|block| block.name.clone()) + .collect(); + let finished = events.into_iter().find_map(|event| match event { + Event::Result(finished) => Some(finished), + _ => None, + }); + let Some(finished) = finished else { + return Outcome { + tool_calls, + ..Outcome::default() + }; + }; + let result = finished.result.unwrap_or_default(); + let (text, errors) = if finished.is_error { + (String::new(), vec![result]) + } else { + (result, Vec::new()) + }; + Outcome { + text, + tool_calls, + usage: finished.usage.map_or_else(Usage::default, |usage| Usage { + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + }), + errors, + exit_code: None, + } + } +} diff --git a/litellm-rust/crates/testkit/src/agent/codex.rs b/litellm-rust/crates/testkit/src/agent/codex.rs new file mode 100644 index 00000000000..6749e81a471 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/codex.rs @@ -0,0 +1,174 @@ +use std::collections::BTreeMap; +use std::path::{Path, PathBuf}; + +use semver::Version; +use serde::Deserialize; + +use super::{ + Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines, + path_string, quoted, v1, +}; +use crate::install::release::github_release; +use crate::install::{Packaging, Release}; +use crate::target::{Arch, Os}; +use crate::{Error, Fetch, Target}; + +const RELEASES: &str = "https://api.github.com/repos/openai/codex/releases/tags"; + +pub struct Codex; + +fn triple(target: Target) -> String { + let arch = match target.arch { + Arch::Aarch64 => "aarch64", + Arch::X86_64 => "x86_64", + }; + match target.os { + Os::Macos => format!("{arch}-apple-darwin"), + Os::Linux => format!("{arch}-unknown-linux-musl"), + } +} + +impl Install for Codex { + fn binary(&self) -> &'static str { + "codex" + } + + async fn release( + &self, + fetch: &impl Fetch, + version: &Version, + target: Target, + ) -> Result { + let triple = triple(target); + github_release( + fetch, + RELEASES, + &format!("rust-v{version}"), + &format!("codex-{triple}.tar.gz"), + Packaging::TarGz { + member: format!("codex-{triple}"), + }, + ) + .await + } +} + +impl Configure for Codex { + fn configure( + &self, + _version: &Version, + settings: &Settings, + home: &Path, + ) -> Result { + if settings.wire != Wire::Responses { + return Err(Error::UnsupportedWire { + agent: "codex", + wire: settings.wire, + }); + } + let config = format!( + "model = {model}\nmodel_provider = \"litellm\"\n\n[model_providers.litellm]\nname = \"LiteLLM\"\nbase_url = {base_url}\nenv_key = \"LITELLM_API_KEY\"\nwire_api = \"responses\"\n", + model = quoted(&settings.model), + base_url = quoted(&v1(settings)), + ); + Ok(LaunchSpec { + env: env([ + ("HOME", path_string(home)), + ("CODEX_HOME", path_string(&home.join(".codex"))), + ("LITELLM_API_KEY", settings.api_key.clone()), + ]), + files: BTreeMap::from([(PathBuf::from(".codex/config.toml"), config)]), + }) + } +} + +#[derive(Deserialize)] +enum EventKind { + #[serde(rename = "item.completed")] + ItemCompleted, + #[serde(rename = "turn.completed")] + TurnCompleted, + #[serde(rename = "turn.failed")] + TurnFailed, + #[serde(other)] + Other, +} + +#[derive(Deserialize)] +struct Event { + #[serde(rename = "type")] + kind: EventKind, + item: Option, + usage: Option, + error: Option, +} + +#[derive(Deserialize)] +struct Item { + #[serde(rename = "type")] + kind: String, + text: Option, +} + +#[derive(Deserialize)] +struct TokenUsage { + input_tokens: u64, + output_tokens: u64, +} + +#[derive(Deserialize)] +struct Failure { + message: String, +} + +const NON_TOOL_ITEMS: [&str; 3] = ["agent_message", "reasoning", "error"]; + +impl Drive for Codex { + fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec { + let sandbox = ["--sandbox", "workspace-write"]; + ["exec", "--json", "--skip-git-repo-check"] + .into_iter() + .chain(sandbox.into_iter().filter(|_| prompt.allow_tools)) + .chain([prompt.text.as_str()]) + .map(str::to_owned) + .collect() + } + + fn parse(&self, _version: &Version, stdout: &str) -> Outcome { + let events: Vec = json_lines(stdout).collect(); + let items: Vec<&Item> = events + .iter() + .filter(|event| matches!(event.kind, EventKind::ItemCompleted)) + .filter_map(|event| event.item.as_ref()) + .collect(); + Outcome { + text: items + .iter() + .rev() + .find(|item| item.kind == "agent_message") + .and_then(|item| item.text.clone()) + .unwrap_or_default(), + tool_calls: items + .iter() + .filter(|item| !NON_TOOL_ITEMS.contains(&item.kind.as_str())) + .map(|item| item.kind.clone()) + .collect(), + usage: events + .iter() + .filter(|event| matches!(event.kind, EventKind::TurnCompleted)) + .filter_map(|event| event.usage.as_ref()) + .map(|usage| Usage { + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + }) + .fold(Usage::default(), |total, turn| total + turn), + errors: events + .iter() + .filter(|event| matches!(event.kind, EventKind::TurnFailed)) + .filter_map(|event| event.error.as_ref()) + .map(|failure| failure.message.clone()) + .collect(), + exit_code: None, + } + } +} diff --git a/litellm-rust/crates/testkit/src/agent/configure.rs b/litellm-rust/crates/testkit/src/agent/configure.rs new file mode 100644 index 00000000000..a095aacc92f --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/configure.rs @@ -0,0 +1,69 @@ +use std::collections::BTreeMap; +use std::path::{Path, PathBuf}; + +use semver::Version; + +use crate::Error; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Wire { + ChatCompletions, + Messages, + Responses, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Settings { + pub base_url: String, + pub api_key: String, + pub model: String, + pub wire: Wire, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct LaunchSpec { + pub env: BTreeMap, + pub files: BTreeMap, +} + +impl LaunchSpec { + pub fn write_files(&self, home: &Path) -> std::io::Result<()> { + self.files.iter().try_for_each(|(relative, contents)| { + let path = home.join(relative); + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent)?; + } + std::fs::write(path, contents) + }) + } +} + +pub trait Configure { + fn configure( + &self, + version: &Version, + settings: &Settings, + home: &Path, + ) -> Result; +} + +pub(crate) fn env( + pairs: impl IntoIterator, +) -> BTreeMap { + pairs + .into_iter() + .map(|(key, value)| (key.to_owned(), value)) + .collect() +} + +pub(crate) fn path_string(path: &Path) -> String { + path.to_string_lossy().into_owned() +} + +pub(crate) fn quoted(value: &str) -> String { + serde_json::Value::from(value).to_string() +} + +pub(crate) fn v1(settings: &Settings) -> String { + format!("{}/v1", settings.base_url.trim_end_matches('/')) +} diff --git a/litellm-rust/crates/testkit/src/agent/drive.rs b/litellm-rust/crates/testkit/src/agent/drive.rs new file mode 100644 index 00000000000..2c238843ed7 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/drive.rs @@ -0,0 +1,57 @@ +use std::ops::Add; + +use semver::Version; + +use crate::Settings; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Prompt { + pub text: String, + pub allow_tools: bool, +} + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct Usage { + pub input_tokens: u64, + pub output_tokens: u64, +} + +impl Add for Usage { + type Output = Self; + + fn add(self, other: Self) -> Self { + Self { + input_tokens: self.input_tokens + other.input_tokens, + output_tokens: self.output_tokens + other.output_tokens, + } + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct Outcome { + pub text: String, + pub tool_calls: Vec, + pub usage: Usage, + pub errors: Vec, + pub exit_code: Option, +} + +impl Outcome { + pub fn succeeded(&self) -> bool { + self.exit_code == Some(0) && self.errors.is_empty() + } +} + +pub trait Drive { + fn args(&self, version: &Version, settings: &Settings, prompt: &Prompt) -> Vec; + + fn parse(&self, version: &Version, stdout: &str) -> Outcome; +} + +pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>( + stdout: &'a str, +) -> impl Iterator + 'a { + stdout + .lines() + .filter_map(|line| serde_json::from_str(line).ok()) +} diff --git a/litellm-rust/crates/testkit/src/agent/install.rs b/litellm-rust/crates/testkit/src/agent/install.rs new file mode 100644 index 00000000000..3f1a0f8950d --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/install.rs @@ -0,0 +1,17 @@ +use std::future::Future; + +use semver::Version; + +use crate::install::Release; +use crate::{Error, Fetch, Target}; + +pub trait Install: Sync { + fn binary(&self) -> &'static str; + + fn release( + &self, + fetch: &impl Fetch, + version: &Version, + target: Target, + ) -> impl Future> + Send; +} diff --git a/litellm-rust/crates/testkit/src/agent/mod.rs b/litellm-rust/crates/testkit/src/agent/mod.rs new file mode 100644 index 00000000000..03a70583f66 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/mod.rs @@ -0,0 +1,20 @@ +mod claude; +mod codex; +mod configure; +mod drive; +mod install; +mod opencode; + +pub use claude::ClaudeCode; +pub use codex::Codex; +pub use configure::{Configure, LaunchSpec, Settings, Wire}; +pub use drive::{Drive, Outcome, Prompt, Usage}; +pub use install::Install; +pub use opencode::Opencode; + +pub(crate) use configure::{env, path_string, quoted, v1}; +pub(crate) use drive::json_lines; + +pub trait Agent: Install + Configure + Drive {} + +impl Agent for T {} diff --git a/litellm-rust/crates/testkit/src/agent/opencode.rs b/litellm-rust/crates/testkit/src/agent/opencode.rs new file mode 100644 index 00000000000..a1b2fa01eef --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/opencode.rs @@ -0,0 +1,187 @@ +use std::collections::BTreeMap; +use std::path::{Path, PathBuf}; + +use semver::Version; +use serde::Deserialize; + +use super::{ + Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines, + path_string, v1, +}; +use crate::install::release::github_release; +use crate::install::{Packaging, Release}; +use crate::target::Os; +use crate::{Error, Fetch, Target}; + +const RELEASES: &str = "https://api.github.com/repos/sst/opencode/releases/tags"; + +pub struct Opencode; + +impl Install for Opencode { + fn binary(&self) -> &'static str { + "opencode" + } + + async fn release( + &self, + fetch: &impl Fetch, + version: &Version, + target: Target, + ) -> Result { + let stem = format!( + "opencode-{}-{}{}", + target.os_name(), + target.arch_name(), + target.musl_suffix() + ); + let member = "opencode".to_owned(); + let (asset, packaging) = match target.os { + Os::Macos => (format!("{stem}.zip"), Packaging::Zip { member }), + Os::Linux => (format!("{stem}.tar.gz"), Packaging::TarGz { member }), + }; + github_release(fetch, RELEASES, &format!("v{version}"), &asset, packaging).await + } +} + +impl Configure for Opencode { + fn configure( + &self, + _version: &Version, + settings: &Settings, + home: &Path, + ) -> Result { + let npm = match settings.wire { + Wire::ChatCompletions => "@ai-sdk/openai-compatible", + Wire::Responses => "@ai-sdk/openai", + Wire::Messages => "@ai-sdk/anthropic", + }; + let config = serde_json::json!({ + "$schema": "https://opencode.ai/config.json", + "model": format!("litellm/{}", settings.model), + "provider": { + "litellm": { + "npm": npm, + "name": "LiteLLM", + "options": { "baseURL": v1(settings), "apiKey": settings.api_key }, + "models": { settings.model.clone(): { "name": settings.model } }, + } + }, + }); + Ok(LaunchSpec { + env: env([ + ("HOME", path_string(home)), + ("XDG_CONFIG_HOME", path_string(&home.join(".config"))), + ("XDG_DATA_HOME", path_string(&home.join(".local/share"))), + ("OPENCODE_DISABLE_AUTOUPDATE", "true".to_owned()), + ]), + files: BTreeMap::from([( + PathBuf::from(".config/opencode/opencode.json"), + config.to_string(), + )]), + }) + } +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum Event { + Text { + part: TextPart, + }, + ToolUse { + part: ToolPart, + }, + StepFinish { + part: StepFinish, + }, + Error { + error: Failure, + }, + #[serde(other)] + Other, +} + +#[derive(Deserialize)] +struct TextPart { + text: String, +} + +#[derive(Deserialize)] +struct ToolPart { + tool: String, +} + +#[derive(Deserialize)] +struct StepFinish { + tokens: Tokens, +} + +#[derive(Deserialize)] +struct Tokens { + input: u64, + output: u64, +} + +#[derive(Deserialize)] +struct Failure { + name: String, + data: Option, +} + +#[derive(Deserialize)] +struct FailureData { + message: Option, +} + +impl Drive for Opencode { + fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec { + ["run", "--format", "json", &prompt.text] + .map(str::to_owned) + .to_vec() + } + + fn parse(&self, _version: &Version, stdout: &str) -> Outcome { + let events: Vec = json_lines(stdout).collect(); + Outcome { + text: events + .iter() + .rev() + .find_map(|event| match event { + Event::Text { part } => Some(part.text.clone()), + _ => None, + }) + .unwrap_or_default(), + tool_calls: events + .iter() + .filter_map(|event| match event { + Event::ToolUse { part } => Some(part.tool.clone()), + _ => None, + }) + .collect(), + usage: events + .iter() + .filter_map(|event| match event { + Event::StepFinish { part } => Some(Usage { + input_tokens: part.tokens.input, + output_tokens: part.tokens.output, + }), + _ => None, + }) + .fold(Usage::default(), |total, step| total + step), + errors: events + .iter() + .filter_map(|event| match event { + Event::Error { error } => Some( + error + .data + .as_ref() + .and_then(|data| data.message.clone()) + .unwrap_or_else(|| error.name.clone()), + ), + _ => None, + }) + .collect(), + exit_code: None, + } + } +} diff --git a/litellm-rust/crates/testkit/src/error.rs b/litellm-rust/crates/testkit/src/error.rs new file mode 100644 index 00000000000..03520d827f1 --- /dev/null +++ b/litellm-rust/crates/testkit/src/error.rs @@ -0,0 +1,56 @@ +use std::io; +use std::path::PathBuf; + +use thiserror::Error; + +use crate::Wire; + +#[derive(Debug, Error)] +pub enum Error { + #[error("unsupported target {0}")] + UnsupportedTarget(String), + #[error("{0} is not a plain x.y.z release version")] + InvalidVersion(String), + #[error("request to {url} failed")] + Request { + url: String, + #[source] + source: reqwest::Error, + }, + #[error("{url} answered with status {status}")] + Status { url: String, status: u16 }, + #[error("release metadata at {url} is malformed")] + Metadata { + url: String, + #[source] + source: serde_json::Error, + }, + #[error("release has no asset named {0}")] + AssetNotFound(String), + #[error("release publishes no sha256 for {0}")] + MissingChecksum(String), + #[error("sha256 mismatch for {asset}: expected {expected}, got {actual}")] + ChecksumMismatch { + asset: String, + expected: String, + actual: String, + }, + #[error("archive does not contain {0}")] + ArchiveMemberNotFound(String), + #[error("archive is unreadable")] + Archive(#[source] io::Error), + #[error("zip archive is unreadable")] + Zip(#[from] zip::result::ZipError), + #[error("{binary} reports version '{reported}', expected {expected}")] + VersionMismatch { + binary: PathBuf, + expected: String, + reported: String, + }, + #[error("{agent} cannot talk to the gateway over {wire:?}")] + UnsupportedWire { agent: &'static str, wire: Wire }, + #[error("agent did not finish within {0:?}")] + Timeout(std::time::Duration), + #[error("io failure")] + Io(#[from] io::Error), +} diff --git a/litellm-rust/crates/testkit/src/install/archive.rs b/litellm-rust/crates/testkit/src/install/archive.rs new file mode 100644 index 00000000000..c8d08f66f47 --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/archive.rs @@ -0,0 +1,52 @@ +use std::io::{Cursor, Read}; + +use flate2::read::GzDecoder; +use sha2::{Digest, Sha256}; + +use super::release::Packaging; +use crate::Error; + +pub(crate) fn verify_sha256(asset: &str, expected: &str, bytes: &[u8]) -> Result<(), Error> { + let actual = format!("{:x}", Sha256::digest(bytes)); + if actual.eq_ignore_ascii_case(expected) { + return Ok(()); + } + Err(Error::ChecksumMismatch { + asset: asset.to_owned(), + expected: expected.to_owned(), + actual, + }) +} + +pub(crate) fn extract_binary(packaging: &Packaging, bytes: &[u8]) -> Result, Error> { + match packaging { + Packaging::Bare => Ok(bytes.to_vec()), + Packaging::TarGz { member } => extract_tar_gz(member, bytes), + Packaging::Zip { member } => extract_zip(member, bytes), + } +} + +fn extract_tar_gz(member: &str, bytes: &[u8]) -> Result, Error> { + let mut archive = tar::Archive::new(GzDecoder::new(bytes)); + for entry in archive.entries().map_err(Error::Archive)? { + let mut entry = entry.map_err(Error::Archive)?; + let path = entry.path().map_err(Error::Archive)?; + if path.file_name().is_some_and(|name| name == member) { + let mut binary = Vec::new(); + entry.read_to_end(&mut binary).map_err(Error::Archive)?; + return Ok(binary); + } + } + Err(Error::ArchiveMemberNotFound(member.to_owned())) +} + +fn extract_zip(member: &str, bytes: &[u8]) -> Result, Error> { + let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?; + let mut file = archive.by_name(member).map_err(|error| match error { + zip::result::ZipError::FileNotFound => Error::ArchiveMemberNotFound(member.to_owned()), + other => Error::Zip(other), + })?; + let mut binary = Vec::new(); + file.read_to_end(&mut binary).map_err(Error::Archive)?; + Ok(binary) +} diff --git a/litellm-rust/crates/testkit/src/install/fetch.rs b/litellm-rust/crates/testkit/src/install/fetch.rs new file mode 100644 index 00000000000..73008f7a0da --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/fetch.rs @@ -0,0 +1,55 @@ +use std::future::Future; + +use crate::Error; + +pub trait Fetch: Sync { + fn get(&self, url: &str) -> impl Future, Error>> + Send; +} + +pub struct HttpFetch { + client: reqwest::Client, + github_token: Option, +} + +impl HttpFetch { + pub fn new(github_token: Option) -> Self { + Self { + client: reqwest::Client::new(), + github_token, + } + } + + pub fn from_env() -> Self { + Self::new(std::env::var("GITHUB_TOKEN").ok()) + } +} + +impl Fetch for HttpFetch { + async fn get(&self, url: &str) -> Result, Error> { + let request = self + .client + .get(url) + .header("user-agent", "litellm-testkit") + .header("accept", "application/json, application/octet-stream"); + let request = match ( + &self.github_token, + url.starts_with("https://api.github.com/"), + ) { + (Some(token), true) => request.bearer_auth(token), + _ => request, + }; + let request_error = |source| Error::Request { + url: url.to_owned(), + source, + }; + let response = request.send().await.map_err(request_error)?; + let status = response.status(); + if !status.is_success() { + return Err(Error::Status { + url: url.to_owned(), + status: status.as_u16(), + }); + } + Ok(response.bytes().await.map_err(request_error)?.to_vec()) + } +} diff --git a/litellm-rust/crates/testkit/src/install/mod.rs b/litellm-rust/crates/testkit/src/install/mod.rs new file mode 100644 index 00000000000..1104bcec102 --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/mod.rs @@ -0,0 +1,118 @@ +mod archive; +mod fetch; +pub(crate) mod release; + +use std::os::unix::fs::PermissionsExt; +use std::path::{Path, PathBuf}; +use std::process::Stdio; +use std::sync::atomic::{AtomicU64, Ordering}; + +use semver::Version; +use tokio::fs; +use tokio::process::Command; + +use crate::{Error, Install, Target}; +use archive::{extract_binary, verify_sha256}; + +static STAGING_COUNTER: AtomicU64 = AtomicU64::new(0); + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Installed { + pub version: Version, + pub binary: PathBuf, +} + +pub struct Installer { + fetch: F, + cache_root: PathBuf, + target: Target, +} + +impl Installer { + pub fn new(fetch: F, cache_root: impl Into, target: Target) -> Self { + Self { + fetch, + cache_root: cache_root.into(), + target, + } + } + + pub async fn install( + &self, + agent: &impl Install, + version: &Version, + ) -> Result { + validate_release(version)?; + let dir = self + .cache_root + .join(agent.binary()) + .join(version.to_string()); + let binary = dir.join(agent.binary()); + let installed = Installed { + version: version.clone(), + binary: binary.clone(), + }; + if fs::try_exists(&binary).await? && probe_version(&binary, version).await.is_ok() { + return Ok(installed); + } + + let release = agent.release(&self.fetch, version, self.target).await?; + let archive = self.fetch.get(&release.url).await?; + verify_sha256(&release.asset, &release.sha256, &archive)?; + let contents = extract_binary(&release.packaging, &archive)?; + + fs::create_dir_all(&dir).await?; + let staging = dir.join(format!( + ".{}.{}.{}.partial", + agent.binary(), + std::process::id(), + STAGING_COUNTER.fetch_add(1, Ordering::Relaxed) + )); + fs::write(&staging, contents).await?; + fs::set_permissions(&staging, std::fs::Permissions::from_mode(0o755)).await?; + fs::rename(&staging, &binary).await?; + + match probe_version(&binary, version).await { + Ok(()) => Ok(installed), + Err(error) => { + fs::remove_file(&binary).await?; + Err(error) + } + } + } +} + +fn validate_release(version: &Version) -> Result<(), Error> { + if version.pre.is_empty() && version.build.is_empty() { + return Ok(()); + } + Err(Error::InvalidVersion(version.to_string())) +} + +async fn probe_version(binary: &Path, expected: &Version) -> Result<(), Error> { + let home = std::env::temp_dir(); + let output = Command::new(binary) + .arg("--version") + .env_clear() + .env("HOME", home) + .env("DISABLE_AUTOUPDATER", "1") + .stdin(Stdio::null()) + .output() + .await?; + let stdout = String::from_utf8_lossy(&output.stdout); + if stdout + .split_whitespace() + .filter_map(|token| Version::parse(token).ok()) + .any(|reported| &reported == expected) + { + return Ok(()); + } + Err(Error::VersionMismatch { + binary: binary.to_owned(), + expected: expected.to_string(), + reported: stdout.trim().to_owned(), + }) +} + +pub use fetch::{Fetch, HttpFetch}; +pub use release::{Packaging, Release}; diff --git a/litellm-rust/crates/testkit/src/install/release.rs b/litellm-rust/crates/testkit/src/install/release.rs new file mode 100644 index 00000000000..a21b9a14f1b --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/release.rs @@ -0,0 +1,65 @@ +use serde::Deserialize; + +use crate::{Error, Fetch}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum Packaging { + Bare, + TarGz { member: String }, + Zip { member: String }, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Release { + pub asset: String, + pub url: String, + pub sha256: String, + pub packaging: Packaging, +} + +#[derive(Deserialize)] +struct GithubRelease { + assets: Vec, +} + +#[derive(Deserialize)] +struct GithubAsset { + name: String, + digest: Option, + browser_download_url: String, +} + +pub(crate) async fn github_release( + fetch: &impl Fetch, + releases_url: &str, + tag: &str, + asset_name: &str, + packaging: Packaging, +) -> Result { + let url = format!("{releases_url}/{tag}"); + let release: GithubRelease = parse(&url, &fetch.get(&url).await?)?; + let asset = release + .assets + .into_iter() + .find(|asset| asset.name == asset_name) + .ok_or_else(|| Error::AssetNotFound(asset_name.to_owned()))?; + let sha256 = asset + .digest + .as_deref() + .and_then(|digest| digest.strip_prefix("sha256:")) + .ok_or_else(|| Error::MissingChecksum(asset_name.to_owned()))? + .to_owned(); + Ok(Release { + asset: asset.name, + url: asset.browser_download_url, + sha256, + packaging, + }) +} + +pub(crate) fn parse Deserialize<'de>>(url: &str, body: &[u8]) -> Result { + serde_json::from_slice(body).map_err(|source| Error::Metadata { + url: url.to_owned(), + source, + }) +} diff --git a/litellm-rust/crates/testkit/src/lib.rs b/litellm-rust/crates/testkit/src/lib.rs new file mode 100644 index 00000000000..9ea6123a176 --- /dev/null +++ b/litellm-rust/crates/testkit/src/lib.rs @@ -0,0 +1,15 @@ +mod agent; +mod error; +mod install; +mod session; +mod target; + +pub use agent::{ + Agent, ClaudeCode, Codex, Configure, Drive, Install, LaunchSpec, Opencode, Outcome, Prompt, + Settings, Usage, Wire, +}; +pub use error::Error; +pub use install::{Fetch, HttpFetch, Installed, Installer, Packaging, Release}; +pub use semver::Version; +pub use session::Session; +pub use target::{Arch, Os, Target}; diff --git a/litellm-rust/crates/testkit/src/session.rs b/litellm-rust/crates/testkit/src/session.rs new file mode 100644 index 00000000000..6b06e756cca --- /dev/null +++ b/litellm-rust/crates/testkit/src/session.rs @@ -0,0 +1,76 @@ +use std::collections::BTreeMap; +use std::path::PathBuf; +use std::process::Stdio; +use std::time::Duration; + +use semver::Version; +use tokio::process::Command; +use tokio::time::timeout; + +use crate::{Configure, Drive, Error, Installed, Outcome, Prompt, Settings}; + +const STDERR_LIMIT_CHARS: usize = 2000; + +pub struct Session { + binary: PathBuf, + home: PathBuf, + version: Version, + settings: Settings, + env: BTreeMap, +} + +impl Session { + pub fn prepare( + agent: &impl Configure, + installed: &Installed, + settings: Settings, + home: impl Into, + ) -> Result { + let home = home.into(); + let spec = agent.configure(&installed.version, &settings, &home)?; + spec.write_files(&home)?; + Ok(Self { + binary: installed.binary.clone(), + home, + version: installed.version.clone(), + settings, + env: spec.env, + }) + } + + pub async fn run( + &self, + agent: &impl Drive, + prompt: &Prompt, + limit: Duration, + ) -> Result { + let child = Command::new(&self.binary) + .args(agent.args(&self.version, &self.settings, prompt)) + .env_clear() + .env("PATH", "/usr/bin:/bin") + .envs(&self.env) + .current_dir(&self.home) + .stdin(Stdio::null()) + .kill_on_drop(true) + .output(); + let output = timeout(limit, child) + .await + .map_err(|_| Error::Timeout(limit))??; + let parsed = agent.parse(&self.version, &String::from_utf8_lossy(&output.stdout)); + let failed_silently = !output.status.success() && parsed.errors.is_empty(); + Ok(Outcome { + errors: if failed_silently { + vec![ + String::from_utf8_lossy(&output.stderr) + .chars() + .take(STDERR_LIMIT_CHARS) + .collect(), + ] + } else { + parsed.errors + }, + exit_code: output.status.code(), + ..parsed + }) + } +} diff --git a/litellm-rust/crates/testkit/src/target.rs b/litellm-rust/crates/testkit/src/target.rs new file mode 100644 index 00000000000..a9d4b012d52 --- /dev/null +++ b/litellm-rust/crates/testkit/src/target.rs @@ -0,0 +1,69 @@ +use target_lexicon::{Architecture, Environment, OperatingSystem, Triple}; + +use crate::Error; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Os { + Macos, + Linux, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Arch { + Aarch64, + X86_64, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct Target { + pub os: Os, + pub arch: Arch, + pub musl: bool, +} + +impl Target { + pub fn host() -> Result { + Self::try_from(&Triple::host()) + } + + pub(crate) const fn os_name(self) -> &'static str { + match self.os { + Os::Macos => "darwin", + Os::Linux => "linux", + } + } + + pub(crate) const fn arch_name(self) -> &'static str { + match self.arch { + Arch::Aarch64 => "arm64", + Arch::X86_64 => "x64", + } + } + + pub(crate) const fn musl_suffix(self) -> &'static str { + if self.musl { "-musl" } else { "" } + } +} + +impl TryFrom<&Triple> for Target { + type Error = Error; + + fn try_from(triple: &Triple) -> Result { + let unsupported = || Error::UnsupportedTarget(triple.to_string()); + let os = match triple.operating_system { + OperatingSystem::Darwin(_) | OperatingSystem::MacOSX(_) => Os::Macos, + OperatingSystem::Linux => Os::Linux, + _ => return Err(unsupported()), + }; + let arch = match triple.architecture { + Architecture::Aarch64(_) => Arch::Aarch64, + Architecture::X86_64 => Arch::X86_64, + _ => return Err(unsupported()), + }; + Ok(Self { + os, + arch, + musl: triple.environment == Environment::Musl, + }) + } +} diff --git a/litellm-rust/crates/testkit/tests/configure.rs b/litellm-rust/crates/testkit/tests/configure.rs new file mode 100644 index 00000000000..ca3587c3474 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/configure.rs @@ -0,0 +1,133 @@ +use std::path::Path; + +use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire}; +use rstest::rstest; + +fn settings(wire: Wire) -> Settings { + Settings { + base_url: "http://localhost:4000/".to_owned(), + api_key: "sk-test \"quoted\"".to_owned(), + model: "some-model".to_owned(), + wire, + } +} + +fn version() -> Version { + Version::new(1, 2, 3) +} + +#[rstest] +#[case(&ClaudeCode, Wire::Messages)] +#[case(&Codex, Wire::Responses)] +#[case(&Opencode, Wire::ChatCompletions)] +fn every_agent_runs_inside_the_given_home(#[case] agent: &impl Configure, #[case] wire: Wire) { + let home = Path::new("/scratch/home"); + + let spec = agent.configure(&version(), &settings(wire), home).unwrap(); + + assert_eq!(spec.env["HOME"], "/scratch/home"); + assert!( + spec.env + .iter() + .filter(|(key, _)| key.ends_with("_HOME") || key.as_str() == "CLAUDE_CONFIG_DIR") + .all(|(_, value)| value.starts_with("/scratch/home")) + ); + assert!(spec.files.keys().all(|path| path.is_relative())); +} + +#[rstest] +#[case::claude_code(&ClaudeCode, &[Wire::ChatCompletions, Wire::Responses])] +#[case::codex(&Codex, &[Wire::ChatCompletions, Wire::Messages])] +fn wires_an_agent_cannot_speak_are_refused( + #[case] agent: &impl Configure, + #[case] refused: &[Wire], +) { + refused.iter().for_each(|wire| { + let result = agent.configure(&version(), &settings(*wire), Path::new("/h")); + + assert!(matches!(result, Err(Error::UnsupportedWire { wire: got, .. }) if got == *wire)); + }); +} + +#[test] +fn claude_code_points_at_the_gateway_root_with_the_key_and_model() { + let spec = ClaudeCode + .configure(&version(), &settings(Wire::Messages), Path::new("/h")) + .unwrap(); + + assert_eq!(spec.env["ANTHROPIC_BASE_URL"], "http://localhost:4000/"); + assert_eq!(spec.env["ANTHROPIC_AUTH_TOKEN"], "sk-test \"quoted\""); + assert_eq!(spec.env["ANTHROPIC_MODEL"], "some-model"); +} + +#[test] +fn codex_config_is_valid_toml_routing_the_responses_api_to_the_gateway() { + let dir = tempfile::tempdir().unwrap(); + let spec = Codex + .configure(&version(), &settings(Wire::Responses), dir.path()) + .unwrap(); + spec.write_files(dir.path()).unwrap(); + + let config: toml::Table = + toml::from_str(&std::fs::read_to_string(dir.path().join(".codex/config.toml")).unwrap()) + .unwrap(); + let provider = &config["model_providers"]["litellm"]; + + assert_eq!(config["model"].as_str(), Some("some-model")); + assert_eq!(config["model_provider"].as_str(), Some("litellm")); + assert_eq!( + provider["base_url"].as_str(), + Some("http://localhost:4000/v1") + ); + assert_eq!(provider["wire_api"].as_str(), Some("responses")); + let key_var = provider["env_key"].as_str().unwrap(); + assert_eq!(spec.env[key_var], "sk-test \"quoted\""); +} + +#[rstest] +#[case(Wire::ChatCompletions)] +#[case(Wire::Responses)] +#[case(Wire::Messages)] +fn opencode_config_is_valid_json_registering_the_gateway_model(#[case] wire: Wire) { + let dir = tempfile::tempdir().unwrap(); + let spec = Opencode + .configure(&version(), &settings(wire), dir.path()) + .unwrap(); + spec.write_files(dir.path()).unwrap(); + + let config: serde_json::Value = serde_json::from_str( + &std::fs::read_to_string(dir.path().join(".config/opencode/opencode.json")).unwrap(), + ) + .unwrap(); + let provider = &config["provider"]["litellm"]; + + assert_eq!(config["model"], "litellm/some-model"); + assert_eq!(provider["options"]["baseURL"], "http://localhost:4000/v1"); + assert_eq!(provider["options"]["apiKey"], "sk-test \"quoted\""); + assert!(provider["models"]["some-model"].is_object()); +} + +#[test] +fn opencode_uses_a_different_provider_package_for_every_wire() { + let package = |wire| { + let dir = tempfile::tempdir().unwrap(); + let spec = Opencode + .configure(&version(), &settings(wire), dir.path()) + .unwrap(); + let config: serde_json::Value = + serde_json::from_str(spec.files.values().next().unwrap()).unwrap(); + config["provider"]["litellm"]["npm"] + .as_str() + .unwrap() + .to_owned() + }; + let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package); + + assert_eq!( + packages + .iter() + .collect::>() + .len(), + packages.len() + ); +} diff --git a/litellm-rust/crates/testkit/tests/install.rs b/litellm-rust/crates/testkit/tests/install.rs new file mode 100644 index 00000000000..7edaf27de6c --- /dev/null +++ b/litellm-rust/crates/testkit/tests/install.rs @@ -0,0 +1,262 @@ +mod support; + +use std::str::FromStr; + +use litellm_testkit::{ClaudeCode, Codex, Error, Installer, Opencode, Target, Version}; +use rstest::rstest; +use serde_json::json; +use support::{FakeFetch, script_printing, sha256, tar_gz, zip_archive}; +use target_lexicon::Triple; + +fn target(triple: &str) -> Target { + Target::try_from(&Triple::from_str(triple).unwrap()).unwrap() +} + +fn linux() -> Target { + target("x86_64-unknown-linux-gnu") +} +fn version() -> Version { + Version::new(9, 8, 7) +} + +fn github_release(asset: &str, download_url: &str, digest: Option) -> Vec { + json!({ + "assets": [ + { "name": "unrelated.txt", "digest": "sha256:00", "browser_download_url": "https://example.test/unrelated" }, + { "name": asset, "digest": digest, "browser_download_url": download_url }, + ] + }) + .to_string() + .into_bytes() +} + +fn claude_routes(binary: &[u8], checksum: &str) -> Vec<(String, Vec)> { + let base = "https://downloads.claude.ai/claude-code-releases/9.8.7"; + let manifest = json!({ "platforms": { "linux-x64": { "checksum": checksum } } }); + vec![ + ( + format!("{base}/manifest.json"), + manifest.to_string().into_bytes(), + ), + (format!("{base}/linux-x64/claude"), binary.to_vec()), + ] +} + +fn codex_routes(archive: Vec, digest: Option) -> Vec<(String, Vec)> { + let release = github_release( + "codex-x86_64-unknown-linux-musl.tar.gz", + "https://example.test/codex.tar.gz", + digest, + ); + vec![ + ( + "https://api.github.com/repos/openai/codex/releases/tags/rust-v9.8.7".to_owned(), + release, + ), + ("https://example.test/codex.tar.gz".to_owned(), archive), + ] +} + +#[tokio::test] +async fn claude_bare_binary_is_installed_and_runnable() { + let binary = script_printing("9.8.7 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary))); + let cache = tempfile::tempdir().unwrap(); + + let installed = Installer::new(&fetch, cache.path(), linux()) + .install(&ClaudeCode, &version()) + .await + .unwrap(); + + assert_eq!(installed.binary, cache.path().join("claude/9.8.7/claude")); + assert_eq!(std::fs::read(&installed.binary).unwrap(), binary); +} + +#[tokio::test] +async fn codex_binary_is_extracted_from_the_tarball_under_its_own_name() { + let binary = script_printing("codex-cli 9.8.7"); + let archive = tar_gz("codex-x86_64-unknown-linux-musl", &binary); + let fetch = FakeFetch::new(codex_routes( + archive.clone(), + Some(format!("sha256:{}", sha256(&archive))), + )); + let cache = tempfile::tempdir().unwrap(); + + let installed = Installer::new(&fetch, cache.path(), linux()) + .install(&Codex, &version()) + .await + .unwrap(); + + assert_eq!(std::fs::read(&installed.binary).unwrap(), binary); + assert_eq!(installed.binary, cache.path().join("codex/9.8.7/codex")); +} + +#[tokio::test] +async fn opencode_binary_is_extracted_from_the_darwin_zip() { + let binary = script_printing("9.8.7"); + let archive = zip_archive("opencode", &binary); + let release = github_release( + "opencode-darwin-arm64.zip", + "https://example.test/opencode.zip", + Some(format!("sha256:{}", sha256(&archive))), + ); + let fetch = FakeFetch::new([ + ( + "https://api.github.com/repos/sst/opencode/releases/tags/v9.8.7".to_owned(), + release, + ), + ("https://example.test/opencode.zip".to_owned(), archive), + ]); + let cache = tempfile::tempdir().unwrap(); + + let installed = Installer::new(&fetch, cache.path(), target("aarch64-apple-darwin")) + .install(&Opencode, &version()) + .await + .unwrap(); + + assert_eq!(std::fs::read(&installed.binary).unwrap(), binary); +} + +#[tokio::test] +async fn tampered_download_is_rejected_and_nothing_is_left_behind() { + let binary = script_printing("9.8.7 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(b"what the vendor signed"))); + let cache = tempfile::tempdir().unwrap(); + + let result = Installer::new(&fetch, cache.path(), linux()) + .install(&ClaudeCode, &version()) + .await; + + assert!(matches!(result, Err(Error::ChecksumMismatch { .. }))); + assert!(!cache.path().join("claude/9.8.7").exists()); +} + +#[tokio::test] +async fn github_asset_without_a_digest_is_refused() { + let archive = tar_gz( + "codex-x86_64-unknown-linux-musl", + &script_printing("codex-cli 9.8.7"), + ); + let fetch = FakeFetch::new(codex_routes(archive, None)); + let cache = tempfile::tempdir().unwrap(); + + let result = Installer::new(&fetch, cache.path(), linux()) + .install(&Codex, &version()) + .await; + + assert!(matches!(result, Err(Error::MissingChecksum(_)))); +} + +#[tokio::test] +async fn binary_reporting_a_different_version_is_removed() { + let binary = script_printing("1.0.0 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary))); + let cache = tempfile::tempdir().unwrap(); + + let result = Installer::new(&fetch, cache.path(), linux()) + .install(&ClaudeCode, &version()) + .await; + + assert!(matches!(result, Err(Error::VersionMismatch { .. }))); + assert!(!cache.path().join("claude/9.8.7/claude").exists()); +} + +#[tokio::test] +async fn second_install_reuses_the_cached_binary_without_downloading() { + let binary = script_printing("9.8.7 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary))); + let cache = tempfile::tempdir().unwrap(); + let installer = Installer::new(&fetch, cache.path(), linux()); + + let first = installer.install(&ClaudeCode, &version()).await.unwrap(); + let calls_after_first = fetch.calls(); + let second = installer.install(&ClaudeCode, &version()).await.unwrap(); + + assert_eq!(first, second); + assert_eq!(fetch.calls(), calls_after_first); +} + +#[tokio::test] +async fn corrupted_cache_entry_is_replaced_by_a_fresh_download() { + let binary = script_printing("9.8.7 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary))); + let cache = tempfile::tempdir().unwrap(); + let installer = Installer::new(&fetch, cache.path(), linux()); + let installed = installer.install(&ClaudeCode, &version()).await.unwrap(); + std::fs::write(&installed.binary, script_printing("0.0.1")).unwrap(); + + installer.install(&ClaudeCode, &version()).await.unwrap(); + + assert_eq!(std::fs::read(&installed.binary).unwrap(), binary); +} + +#[rstest] +#[case("9.8.7-beta.1")] +#[case("9.8.7+build.5")] +#[tokio::test] +async fn pre_releases_never_reach_the_network_or_the_filesystem(#[case] version: &str) { + let fetch = FakeFetch::new([]); + let cache = tempfile::tempdir().unwrap(); + + let result = Installer::new(&fetch, cache.path(), linux()) + .install(&ClaudeCode, &Version::parse(version).unwrap()) + .await; + + assert!(matches!(result, Err(Error::InvalidVersion(_)))); + assert_eq!(fetch.calls(), 0); + assert_eq!(std::fs::read_dir(cache.path()).unwrap().count(), 0); +} + +#[tokio::test] +async fn musl_linux_picks_the_musl_claude_build() { + let binary = script_printing("9.8.7 (Claude Code)"); + let base = "https://downloads.claude.ai/claude-code-releases/9.8.7"; + let manifest = json!({ "platforms": { + "linux-x64": { "checksum": sha256(b"glibc build") }, + "linux-x64-musl": { "checksum": sha256(&binary) }, + } }); + let fetch = FakeFetch::new([ + ( + format!("{base}/manifest.json"), + manifest.to_string().into_bytes(), + ), + (format!("{base}/linux-x64-musl/claude"), binary.clone()), + ]); + let cache = tempfile::tempdir().unwrap(); + + let installed = Installer::new(&fetch, cache.path(), target("x86_64-unknown-linux-musl")) + .install(&ClaudeCode, &version()) + .await + .unwrap(); + + assert_eq!(std::fs::read(&installed.binary).unwrap(), binary); +} + +#[rstest] +#[case("x86_64-pc-windows-msvc")] +#[case("riscv64gc-unknown-linux-gnu")] +#[case("wasm32-unknown-unknown")] +fn targets_no_agent_ships_for_are_rejected(#[case] triple: &str) { + let result = Target::try_from(&Triple::from_str(triple).unwrap()); + + assert!(matches!(result, Err(Error::UnsupportedTarget(_)))); +} + +#[tokio::test] +async fn concurrent_installs_of_the_same_version_both_succeed() { + let binary = script_printing("9.8.7 (Claude Code)"); + let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary))); + let cache = tempfile::tempdir().unwrap(); + let installer = Installer::new(&fetch, cache.path(), linux()); + + let wanted = version(); + let installs = + futures_util::future::join_all((0..8).map(|_| installer.install(&ClaudeCode, &wanted))) + .await; + + assert!(installs.iter().all(Result::is_ok)); + assert_eq!( + std::fs::read(&installs[0].as_ref().unwrap().binary).unwrap(), + binary + ); +} diff --git a/litellm-rust/crates/testkit/tests/live.rs b/litellm-rust/crates/testkit/tests/live.rs new file mode 100644 index 00000000000..b805596879a --- /dev/null +++ b/litellm-rust/crates/testkit/tests/live.rs @@ -0,0 +1,133 @@ +//! Drives the real agents through a real gateway. Run with `cargo test -p litellm-testkit --test live -- --ignored` +//! after exporting `TESTKIT_GATEWAY_URL`, `TESTKIT_GATEWAY_KEY`, one `TESTKIT_MODEL_` per wire +//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT__VERSION` per agent +//! (`CLAUDE`, `CODEX`, `OPENCODE`). `TESTKIT_CACHE_DIR` and `GITHUB_TOKEN` are optional. + +use std::path::PathBuf; +use std::time::Duration; + +use litellm_testkit::{ + Agent, ClaudeCode, Codex, HttpFetch, Installer, Opencode, Outcome, Prompt, Session, Settings, + Target, Version, Wire, +}; +use rstest::rstest; + +const LIMIT: Duration = Duration::from_secs(180); + +fn required(name: &str) -> String { + std::env::var(name).unwrap_or_else(|_| panic!("{name} must be set to run the live tests")) +} + +fn model_var(wire: Wire) -> &'static str { + match wire { + Wire::Messages => "TESTKIT_MODEL_MESSAGES", + Wire::Responses => "TESTKIT_MODEL_RESPONSES", + Wire::ChatCompletions => "TESTKIT_MODEL_CHAT_COMPLETIONS", + } +} + +async fn drive( + agent: &impl Agent, + version_var: &str, + wire: Wire, + model: Option<&str>, + prompt: Prompt, +) -> Outcome { + let cache = std::env::var("TESTKIT_CACHE_DIR") + .map(PathBuf::from) + .unwrap_or_else(|_| std::env::temp_dir().join("litellm-testkit-cache")); + let installer = Installer::new(HttpFetch::from_env(), cache, Target::host().unwrap()); + let installed = installer + .install(agent, &Version::parse(&required(version_var)).unwrap()) + .await + .unwrap(); + let settings = Settings { + base_url: required("TESTKIT_GATEWAY_URL"), + api_key: required("TESTKIT_GATEWAY_KEY"), + model: model.map_or_else(|| required(model_var(wire)), str::to_owned), + wire, + }; + let home = tempfile::tempdir().unwrap(); + let session = Session::prepare(agent, &installed, settings, home.path()).unwrap(); + session.run(agent, &prompt, LIMIT).await.unwrap() +} + +fn text_prompt() -> Prompt { + Prompt { + text: "Reply with the single word: pong".to_owned(), + allow_tools: false, + } +} + +fn tool_prompt() -> Prompt { + Prompt { + text: "Run the shell command 'echo tool-ok' and reply with exactly its output.".to_owned(), + allow_tools: true, + } +} + +#[rstest] +#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)] +#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)] +#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)] +#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)] +#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)] +#[ignore = "needs a live gateway, see the module docs"] +#[tokio::test] +async fn plain_prompt_gets_an_answer_and_token_usage( + #[case] agent: &impl Agent, + #[case] version_var: &str, + #[case] wire: Wire, +) { + let outcome = drive(agent, version_var, wire, None, text_prompt()).await; + + assert!(outcome.succeeded(), "{outcome:?}"); + assert!(outcome.text.to_lowercase().contains("pong"), "{outcome:?}"); + assert!(outcome.usage.output_tokens > 0, "{outcome:?}"); +} + +#[rstest] +#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)] +#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)] +#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)] +#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)] +#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)] +#[ignore = "needs a live gateway, see the module docs"] +#[tokio::test] +async fn tool_use_is_reported_and_its_result_reaches_the_answer( + #[case] agent: &impl Agent, + #[case] version_var: &str, + #[case] wire: Wire, +) { + let outcome = drive(agent, version_var, wire, None, tool_prompt()).await; + + assert!(outcome.succeeded(), "{outcome:?}"); + assert!(!outcome.tool_calls.is_empty(), "{outcome:?}"); + assert!(outcome.text.contains("tool-ok"), "{outcome:?}"); +} + +#[rstest] +#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)] +#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)] +#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)] +#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)] +#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)] +#[ignore = "needs a live gateway, see the module docs"] +#[tokio::test] +async fn model_the_gateway_rejects_is_reported_as_an_error( + #[case] agent: &impl Agent, + #[case] version_var: &str, + #[case] wire: Wire, +) { + let outcome = drive( + agent, + version_var, + wire, + Some("testkit-no-such-model"), + text_prompt(), + ) + .await; + + assert!(!outcome.succeeded(), "{outcome:?}"); + assert!(!outcome.errors.is_empty(), "{outcome:?}"); +} diff --git a/litellm-rust/crates/testkit/tests/session.rs b/litellm-rust/crates/testkit/tests/session.rs new file mode 100644 index 00000000000..cd5e0dcbc71 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/session.rs @@ -0,0 +1,155 @@ +use std::os::unix::fs::PermissionsExt; +use std::path::{Path, PathBuf}; +use std::time::Duration; + +use litellm_testkit::{ + Configure, Drive, Error, Installed, LaunchSpec, Outcome, Prompt, Session, Settings, Version, + Wire, +}; + +struct Scripted; + +impl Configure for Scripted { + fn configure( + &self, + version: &Version, + _settings: &Settings, + home: &Path, + ) -> Result { + Ok(LaunchSpec { + env: [ + ("AGENT_HOME".to_owned(), home.to_string_lossy().into_owned()), + ("AGENT_SAW_VERSION".to_owned(), version.to_string()), + ] + .into(), + files: [( + PathBuf::from("conf/agent.toml"), + "configured = true\n".to_owned(), + )] + .into(), + }) + } +} + +impl Drive for Scripted { + fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec { + vec!["--prompt".to_owned(), prompt.text.clone()] + } + + fn parse(&self, _version: &Version, stdout: &str) -> Outcome { + Outcome { + text: stdout.to_owned(), + ..Outcome::default() + } + } +} + +fn settings() -> Settings { + Settings { + base_url: "http://gateway.test".to_owned(), + api_key: "sk-test".to_owned(), + model: "some-model".to_owned(), + wire: Wire::Messages, + } +} + +fn prompt(text: &str) -> Prompt { + Prompt { + text: text.to_owned(), + allow_tools: false, + } +} + +fn session(script: &str) -> (Session, tempfile::TempDir) { + let dir = tempfile::tempdir().unwrap(); + let binary = dir.path().join("agent"); + std::fs::write(&binary, format!("#!/bin/sh\n{script}\n")).unwrap(); + std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o755)).unwrap(); + let home = dir.path().join("home"); + std::fs::create_dir(&home).unwrap(); + let installed = Installed { + version: Version::new(4, 5, 6), + binary, + }; + ( + Session::prepare(&Scripted, &installed, settings(), home).unwrap(), + dir, + ) +} + +const LIMIT: Duration = Duration::from_secs(20); + +#[tokio::test] +async fn prepare_writes_the_config_files_under_home() { + let (_session, dir) = session("true"); + + let written = std::fs::read_to_string(dir.path().join("home/conf/agent.toml")).unwrap(); + + assert_eq!(written, "configured = true\n"); +} + +#[tokio::test] +async fn configure_and_drive_are_given_the_installed_version() { + let (session, _dir) = session("echo \"$AGENT_SAW_VERSION\""); + + let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap(); + + assert_eq!(outcome.text.trim(), "4.5.6"); +} + +#[tokio::test] +async fn agent_runs_in_home_with_only_its_own_environment() { + let (session, dir) = session("pwd -P; env"); + + let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap(); + + let home = dir.path().join("home").canonicalize().unwrap(); + assert_eq!(outcome.text.lines().next().unwrap(), home.to_string_lossy()); + assert!(outcome.text.contains("AGENT_HOME=")); + assert!( + !outcome.text.contains("CARGO_"), + "test runner environment leaked into the agent" + ); +} + +#[tokio::test] +async fn prompt_reaches_the_agent_as_one_untouched_argument() { + let (session, _dir) = session("printf '%s|' \"$@\""); + let text = "two spaces; $(echo injected) 'quoted'"; + + let outcome = session.run(&Scripted, &prompt(text), LIMIT).await.unwrap(); + + assert_eq!(outcome.text, format!("--prompt|{text}|")); +} + +#[tokio::test] +async fn clean_exit_is_a_success() { + let (session, _dir) = session("echo done"); + + let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap(); + + assert_eq!(outcome.exit_code, Some(0)); + assert!(outcome.succeeded()); +} + +#[tokio::test] +async fn failing_exit_without_a_parsed_error_reports_stderr() { + let (session, _dir) = session("echo boom >&2; exit 3"); + + let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap(); + + assert_eq!(outcome.exit_code, Some(3)); + assert!(!outcome.succeeded()); + assert_eq!(outcome.errors, ["boom\n"]); +} + +#[tokio::test] +async fn agent_that_outlives_the_limit_is_stopped() { + let (session, _dir) = session("sleep 30"); + + let result = session + .run(&Scripted, &prompt("hi"), Duration::from_millis(200)) + .await; + + assert!(matches!(result, Err(Error::Timeout(_)))); +} diff --git a/litellm-rust/crates/testkit/tests/support/mod.rs b/litellm-rust/crates/testkit/tests/support/mod.rs new file mode 100644 index 00000000000..f4a13759941 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/support/mod.rs @@ -0,0 +1,70 @@ +use std::collections::HashMap; +use std::io::Write; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use litellm_testkit::{Error, Fetch}; +use sha2::{Digest, Sha256}; + +pub struct FakeFetch { + routes: HashMap>, + calls: AtomicUsize, +} + +impl FakeFetch { + pub fn new(routes: impl IntoIterator)>) -> Self { + Self { + routes: routes.into_iter().collect(), + calls: AtomicUsize::new(0), + } + } + + pub fn calls(&self) -> usize { + self.calls.load(Ordering::SeqCst) + } +} + +impl Fetch for FakeFetch { + async fn get(&self, url: &str) -> Result, Error> { + self.calls.fetch_add(1, Ordering::SeqCst); + self.routes.get(url).cloned().ok_or_else(|| Error::Status { + url: url.to_owned(), + status: 404, + }) + } +} + +impl Fetch for &FakeFetch { + async fn get(&self, url: &str) -> Result, Error> { + (*self).get(url).await + } +} + +pub fn sha256(bytes: &[u8]) -> String { + format!("{:x}", Sha256::digest(bytes)) +} + +pub fn script_printing(output: &str) -> Vec { + format!("#!/bin/sh\necho '{output}'\n").into_bytes() +} + +pub fn tar_gz(member: &str, contents: &[u8]) -> Vec { + let mut builder = tar::Builder::new(Vec::new()); + let mut header = tar::Header::new_gnu(); + header.set_size(contents.len() as u64); + header.set_mode(0o755); + header.set_cksum(); + builder.append_data(&mut header, member, contents).unwrap(); + let tarball = builder.into_inner().unwrap(); + let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default()); + encoder.write_all(&tarball).unwrap(); + encoder.finish().unwrap() +} + +pub fn zip_archive(member: &str, contents: &[u8]) -> Vec { + let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new())); + writer + .start_file(member, zip::write::SimpleFileOptions::default()) + .unwrap(); + writer.write_all(contents).unwrap(); + writer.finish().unwrap().into_inner() +} diff --git a/litellm/__init__.py b/litellm/__init__.py index 8b1b5a5d008..e334fbe8ca8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -657,6 +657,7 @@ azure_anthropic_models: Set = set() azure_text_models: Set = set() anyscale_models: Set = set() cerebras_models: Set = set() +nadir_models: Set = set() # mutable-ok: provider registry, filled from model_cost at import like every sibling provider galadriel_models: Set = set() nvidia_nim_models: Set = set() nvidia_riva_models: Set = set() @@ -893,6 +894,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None: anyscale_models.add(key) elif value.get("litellm_provider") == "cerebras": cerebras_models.add(key) + elif value.get("litellm_provider") == "nadir": + nadir_models.add(key) elif value.get("litellm_provider") == "galadriel": galadriel_models.add(key) elif value.get("litellm_provider") == "nvidia_nim": @@ -1083,6 +1086,7 @@ model_list = list( | azure_anthropic_models | anyscale_models | cerebras_models + | nadir_models | galadriel_models | nvidia_nim_models | nvidia_riva_models @@ -1191,6 +1195,7 @@ def _build_models_by_provider() -> dict: "azure_text": azure_text_models, "anyscale": anyscale_models, "cerebras": cerebras_models, + "nadir": nadir_models, "galadriel": galadriel_models, "nvidia_nim": nvidia_nim_models, "nvidia_riva": nvidia_riva_models, @@ -1994,6 +1999,7 @@ if TYPE_CHECKING: FeatherlessAIConfig as FeatherlessAIConfig, ) from .llms.cerebras.chat import CerebrasConfig as CerebrasConfig + from .llms.nadir.chat.transformation import NadirConfig as NadirConfig from .llms.baseten.chat import BasetenConfig as BasetenConfig from .llms.sambanova.chat import SambanovaConfig as SambanovaConfig from .llms.sambanova.embedding.transformation import ( diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index d3236a04ae0..42513321391 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -264,6 +264,7 @@ LLM_CONFIG_NAMES: Final = ( "NvidiaNimEmbeddingConfig", "FeatherlessAIConfig", "CerebrasConfig", + "NadirConfig", "BasetenConfig", "SambanovaConfig", "SambaNovaEmbeddingConfig", @@ -1061,6 +1062,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { "FeatherlessAIConfig", ), "CerebrasConfig": (".llms.cerebras.chat", "CerebrasConfig"), + "NadirConfig": (".llms.nadir.chat.transformation", "NadirConfig"), "BasetenConfig": (".llms.baseten.chat", "BasetenConfig"), "SambanovaConfig": (".llms.sambanova.chat", "SambanovaConfig"), "SambaNovaEmbeddingConfig": ( diff --git a/litellm/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py index eee7e2ea289..6e4635dd83a 100644 --- a/litellm/caching/evicted_client_closer.py +++ b/litellm/caching/evicted_client_closer.py @@ -136,6 +136,8 @@ def _has_connection_in_flight(client: object) -> bool: window as the only guard, exactly as it was before this check existed. """ try: + if getattr(getattr(client, "connection_pool", None), "_in_use_connections", None): + return True transport: Final = _transport_of(client) pooled_busy: Final = _pool_has_busy_connection(transport) if pooled_busy is not None: diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7cec84e0ebb..0b56c28f9b1 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -21,6 +21,7 @@ from collections.abc import Awaitable, Callable, Iterator, Sequence from contextvars import ContextVar from dataclasses import dataclass from datetime import timedelta +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast from pydantic import TypeAdapter @@ -290,6 +291,29 @@ def _opaque_kwarg_key(value: object) -> str: return f"{type(value).__name__}-{id(value)}" +_CLUSTER_ONLY_CONNECTION_KWARGS: Final[frozenset[str]] = frozenset({"response_callbacks"}) + + +def _cluster_node_pubsub_client( # pyright: ignore[reportUnknownParameterType] # redis generics + cluster: async_redis_cluster_client, # pyright: ignore[reportUnknownParameterType] # redis generics +) -> async_redis_client: + """Plain async client on one cluster node; classic PUBLISH/SUBSCRIBE is broadcast cluster-wide.""" + from redis.asyncio import ConnectionPool, Redis + + node: Final = cluster.get_default_node() or next(iter(cluster.nodes_manager.startup_nodes.values()), None) + if node is None: # pyright: ignore[reportUnnecessaryComparison] # get_default_node is None before cluster init + raise ValueError("cannot derive a pub/sub client: redis cluster has no default node and no startup nodes") + node_kwargs: Final = MappingProxyType( + { + key: value # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs + for key, value in cluster.connection_kwargs.items() # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs + if key not in _CLUSTER_ONLY_CONNECTION_KWARGS + } + ) + pool: Final = ConnectionPool(host=node.host, port=node.port, **node_kwargs) # pyright: ignore[reportCallIssue, reportArgumentType] # cluster kwargs validated by redis-py at runtime + return Redis.from_pool(pool) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics + + @functools.lru_cache(maxsize=1) def _redis_health_error_types() -> tuple[type, ...]: """Exception types that mean the Redis backend itself is unhealthy. @@ -738,6 +762,28 @@ class RedisCache(BaseCache): self.redis_async_client = redis_async_client return redis_async_client + def init_pubsub_client(self) -> async_redis_client: # pyright: ignore[reportUnknownParameterType] # redis generics + from redis.asyncio import RedisCluster + + from litellm import in_memory_llm_clients_cache + + client: Final = self.init_async_client() # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics + if not isinstance(client, RedisCluster): + return client # pyright: ignore[reportUnknownVariableType] # redis generics + cache_key: Final = f"{self._get_async_client_cache_key()}-pubsub" + cached_client: Final = in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped in-memory client cache + key=cache_key + ) + if cached_client is not None: + return cast( # cast-ok: per-loop pub/sub client stored by this method # pyright: ignore[reportUnknownVariableType] # redis generics + async_redis_client, cached_client + ) + pubsub_client: Final = _cluster_node_pubsub_client( # pyright: ignore[reportUnknownVariableType] # redis generics + cluster=client + ) + in_memory_llm_clients_cache.set_cache(key=cache_key, value=pubsub_client, litellm_owned_client=True) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped in-memory client cache + return pubsub_client # pyright: ignore[reportUnknownVariableType] # redis generics + def _async_commands(self) -> _AsyncRedisCommands: return self.init_async_client() @@ -1785,7 +1831,21 @@ class RedisCache(BaseCache): self.redis_client.flushall() async def disconnect(self): - await self.async_redis_conn_pool.disconnect(inuse_connections=True) + from litellm import in_memory_llm_clients_cache + + if self.async_redis_conn_pool is not None: + await self.async_redis_conn_pool.disconnect(inuse_connections=True) + cached_pubsub_client: Final = cast( # cast-ok: only this module stores clients under this key # pyright: ignore[reportUnknownVariableType] # redis generics + async_redis_client | None, + in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType] # untyped in-memory client cache + key=f"{self._get_async_client_cache_key()}-pubsub" + ), + ) + if cached_pubsub_client is not None: + try: + await cached_pubsub_client.aclose() # pyright: ignore[reportUnknownMemberType, reportAttributeAccessIssue] # redis stubs leave aclose unknown + except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection + verbose_logger.debug("Error closing cached pub/sub Redis client: %s", e) try: self.redis_client.close() except Exception as e: diff --git a/litellm/constants.py b/litellm/constants.py index a86be55d654..e7ba1f6b07f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60")) +# mcp_tool_permissions entry that grants every current and future tool on a server +MCP_ALL_TOOLS_WILDCARD: Final = "*" + # Default npm cache directory for STDIO MCP servers. # npm/npx needs a writable cache dir; in containers the default (~/.npm) # may not exist or be read-only. /tmp is always writable. @@ -327,6 +330,7 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float( WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123 DEEPGRAM_DEFAULT_API_BASE: Final = "https://api.deepgram.com/v1" +NADIR_DEFAULT_API_BASE: Final = "https://api.getnadir.com/v1" DEEPGRAM_LISTEN_DEFAULT_MODEL: Final = "nova-3" BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update" @@ -594,6 +598,7 @@ FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO: Final = 0.5 #### Logging callback constants #### REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM" MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50)) +LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS: Final = 10_000 # Backpressure + lifetime bounds for the /v1/messages streaming relay (see # BaseAnthropicMessagesStreamingIterator.async_sse_wrapper). The relay queue is # bounded so a slow client throttles the upstream pump instead of letting it @@ -708,6 +713,7 @@ LITELLM_CHAT_PROVIDERS: Final = [ "gigachat", "nvidia_nim", "cerebras", + "nadir", "baseten", "ai21_chat", "volcengine", @@ -901,6 +907,7 @@ openai_compatible_endpoints: Final[list] = [ "codestral.mistral.ai/v1/fim/completions", "api.groq.com/openai/v1", "https://integrate.api.nvidia.com/v1", + NADIR_DEFAULT_API_BASE, "api.deepseek.com/v1", "api.together.ai/v1", "api.together.xyz/v1", diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 6cc0d9444cd..a279b9f0903 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -790,6 +790,16 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None return value if isinstance(value, str) and value else None +_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"}) + + +def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool: + return any( + value is not None and (field in _NON_TOKEN_RATE_FIELDS or ("cost_per" in field and "token" in field)) + for field, value in entry.items() + ) + + def _select_model_name_for_cost_calc( model: str | None, completion_response: object | None, @@ -828,12 +838,7 @@ def _select_model_name_for_cost_calc( if custom_pricing is True: if router_model_id is not None and router_model_id in litellm.model_cost: entry: Final = litellm.model_cost[router_model_id] - if ( - entry.get("input_cost_per_token") is not None - or entry.get("input_cost_per_second") is not None - or entry.get("input_cost_per_query") is not None - or entry.get("tiered_pricing") is not None - ): + if _cost_map_entry_prices_anything(entry): return_model = router_model_id else: return_model = model @@ -1699,6 +1704,8 @@ def completion_cost( litellm_model_name=model, data_residency=data_residency, litellm_logging_obj=litellm_logging_obj, + custom_pricing_model=selected_model if custom_pricing else None, + base_pricing_model=(selected_model if base_model is not None and not custom_pricing else None), ) elif call_type == _MCP_CALL_TYPE: from litellm.proxy._experimental.mcp_server.cost_calculator import ( @@ -2870,14 +2877,20 @@ def _candidate_realtime_token_costs( def _cost_map_entry_declares_pricing(model_name: str, custom_llm_provider: str) -> bool: + """Whether the entry behind ``model_name`` sets any rate of its own, even a zero one. + + The name is resolved the way ``get_model_info`` resolves it before the raw entry is read, + because a deployment-scoped name arrives here already carrying its provider prefix. Two raw + lookups cannot strip that prefix, so a zero-rated override read as declaring nothing, and a + session that should bill nothing fell through to the public rates instead. + """ + resolved: Final = _get_model_info_or_none(model_name, custom_llm_provider) entries: Final = ( + litellm.model_cost.get(resolved.get("key")) if resolved is not None else None, litellm.model_cost.get(model_name), litellm.model_cost.get(f"{custom_llm_provider}/{model_name}"), ) - return any( - entry is not None and any("cost_per" in field and value is not None for field, value in entry.items()) - for entry in entries - ) + return any(entry is not None and _cost_map_entry_prices_anything(entry) for entry in entries) def _first_priced_realtime_token_costs( @@ -2917,6 +2930,8 @@ def handle_realtime_stream_cost_calculation( litellm_model_name: str, data_residency: str | None = None, litellm_logging_obj: LitellmLoggingObject | None = None, + custom_pricing_model: str | None = None, + base_pricing_model: str | None = None, ) -> float: """ Handles the cost calculation for realtime stream responses. @@ -2925,9 +2940,13 @@ def handle_realtime_stream_cost_calculation( Args: results: A list of OpenAIRealtimeStreamBaseObject objects + custom_pricing_model: deployment-scoped pricing key from the deployment's + custom rates, tried ahead of the session-reported model + base_pricing_model: the deployment's resolved base_model, tried ahead of the + session-reported model but after custom rates """ received_model = None - potential_model_names: Final = [] + potential_model_names: Final = [custom_pricing_model, base_pricing_model] for result in results: if result["type"] == "session.created": received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None) @@ -2945,6 +2964,7 @@ def handle_realtime_stream_cost_calculation( results=results, custom_llm_provider=custom_llm_provider, litellm_model_name=litellm_model_name, + custom_pricing_model=custom_pricing_model, ) if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 @@ -2968,6 +2988,7 @@ def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, litellm_model_name: str, + custom_pricing_model: str | None = None, ) -> float: """ Cost for realtime transcription sessions (e.g. gpt-realtime-whisper). @@ -2985,15 +3006,15 @@ def handle_realtime_transcription_cost_calculation( return 0.0 model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name - try: - model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider) - except Exception: - model_info = None + model_info: Final = _get_model_info_or_none(model_name, custom_llm_provider) + override_info: Final = ( + _get_model_info_or_none(custom_pricing_model, custom_llm_provider) if custom_pricing_model is not None else None + ) total_cost = 0.0 for event in completed_events: usage = event.get("usage") or {} - total_cost += _transcription_usage_cost(usage, model_info) + total_cost += _transcription_usage_cost(usage, model_info, override_info) return total_cost @@ -3018,23 +3039,57 @@ def _get_transcription_model_name_from_results( return None -def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float: - if model_info is None: +def _get_model_info_or_none(model: str, custom_llm_provider: str) -> ModelInfo | None: + try: + return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + except Exception: + return None + + +def _declared_transcription_rate(info: ModelInfo | None, keys: tuple[str, ...]) -> float | None: + """First of ``keys`` this entry prices, read off the raw ``litellm.model_cost`` entry + because ``get_model_info`` synthesizes zero token rates for entries that omit them.""" + if info is None: + return None + declared: Final = litellm.model_cost.get(info.get("key")) + if declared is None: + return None + return next( + (float(value) for key in keys if declared.get(key) is not None and (value := info.get(key)) is not None), + None, + ) + + +def _transcription_rate(keys: tuple[str, ...], override: ModelInfo | None, base: ModelInfo | None) -> float: + rates: Final = (_declared_transcription_rate(info, keys) for info in (override, base)) + return next((rate for rate in rates if rate is not None), 0.0) + + +def _transcription_usage_cost( + usage: dict, + model_info: ModelInfo | None, + override_info: ModelInfo | None = None, +) -> float: + if model_info is None and override_info is None: return 0.0 + usage_type: Final = usage.get("type") if usage_type == "duration": seconds: Final = usage.get("seconds") or 0.0 - per_second: Final = model_info.get("input_cost_per_second") or 0.0 - return float(seconds) * float(per_second) + return float(seconds) * _transcription_rate(("input_cost_per_second",), override_info, model_info) if usage_type == "tokens": input_token_details: Final = usage.get("input_token_details") or {} audio_tokens: Final = input_token_details.get("audio_tokens") or 0 text_tokens: Final = input_token_details.get("text_tokens") or 0 output_tokens: Final = usage.get("output_tokens") or 0 - audio_cost: Final = float(audio_tokens) * float( - model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0 + audio_cost: Final = float(audio_tokens) * _transcription_rate( + ("input_cost_per_audio_token", "input_cost_per_token"), override_info, model_info + ) + text_cost: Final = float(text_tokens) * _transcription_rate( + ("input_cost_per_token",), override_info, model_info + ) + output_cost: Final = float(output_tokens) * _transcription_rate( + ("output_cost_per_token",), override_info, model_info ) - text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0) - output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0) return audio_cost + text_cost + output_cost return 0.0 diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1206f9abcbd..01670be74c8 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -36,6 +36,7 @@ from mcp.types import ( REQUEST_TIMEOUT, GetPromptRequestParams, GetPromptResult, + InputRequiredResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult, @@ -44,7 +45,6 @@ from mcp.types import ( Prompt, ResourceTemplate, ServerNotification, - TextContent, ) from mcp.types import CallToolRequestParams as MCPCallToolRequestParams from mcp.types import CallToolResult as MCPCallToolResult @@ -61,6 +61,7 @@ from litellm.constants import ( from litellm.experimental_mcp_client.tools import list_tools_with_pagination from litellm.llms.custom_httpx.http_handler import get_ssl_configuration from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response +from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result from litellm.types.llms.custom_http import VerifyTypes from litellm.types.mcp import ( MCPAuth, @@ -828,17 +829,15 @@ class MCPClient: @staticmethod def error_tool_result(exc: Exception) -> MCPCallToolResult: """The error result ``call_tool`` returns when it swallows a failure (no re-execution).""" - return MCPCallToolResult( - content=[TextContent(type="text", text=f"{type(exc).__name__}: {exc}")], - is_error=True, - ) + return error_text_result(exc) async def call_tool( self, call_tool_request_params: MCPCallToolRequestParams, host_progress_callback: Callable | None = None, raise_on_error: bool = False, - ) -> MCPCallToolResult: + allow_input_required: bool = False, + ) -> MCPCallToolResult | InputRequiredResult: """ Call an MCP Tool. @@ -847,6 +846,9 @@ class MCPClient: ``isError=True`` result. The token-exchange (OBO) tool-call path uses this to detect an upstream 401 so it can re-mint the exchanged token and retry once; every other caller keeps the default and gets graceful ``isError`` degradation. + allow_input_required: When True, a 2026-07-28 upstream may answer with an interim + ``InputRequiredResult`` and it is returned as is. The SDK rejects it otherwise, so + callers only opt in when the downstream side can carry it. """ verbose_logger.info("MCP client calling tool '%s'", call_tool_request_params.name) @@ -869,6 +871,7 @@ class MCPClient: name=call_tool_request_params.name, arguments=call_tool_request_params.arguments, progress_callback=on_progress, + allow_input_required=allow_input_required, ) try: diff --git a/litellm/integrations/SlackAlerting/utils.py b/litellm/integrations/SlackAlerting/utils.py index 297d069a868..eb3a7f80f72 100644 --- a/litellm/integrations/SlackAlerting/utils.py +++ b/litellm/integrations/SlackAlerting/utils.py @@ -3,9 +3,11 @@ Utils used for slack alerting """ import asyncio +from collections.abc import Callable from typing import TYPE_CHECKING, Any, Final import litellm +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import AlertType from litellm.secret_managers.main import get_secret @@ -66,25 +68,27 @@ async def _add_langfuse_trace_id_to_alert( -> trace_id -> litellm_call_id """ - if "langfuse" not in litellm.logging_callback_manager._get_all_callbacks(): + from litellm.integrations.langfuse.langfuse import LangFuseLogger, resolve_langfuse_host + + callbacks: Final[list[CustomLogger | Callable[..., object] | str]] = ( + litellm.logging_callback_manager._get_all_callbacks() + ) + if not any(callback == "langfuse" or isinstance(callback, LangFuseLogger) for callback in callbacks): return None - ######################################################### - # Only run if langfuse is added as a callback - ######################################################### - if request_data is not None and request_data.get("litellm_logging_obj", None) is not None: - trace_id: str | None = None - litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"] + if request_data is None or request_data.get("litellm_logging_obj", None) is None: + return None - for _ in range(3): - trace_id = litellm_logging_obj._get_trace_id(service_name="langfuse") - if trace_id is not None: - break - await asyncio.sleep(3) # wait 3s before retrying for trace id - ######################################################### - langfuse_object: Final = litellm_logging_obj._get_callback_object(service_name="langfuse") - if langfuse_object is not None: - base_url: Final = langfuse_object.Langfuse.base_url - return f"{base_url}/trace/{trace_id}" + litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"] + instance_host: Final = next( + (callback.langfuse_host for callback in callbacks if isinstance(callback, LangFuseLogger)), None + ) + host: Final = resolve_langfuse_host( + litellm_logging_obj.standard_callback_dynamic_params.get("langfuse_host") or instance_host + ) + for _ in range(3): + if (trace_id := litellm_logging_obj._get_trace_id(service_name="langfuse")) is not None: + return f"{host}/trace/{trace_id}" + await asyncio.sleep(3) # wait 3s before retrying for trace id return None diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 96d711337fb..9b860840e69 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -1,14 +1,14 @@ #### What this does #### # On success, logs events to Langfuse -import inspect import os import re import traceback from collections.abc import Callable, Iterable, Mapping, Sequence from datetime import datetime from functools import lru_cache +from importlib.metadata import PackageNotFoundError, version from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable from packaging.version import Version @@ -45,13 +45,13 @@ from litellm.types.utils import ( ) if TYPE_CHECKING: - from langfuse.client import Langfuse, StatefulTraceClient - + from litellm.integrations.langfuse.langfuse_sdk import LangfuseApiClient, LangfuseObservation, LangfuseTracing from litellm.litellm_core_utils.litellm_logging import DynamicLoggingCache else: DynamicLoggingCache = Any - StatefulTraceClient = Any - Langfuse = Any + LangfuseApiClient = Any + LangfuseObservation = Any + LangfuseTracing = Any _DENIED_STEERING_KEYS: Final = frozenset({"headers", "endpoint", "caching_groups", "previous_models"}) @@ -142,6 +142,20 @@ def _logging_id(start_time: datetime | None, response_obj: object) -> str | None return litellm.utils.get_logging_id(start_time, response_obj) +@runtime_checkable +class _ResponseWithId(Protocol): + """Response payloads (ModelResponse and friends, or a plain dict) expose their provider id via ``get``.""" + + def get(self, key: Literal["id"], default: None = None, /) -> object: ... + + +def _lookup_ids(litellm_call_id: str | None, response_obj: object) -> Mapping[str, str]: + """v2 carried the response id inside the generation id; v4 hashes ids to 16 hex chars, so they ride in metadata.""" + response_id: Final[object] = response_obj.get("id") if isinstance(response_obj, _ResponseWithId) else None + ids: Final[tuple[tuple[str, object], ...]] = (("litellm_call_id", litellm_call_id), ("response_id", response_id)) + return MappingProxyType({key: str(value) for key, value in ids if value is not None}) + + def _as_steering_flag(value: object) -> bool: """A string ``str_to_bool`` does not recognise falls back to its truthiness.""" if isinstance(value, str): @@ -158,6 +172,68 @@ def _as_steering_key_sequence(value: object) -> tuple[str, ...]: return () +MINIMUM_LANGFUSE_VERSION: Final = "4.7" +UNSUPPORTED_LANGFUSE_VERSION: Final = "5" +PROMPT_CACHE_TTL_ENV: Final = "LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS" + + +def installed_langfuse_version() -> str: + """Only ``importlib.metadata`` reads correctly on every major. + + ``langfuse.version`` was removed in v4, ``langfuse.__version__`` does not + exist in v3, and in v2 it reports a different value from the distribution + that is actually installed. + """ + return version("langfuse") + + +def raise_if_unsupported_langfuse_version(installed_version: str) -> None: + """Fail at logger construction rather than dropping every event at request time. + + v4 moved the callback onto OpenTelemetry, so on an older SDK the import of + `LangfuseOtelSpanAttributes` raises inside the per-request handler and the + broad except there turns it into silent total data loss. + """ + installed: Final = Version(installed_version) + # compare majors, not versions: "5.0.0rc1" sorts below "5" but is just as unsupported + if Version(MINIMUM_LANGFUSE_VERSION) <= installed and installed.major < Version(UNSUPPORTED_LANGFUSE_VERSION).major: + return + raise ImportError( + f"\033[91mlitellm requires langfuse>={MINIMUM_LANGFUSE_VERSION},<{UNSUPPORTED_LANGFUSE_VERSION} for the " + f"'langfuse' callback, but {installed_version} is installed. Run " + f"'pip install \"langfuse>={MINIMUM_LANGFUSE_VERSION},<{UNSUPPORTED_LANGFUSE_VERSION}\"' to upgrade, or use " + f"the 'langfuse_otel' callback, which does not depend on the langfuse SDK\033[0m" + ) + + +def whole_number(raw: str) -> int | None: + try: + return int(raw) + except ValueError: + return None + + +def raise_if_unusable_prompt_cache_ttl() -> None: + """The v4 SDK runs ``int()`` on this variable while it is being imported, so a value that is not a whole + number has to be named here, before that import fails with a bare ``ValueError`` on every request.""" + raw: Final = os.environ.get(PROMPT_CACHE_TTL_ENV) + if raw is None or whole_number(raw) is not None: + return + raise ValueError(f"\033[91m{PROMPT_CACHE_TTL_ENV}={raw!r} must be a whole number of seconds\033[0m") + + +def _optional_str(value: object) -> str | None: + """v4 sets attribute values raw; a non-string version would be dropped by the server.""" + return str(value) if value is not None else None + + +def _trace_public_flag(value: object) -> bool | None: + """``trace_public`` reaches here as a bool from metadata or a string from a ``langfuse_*`` header.""" + if value is None: + return None + return _as_steering_flag(value) + + def resolve_langfuse_credentials( langfuse_public_key=None, langfuse_secret=None, @@ -172,9 +248,29 @@ def resolve_langfuse_credentials( secret_key = langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY") public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") - resolved_host: Final = langfuse_host or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com") + return public_key, secret_key, resolve_langfuse_host(langfuse_host) - return public_key, secret_key, resolved_host + +def resolve_langfuse_host(langfuse_host: object = None) -> str: + """The Langfuse base URL for ``langfuse_host`` with the env fallbacks, always carrying a scheme.""" + resolved: Final = str( + langfuse_host or os.getenv("LANGFUSE_HOST") or os.getenv("LANGFUSE_BASE_URL") or "https://cloud.langfuse.com" + ) + return resolved if resolved.startswith(("http://", "https://")) else f"http://{resolved}" + + +def warn_if_upstream_langfuse_configured() -> None: + if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is None: + return + verbose_logger.warning( + "UPSTREAM_LANGFUSE_* is no longer supported: the langfuse callback moved to SDK v4, " + "which has no second ingestion client. The values are ignored." + ) + + +def parse_langfuse_debug(raw_value: str | None) -> bool: + """Parse the LANGFUSE_DEBUG value into the boolean flag the langfuse client expects.""" + return raw_value is not None and raw_value.strip().lower() in ("true", "1") @lru_cache(maxsize=8) @@ -199,29 +295,29 @@ class LangFuseLogger: allow_env_credentials: bool = True, ): try: - import langfuse - from langfuse import Langfuse - except Exception as e: + self.langfuse_sdk_version: str = installed_langfuse_version() + except PackageNotFoundError as e: raise Exception( - f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m" - ) + f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\033[0m" + ) from e + raise_if_unsupported_langfuse_version(self.langfuse_sdk_version) + raise_if_unusable_prompt_cache_ttl() + from litellm.integrations.langfuse.langfuse_sdk import configured_release + self.public_key, self.secret_key, self.langfuse_host = resolve_langfuse_credentials( langfuse_public_key=langfuse_public_key, langfuse_secret=langfuse_secret, langfuse_host=langfuse_host, allow_env_credentials=allow_env_credentials, ) - if not (self.langfuse_host.startswith("http://") or self.langfuse_host.startswith("https://")): - # add http:// if unset, assume communicating over private network - e.g. render - self.langfuse_host = "http://" + self.langfuse_host _env_override: Final = str(langfuse_environment).strip() if langfuse_environment is not None else None if _env_override: validate_langfuse_environment_value(_env_override) self.langfuse_environment: str | None = _env_override else: self.langfuse_environment = self.resolve_deployment_environment() - self.langfuse_release = os.getenv("LANGFUSE_RELEASE") - self.langfuse_debug = os.getenv("LANGFUSE_DEBUG") + self.langfuse_release = configured_release() + self.langfuse_debug = parse_langfuse_debug(os.getenv("LANGFUSE_DEBUG")) self.langfuse_flush_interval = LangFuseLogger._get_langfuse_flush_interval(flush_interval) if should_use_langfuse_mock(): @@ -232,22 +328,9 @@ class LangFuseLogger: self.langfuse_client = self._http_handler.client self.is_mock_mode = False - parameters: Final = { - "public_key": self.public_key, - "secret_key": self.secret_key, - "host": self.langfuse_host, - "release": self.langfuse_release, - "debug": self.langfuse_debug, - "flush_interval": self.langfuse_flush_interval, # flush interval in seconds - "httpx_client": self.langfuse_client, - } - self.langfuse_sdk_version: str = langfuse.version.__version__ - - if "environment" in inspect.signature(Langfuse.__init__).parameters: - parameters["environment"] = self.langfuse_environment - if Version(self.langfuse_sdk_version) >= Version("2.6.0"): - parameters["sdk_integration"] = "litellm" - self.Langfuse: Langfuse = self.safe_init_langfuse_client(parameters) + self.api_client: LangfuseApiClient + self.tracing: LangfuseTracing + self.api_client, self.tracing = self.safe_init_langfuse_client() # set the current langfuse project id in the environ # this is used by Alerting to link to the correct project @@ -256,49 +339,62 @@ class LangFuseLogger: verbose_logger.debug("Langfuse Mock: Using mock project ID") else: try: - project_id = self.Langfuse.client.projects.get().data[0].id - os.environ["LANGFUSE_PROJECT_ID"] = project_id + project_id: Final = self.api_client.project_id() + if project_id is not None: + os.environ["LANGFUSE_PROJECT_ID"] = project_id except Exception: - project_id = None + verbose_logger.debug("Langfuse project id unavailable, alerting links will omit it") - if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None: - upstream_langfuse_debug_env: Final = os.getenv("UPSTREAM_LANGFUSE_DEBUG") - upstream_langfuse_debug: Final = ( - str_to_bool(upstream_langfuse_debug_env) if upstream_langfuse_debug_env is not None else None - ) - self.upstream_langfuse_secret_key = os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") - self.upstream_langfuse_public_key = os.getenv("UPSTREAM_LANGFUSE_PUBLIC_KEY") - self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST") - self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE") - self.upstream_langfuse_debug = upstream_langfuse_debug_env - self.upstream_langfuse = Langfuse( - public_key=self.upstream_langfuse_public_key, - secret_key=self.upstream_langfuse_secret_key, - host=self.upstream_langfuse_host, - release=self.upstream_langfuse_release, - debug=(upstream_langfuse_debug if upstream_langfuse_debug is not None else False), - ) - else: - self.upstream_langfuse = None + warn_if_upstream_langfuse_configured() - def safe_init_langfuse_client(self, parameters: dict) -> Langfuse: + def safe_init_langfuse_client(self) -> "tuple[LangfuseApiClient, LangfuseTracing]": + """Build the REST client and export channel while the process is under its logger budget. + + The budget dates from the SDK client, which started a consumer thread per instance and once + pinned a CPU at 100% when many were built; it still bounds the number of per-key loggers. """ - Safely init a langfuse client if the number of initialized clients is less than the max - - Note: - - Langfuse initializes 1 thread everytime a client is initialized. - - We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times. - """ - from langfuse import Langfuse - if litellm.initialized_langfuse_clients >= MAX_LANGFUSE_INITIALIZED_CLIENTS: raise Exception( f"Max langfuse clients reached: {litellm.initialized_langfuse_clients} is greater than {MAX_LANGFUSE_INITIALIZED_CLIENTS}" ) - langfuse_client: Final = Langfuse(**parameters) + from litellm.integrations.langfuse.langfuse_sdk import ( + acquire_langfuse_tracing, + build_langfuse_client, + release_langfuse_tracing, + ) + + tracing: Final = acquire_langfuse_tracing( + public_key=str(self.public_key), + secret_key=str(self.secret_key), + base_url=self.langfuse_host, + environment=self.langfuse_environment, + release=self.langfuse_release, + flush_interval=self.langfuse_flush_interval, + mock_mode=self.is_mock_mode, + ) + try: + api_client: Final = build_langfuse_client( + public_key=self.public_key, + secret_key=self.secret_key, + base_url=self.langfuse_host, + httpx_client=self.langfuse_client, + ) + except Exception: + release_langfuse_tracing(tracing, grace_seconds=0.0) + raise litellm.initialized_langfuse_clients += 1 verbose_logger.debug("Created langfuse client number %s", litellm.initialized_langfuse_clients) - return langfuse_client + return api_client, tracing + + def flush(self) -> None: + """Push every queued observation to Langfuse before the process goes away.""" + self.tracing.flush() + + def stop(self) -> None: + """Give the export channel back; ``DynamicLoggingCache`` calls this when a per-key logger expires.""" + from litellm.integrations.langfuse.langfuse_sdk import release_langfuse_tracing + + release_langfuse_tracing(self.tracing) @staticmethod def add_metadata_from_header(litellm_params: dict, metadata: dict) -> dict[str, object]: @@ -349,7 +445,7 @@ class LangFuseLogger: user_id: str | None = None, level: str = "DEFAULT", status_message: str | None = None, - ) -> dict: + ) -> LangfuseLoggedEvent: """ Logs a success or error event on Langfuse """ @@ -411,10 +507,10 @@ class LangFuseLogger: verbose_logger.debug("Langfuse Layer Logging - final response object: %s", response_obj) verbose_logger.info("Langfuse Layer Logging - logging success") - return {"trace_id": trace_id, "generation_id": generation_id} + return LangfuseLoggedEvent(trace_id=trace_id, generation_id=generation_id) except Exception as e: verbose_logger.exception("Langfuse Layer Error(): Exception occured - %s", e) - return {"trace_id": None, "generation_id": None} + return LangfuseLoggedEvent(trace_id=None, generation_id=None) def _get_langfuse_input_output_content( self, @@ -518,18 +614,14 @@ class LangFuseLogger: level: str, litellm_call_id: str | None, ) -> tuple: - verbose_logger.debug("Langfuse Layer Logging - logging to langfuse v2") + verbose_logger.debug("Langfuse Layer Logging - logging to langfuse via sdk v%s", self.langfuse_sdk_version) try: standard_logging_object: Final[StandardLoggingPayload | None] = cast( StandardLoggingPayload | None, kwargs.get("standard_logging_object", None), ) - tags = ( - self._get_langfuse_tags(standard_logging_object=standard_logging_object) - if self._supports_tags() - else [] - ) + tags = self._get_langfuse_tags(standard_logging_object=standard_logging_object) allowlisted_metadata: Final[StandardLoggingMetadata | Mapping[str, object]] = ( standard_logging_object["metadata"] if standard_logging_object is not None else _NO_METADATA @@ -581,17 +673,17 @@ class LangFuseLogger: # This allows continuing an existing trace while still returning the correct trace_id if existing_trace_id is not None: trace_id = existing_trace_id - resolved_trace_id: Final = ( + call_trace_id: Final = ( litellm_call_id or trace_id if existing_trace_id is None and _is_session_header_trace(trace_id, session_id, litellm_params.get("proxy_server_request")) else trace_id ) - if resolved_trace_id != trace_id: + if call_trace_id != trace_id: verbose_logger.debug( "Langfuse: trace_id %s came from a session header; using call id %s so each call gets its own trace", trace_id, - resolved_trace_id, + call_trace_id, ) requested_trace_keys: Final = _as_steering_key_sequence(clean_metadata.pop("update_trace_keys", ())) update_trace_keys: Final = ( @@ -647,7 +739,7 @@ class LangFuseLogger: trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm" else: # don't overwrite an existing trace trace_params = { - "id": resolved_trace_id, + "id": call_trace_id, "name": trace_name, "session_id": session_id, "input": masked_input if not mask_input else "redacted-by-litellm", @@ -659,10 +751,7 @@ class LangFuseLogger: for key in list(filter(lambda key: key.startswith("trace_"), clean_metadata.keys())): trace_params[key.replace("trace_", "")] = clean_metadata.pop(key, None) - if level == "ERROR": - trace_params["status_message"] = masked_output - else: - trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm" + trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm" if debug is True or (isinstance(debug, str) and debug.lower() == "true"): debug_metadata: Final = { @@ -697,17 +786,16 @@ class LangFuseLogger: ("api_base", api_base, bool(api_base)), ("vertex_location", vertex_location, bool(vertex_location)), ("aws_region_name", aws_region_name, bool(aws_region_name)), - ("cache_hit", kwargs.get("cache_hit") or False, self._supports_tags() and "cache_hit" in kwargs), + ("cache_hit", kwargs.get("cache_hit") or False, "cache_hit" in kwargs), ) enrichments: Final[Mapping[str, object]] = { key: value for key, value, include in candidate_enrichments if include } - if self._supports_tags(): - if "cache_hit" in kwargs and kwargs["cache_hit"] is None: - kwargs["cache_hit"] = False # rebind-ok: pre-existing normalization other integrations rely on - if existing_trace_id is None: - trace_params.update({"tags": tags}) + if "cache_hit" in kwargs and kwargs["cache_hit"] is None: + kwargs["cache_hit"] = False # rebind-ok: pre-existing normalization other integrations rely on + if existing_trace_id is None: + trace_params.update({"tags": tags}) proxy_server_request: Final = litellm_params.get("proxy_server_request", None) if proxy_server_request: @@ -721,17 +809,6 @@ class LangFuseLogger: if key.lower() not in _REDACTED_PROXY_HEADERS: clean_headers[key] = value - trace: Final[StatefulTraceClient] = self.Langfuse.trace(**trace_params) - - # Log provider specific information as a span - log_provider_specific_information_as_span(trace, enrichments) - - # Log guardrail information as a span - self._log_guardrail_information_as_span( - trace=trace, - standard_logging_object=standard_logging_object, - ) - generation_id = None usage = None usage_details = None @@ -753,7 +830,7 @@ class LangFuseLogger: usage = { "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, - "total_cost": cost if self._supports_costs() else None, + "total_cost": cost, } # According to langfuse documentation: "the input value must be reduced by the number of cache_read_input_tokens" input_tokens: Final = prompt_tokens - cache_read_input_tokens @@ -765,15 +842,15 @@ class LangFuseLogger: cache_read_input_tokens=cache_read_input_tokens, ) - generation_name = clean_metadata.pop("generation_name", None) - if generation_name is None: - # if `generation_name` is None, use sensible default values - # If using litellm proxy user `key_alias` if not None - # If `key_alias` is None, just log `litellm-{call_type}` as the generation name - _user_api_key_alias: Final = cast(str | None, clean_metadata.get("user_api_key_alias", None)) - generation_name = f"litellm-{cast(str, kwargs.get('call_type', 'completion'))}" - if _user_api_key_alias is not None: - generation_name = f"litellm:{_user_api_key_alias}" + requested_generation_name: Final = clean_metadata.pop("generation_name", None) + _user_api_key_alias: Final = cast(str | None, clean_metadata.get("user_api_key_alias", None)) + generation_name: Final = ( + str(requested_generation_name) + if requested_generation_name is not None + else f"litellm:{_user_api_key_alias}" + if _user_api_key_alias is not None + else f"litellm-{cast(str, kwargs.get('call_type', 'completion'))}" + ) if response_obj is not None: system_fingerprint = getattr(response_obj, "system_fingerprint", None) @@ -789,53 +866,97 @@ class LangFuseLogger: generation_params = { "name": generation_name, "id": clean_metadata.pop("generation_id", generation_id), - "start_time": start_time, - "end_time": end_time, - "model": model_name, - "model_parameters": optional_params, "input": masked_input if not mask_input else "redacted-by-litellm", "output": masked_output if not mask_output else "redacted-by-litellm", - "usage": usage, - "usage_details": usage_details, - "metadata": { - **log_requester_metadata(redact_user_api_key_info(metadata=allowlisted_metadata)), + "cost_details": {"total": cost} # mutable-ok: langfuse serializes this payload + if usage is not None and isinstance(cost, (int, float)) + else None, + "metadata": { # mutable-ok: langfuse serializes this payload, a proxy is not json-encodable + **log_requester_metadata(redact_user_api_key_info(metadata=allowlisted_metadata)), # pyright: ignore[reportArgumentType] # TypedDict in, plain metadata dict out **enrichments, + **_lookup_ids(litellm_call_id, response_obj), }, - "level": level, - "version": clean_metadata.pop("version", None), + "version": _optional_str(clean_metadata.pop("version", None)), } parent_observation_id: Final = metadata.get("parent_observation_id", None) - if parent_observation_id is not None: - generation_params["parent_observation_id"] = parent_observation_id - - if self._supports_prompt(): - generation_params = _add_prompt_to_generation_params( - generation_params=generation_params, - clean_metadata=clean_metadata, - prompt_management_metadata=prompt_management_metadata, - langfuse_client=self.Langfuse, - ) + generation_params = _add_prompt_to_generation_params( + generation_params=generation_params, + clean_metadata=clean_metadata, + prompt_management_metadata=prompt_management_metadata, + langfuse_client=self.api_client, + ) if masked_output is not None and isinstance(masked_output, str) and level == "ERROR": generation_params["status_message"] = masked_output - if self._supports_completion_start_time(): - generation_params["completion_start_time"] = kwargs.get("completion_start_time", None) + # langfuse ships in the proxy-runtime extra, so this module must import cleanly without it + from litellm.integrations.langfuse.langfuse_sdk import ( + observation_attributes, + resolve_observation_id, + resolve_trace_id, + start_generation, + trace_attributes, + ) - generation_client: Final = trace.generation(**generation_params) + resolved_trace_id: Final = resolve_trace_id(call_trace_id) # pyright: ignore[reportArgumentType] # metadata value, str or None at runtime + continued_trace: Final = existing_trace_id is not None + generation_is_trace_root: Final = not continued_trace and parent_observation_id is None + trace_public: Final = _trace_public_flag(trace_params.get("public")) + trace_input: Final = trace_params.get("input") + trace_output: Final = trace_params.get("output") + trace_level_attributes: Final = trace_attributes( + name=trace_params.get("name"), + user_id=trace_params.get("user_id"), + session_id=trace_params.get("session_id"), + version=trace_params.get("version"), + release=trace_params.get("release"), + tags=trace_params.get("tags"), + metadata=trace_params.get("metadata"), + public=trace_public, + input=None if generation_is_trace_root and trace_input == generation_params["input"] else trace_input, + output=None + if generation_is_trace_root and trace_output == generation_params["output"] + else trace_output, + ) + generation_attributes: Final = observation_attributes( + observation_type="generation", + input=generation_params["input"], + output=generation_params["output"], + metadata=generation_params["metadata"], + level=level, + status_message=generation_params.get("status_message"), + version=generation_params["version"], + model=model_name, + model_parameters=optional_params, + usage_details=usage_details, + cost_details=generation_params["cost_details"], + completion_start_time=kwargs.get("completion_start_time", None), + prompt=generation_params.get("prompt"), + ) + generation: Final = start_generation( + tracing=self.tracing, + trace_id=resolved_trace_id, + parent_observation_id=resolve_observation_id(parent_observation_id), # pyright: ignore[reportArgumentType] # metadata value, str or None at runtime + existing_trace=continued_trace, + observation_id=resolve_observation_id(generation_params["id"]), + name=generation_params["name"], # pyright: ignore[reportArgumentType] # always the str set a few lines up + start_time=start_time, + public=trace_public, + attributes=MappingProxyType({**generation_attributes, **trace_level_attributes}), + ) + try: + log_provider_specific_information_as_span( + tracing=self.tracing, parent=generation, enrichments=enrichments + ) + self._log_guardrail_information_as_span( + tracing=self.tracing, parent=generation, standard_logging_object=standard_logging_object + ) + finally: + generation.end(end_time) - # Return the trace_id we set (which should be litellm_call_id when no explicit trace_id provided) - # We explicitly set trace_id in trace_params["id"], so langfuse should use it - # Verify langfuse accepted our trace_id; if it differs, log a warning but still return our intended value - # to match expected test behavior - if hasattr(generation_client, "trace_id") and generation_client.trace_id: - if generation_client.trace_id != resolved_trace_id: - verbose_logger.warning( - "Langfuse trace_id mismatch: set %s, but langfuse returned %s. Using our intended trace_id for consistency.", - resolved_trace_id, - generation_client.trace_id, - ) - return resolved_trace_id, generation_id + # log_event_on_langfuse tuple-unpacks this and re-wraps it in the dict callers cache. + # The observation id is the requested generation_id after resolve_observation_id. + return resolved_trace_id, generation.id except Exception: verbose_logger.error("Langfuse Layer Error - %s", traceback.format_exc()) return None, None @@ -904,27 +1025,11 @@ class LangFuseLogger: _cache_key = _hidden_params.get("cache_key", None) if _cache_key is None and litellm.cache is not None: # fallback to using "preset_cache_key" - _preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) + _preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) # pyright: ignore[reportPrivateUsage] # kwargs-ok: no public preset-cache-key accessor _cache_key = _preset_cache_key tags.append(f"cache_key:{_cache_key}") return tags - def _supports_tags(self): - """Check if current langfuse version supports tags""" - return Version(self.langfuse_sdk_version) >= Version("2.6.3") - - def _supports_prompt(self): - """Check if current langfuse version supports prompt""" - return Version(self.langfuse_sdk_version) >= Version("2.7.3") - - def _supports_costs(self): - """Check if current langfuse version supports costs""" - return Version(self.langfuse_sdk_version) >= Version("2.7.3") - - def _supports_completion_start_time(self): - """Check if current langfuse version supports completion start time""" - return Version(self.langfuse_sdk_version) >= Version("2.7.3") - @staticmethod def _apply_masking_function(data: object, masking_function: Callable[[object], object]) -> object: """ @@ -973,23 +1078,24 @@ class LangFuseLogger: @staticmethod def _get_langfuse_flush_interval(flush_interval: int) -> int: - """ - Get the langfuse flush interval to initialize the Langfuse client - - Reads `LANGFUSE_FLUSH_INTERVAL` from the environment variable. - If not set, uses the flush interval passed in as an argument. - - Args: - flush_interval: The flush interval to use if LANGFUSE_FLUSH_INTERVAL is not set - - Returns: - [int] The flush interval to use to initialize the Langfuse client - """ - return int(os.getenv("LANGFUSE_FLUSH_INTERVAL") or flush_interval) + """``LANGFUSE_FLUSH_INTERVAL`` in whole seconds above 0 (the export scheduler's delay), else ``flush_interval``.""" + raw: Final = os.getenv("LANGFUSE_FLUSH_INTERVAL") + if not raw: + return flush_interval + parsed: Final = int(raw) if raw.strip().isdigit() else None + if parsed is None or parsed <= 0: + verbose_logger.warning( + "LANGFUSE_FLUSH_INTERVAL=%r is not a whole number of seconds above 0; flushing every %d s", + raw, + flush_interval, + ) + return flush_interval + return parsed def _log_guardrail_information_as_span( self, - trace: StatefulTraceClient, + tracing: "LangfuseTracing", + parent: "LangfuseObservation", standard_logging_object: StandardLoggingPayload | None, ): """ @@ -1011,6 +1117,8 @@ class LangFuseLogger: ) return + from litellm.integrations.langfuse.langfuse_sdk import observation_attributes, start_child_span + for guardrail_entry in guardrail_information: if not isinstance(guardrail_entry, dict): verbose_logger.debug( @@ -1019,30 +1127,35 @@ class LangFuseLogger: ) continue - span = trace.span( + span = start_child_span( + tracing=tracing, + parent=parent, name="guardrail", - input=guardrail_entry.get("guardrail_request", None), - output=guardrail_entry.get("guardrail_response", None), - metadata={ - "guardrail_name": guardrail_entry.get("guardrail_name", None), - "guardrail_mode": guardrail_entry.get("guardrail_mode", None), - "guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None), - }, start_time=guardrail_entry.get("start_time", None), - end_time=guardrail_entry.get("end_time", None), + attributes=observation_attributes( + observation_type="span", + input=guardrail_entry.get("guardrail_request", None), + output=guardrail_entry.get("guardrail_response", None), + metadata=MappingProxyType( + { + "guardrail_name": guardrail_entry.get("guardrail_name", None), + "guardrail_mode": guardrail_entry.get("guardrail_mode", None), + "guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None), + } + ), + ), ) verbose_logger.debug("Logged guardrail information as span: %s", span) - span.end() + span.end(guardrail_entry.get("end_time", None)) def _add_prompt_to_generation_params( generation_params: dict, clean_metadata: dict, prompt_management_metadata: StandardLoggingPromptManagementMetadata | None, - langfuse_client: object, + langfuse_client: "LangfuseApiClient", ) -> dict: - from langfuse import Langfuse from langfuse.model import ( ChatPromptClient, Prompt_Chat, @@ -1050,8 +1163,6 @@ def _add_prompt_to_generation_params( TextPromptClient, ) - langfuse_client = cast(Langfuse, langfuse_client) - user_prompt: Final = clean_metadata.pop("prompt", None) if user_prompt is None and prompt_management_metadata is None: pass @@ -1075,7 +1186,7 @@ def _add_prompt_to_generation_params( if "labels" in prompt_text_params and "tags" in prompt_text_params: _data["labels"] = user_prompt.get("labels", []) or [] _data["tags"] = user_prompt.get("tags", []) or [] - _prompt_obj = Prompt_Text(**_data) + _prompt_obj = Prompt_Text(**_data) # pyright: ignore[reportArgumentType] # kwargs-ok: shape mirrors the pydantic model, values from the user's prompt dict generation_params["prompt"] = TextPromptClient(prompt=_prompt_obj) elif isinstance(user_prompt["prompt"], list): @@ -1090,7 +1201,7 @@ def _add_prompt_to_generation_params( _data["labels"] = user_prompt.get("labels", []) or [] _data["tags"] = user_prompt.get("tags", []) or [] - _prompt_obj = Prompt_Chat(**_data) + _prompt_obj = Prompt_Chat(**_data) # pyright: ignore[reportArgumentType] # kwargs-ok: shape mirrors the pydantic model, values from the user's prompt dict generation_params["prompt"] = ChatPromptClient(prompt=_prompt_obj) else: @@ -1110,21 +1221,14 @@ def _add_prompt_to_generation_params( def log_provider_specific_information_as_span( - trace, - clean_metadata: Mapping[str, Any], + *, + tracing: "LangfuseTracing", + parent: "LangfuseObservation", + enrichments: Mapping[str, Any], ): - """ - Logs provider-specific information as spans. + """Logs provider-specific information as spans under the generation.""" - Parameters: - trace: The tracing object used to log spans. - clean_metadata: A dictionary containing metadata to be logged. - - Returns: - None - """ - - _hidden_params: Final[Mapping[str, object] | None] = clean_metadata.get("hidden_params", None) + _hidden_params: Final[Mapping[str, object] | None] = enrichments.get("hidden_params", None) if _hidden_params is None: return @@ -1135,22 +1239,27 @@ def log_provider_specific_information_as_span( for elem in vertex_ai_grounding_metadata: if isinstance(elem, dict): for key, value in elem.items(): - trace.span( - name=key, - input=value, - ) + _end_grounding_span(tracing=tracing, parent=parent, name=key, value=value) else: - trace.span( - name="vertex_ai_grounding_metadata", - input=elem, - ) + _end_grounding_span(tracing=tracing, parent=parent, name="vertex_ai_grounding_metadata", value=elem) else: - trace.span( - name="vertex_ai_grounding_metadata", - input=vertex_ai_grounding_metadata, + _end_grounding_span( + tracing=tracing, parent=parent, name="vertex_ai_grounding_metadata", value=vertex_ai_grounding_metadata ) +def _end_grounding_span(*, tracing: "LangfuseTracing", parent: "LangfuseObservation", name: str, value: object) -> None: + from litellm.integrations.langfuse.langfuse_sdk import observation_attributes, start_child_span + + start_child_span( + tracing=tracing, + parent=parent, + name=name, + start_time=None, + attributes=observation_attributes(observation_type="span", input=value), + ).end() + + def log_requester_metadata(clean_metadata: Mapping[str, Any]): returned_metadata: Final = {} requester_metadata: Final = clean_metadata.get("requester_metadata") or {} diff --git a/litellm/integrations/langfuse/langfuse_prompt_management.py b/litellm/integrations/langfuse/langfuse_prompt_management.py index 90db0626e23..3786087ba91 100644 --- a/litellm/integrations/langfuse/langfuse_prompt_management.py +++ b/litellm/integrations/langfuse/langfuse_prompt_management.py @@ -2,16 +2,14 @@ Call Hook for LiteLLM Proxy which allows Langfuse prompt management. """ -import inspect -import os from functools import lru_cache from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, cast -from packaging.version import Version - from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prompt_management_base import PromptManagementClient from litellm.litellm_core_utils.asyncify import run_async_function +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.types.integrations.langfuse import LangfuseLoggedEvent from litellm.types.llms.openai import AllMessageValues, ChatCompletionSystemMessage from litellm.types.prompts.init_prompts import PromptSpec from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPayload @@ -19,17 +17,27 @@ from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPa from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import ( DynamicLoggingCache, ) +from ...litellm_core_utils.specialty_caches.service_trace_id_cache import in_memory_trace_id_cache from ..prompt_management_base import PromptManagementBase -from .langfuse import LangFuseLogger, resolve_langfuse_credentials +from .langfuse import ( + LangFuseLogger, + installed_langfuse_version, + raise_if_unsupported_langfuse_version, + raise_if_unusable_prompt_cache_ttl, + resolve_langfuse_credentials, + warn_if_upstream_langfuse_configured, +) from .langfuse_handler import LangFuseHandler +from .langfuse_mock_client import create_mock_langfuse_client, should_use_langfuse_mock if TYPE_CHECKING: - from langfuse import Langfuse - from langfuse.client import ChatPromptClient, TextPromptClient + from langfuse.model import ChatPromptClient, TextPromptClient from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj - LangfuseClass: TypeAlias = Langfuse + from .langfuse_sdk import LangfuseApiClient + + LangfuseClass: TypeAlias = LangfuseApiClient PROMPT_CLIENT = TextPromptClient | ChatPromptClient else: @@ -49,23 +57,24 @@ def langfuse_client_init( allow_env_credentials: bool = True, ) -> LangfuseClass: """ - Initialize Langfuse client with caching to prevent multiple initializations. + Initialize the Langfuse REST client with caching to prevent multiple initializations. Args: langfuse_public_key (str, optional): Public key for Langfuse. Defaults to None. langfuse_secret (str, optional): Secret key for Langfuse. Defaults to None. langfuse_host (str, optional): Host URL for Langfuse. Defaults to None. - flush_interval (int, optional): Flush interval in seconds. Defaults to 1. + flush_interval (int, optional): Kept in the signature so cached callers keep their cache key. Returns: - Langfuse: Initialized Langfuse client instance + LangfuseApiClient: prompt, auth and project lookups for one credential set Raises: Exception: If langfuse package is not installed """ + raise_if_unsupported_langfuse_version(installed_langfuse_version()) + raise_if_unusable_prompt_cache_ttl() try: - import langfuse - from langfuse import Langfuse + from .langfuse_sdk import build_langfuse_client except Exception as e: raise Exception( f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m" @@ -83,39 +92,22 @@ def langfuse_client_init( # add http:// if unset, assume communicating over private network - e.g. render langfuse_host = "http://" + langfuse_host - langfuse_release: Final = os.getenv("LANGFUSE_RELEASE") - langfuse_debug: Final = os.getenv("LANGFUSE_DEBUG") + warn_if_upstream_langfuse_configured() - parameters: Final = { - "public_key": public_key, - "secret_key": secret_key, - "host": langfuse_host, - "release": langfuse_release, - "debug": langfuse_debug, - "flush_interval": LangFuseLogger._get_langfuse_flush_interval(flush_interval), # flush interval in seconds - } + httpx_client: Final = create_mock_langfuse_client() if should_use_langfuse_mock() else HTTPHandler().client + return build_langfuse_client( + public_key=public_key, + secret_key=secret_key, + base_url=langfuse_host, + httpx_client=httpx_client, + ) - if Version(langfuse.version.__version__) >= Version("2.6.0"): - parameters["sdk_integration"] = "litellm" - if Version(langfuse.version.__version__) >= Version("2.7.3"): - import httpx - - import litellm - - from ...llms.custom_httpx.http_handler import get_ssl_configuration - - parameters["httpx_client"] = httpx.Client( - verify=get_ssl_configuration(), - cert=os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate), - ) - - if "environment" in inspect.signature(Langfuse.__init__).parameters: - parameters["environment"] = LangFuseLogger.resolve_deployment_environment() - - client: Final = Langfuse(**parameters) - - return client +def _remember_trace_id(litellm_call_id: object, logged: LangfuseLoggedEvent) -> None: + trace_id: Final = logged["trace_id"] + if not isinstance(litellm_call_id, str) or trace_id is None: + return + in_memory_trace_id_cache.set_cache(litellm_call_id=litellm_call_id, service_name="langfuse", trace_id=trace_id) class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogger): @@ -126,15 +118,33 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_host=None, flush_interval=1, ): - import langfuse - self.langfuse_sdk_version = langfuse.version.__version__ - self.Langfuse = langfuse_client_init( + self.langfuse_sdk_version = installed_langfuse_version() + raise_if_unsupported_langfuse_version(self.langfuse_sdk_version) + raise_if_unusable_prompt_cache_ttl() + + from .langfuse_sdk import acquire_langfuse_tracing, configured_release + + self.api_client = langfuse_client_init( langfuse_public_key=langfuse_public_key, langfuse_secret=langfuse_secret, langfuse_host=langfuse_host, flush_interval=flush_interval, ) + self.public_key, self.secret_key, self.langfuse_host = resolve_langfuse_credentials( + langfuse_public_key=langfuse_public_key, + langfuse_secret=langfuse_secret, + langfuse_host=langfuse_host, + ) + self.tracing = acquire_langfuse_tracing( + public_key=str(self.public_key), + secret_key=str(self.secret_key), + base_url=self.langfuse_host, + environment=LangFuseLogger.resolve_deployment_environment(), + release=configured_release(), + flush_interval=LangFuseLogger._get_langfuse_flush_interval(flush_interval), # pyright: ignore[reportPrivateUsage] # shared env-fallback helper, not part of the logger's API + mock_mode=should_use_langfuse_mock(), + ) @property def integration_name(self): @@ -228,11 +238,8 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_host=dynamic_callback_params.get("langfuse_host"), allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None, ) - langfuse_prompt_client: Final = self._get_prompt_from_id( - langfuse_prompt_id=prompt_id, - langfuse_client=langfuse_client, - ) - return langfuse_prompt_client is not None + self._get_prompt_from_id(langfuse_prompt_id=prompt_id, langfuse_client=langfuse_client) + return True def _compile_prompt_helper( self, @@ -311,13 +318,14 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge standard_callback_dynamic_params=standard_callback_dynamic_params, in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache, ) - langfuse_logger_to_use.log_event_on_langfuse( + logged: Final = langfuse_logger_to_use.log_event_on_langfuse( kwargs=kwargs, response_obj=response_obj, start_time=start_time, end_time=end_time, user_id=kwargs.get("user", None), ) + _remember_trace_id(litellm_call_id=kwargs.get("litellm_call_id"), logged=logged) except Exception as e: from litellm._logging import verbose_logger @@ -339,7 +347,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge status_message = str(kwargs.get("exception", "Unknown error")) if standard_logging_object is not None: status_message = standard_logging_object.get("error_str", None) or status_message - langfuse_logger_to_use.log_event_on_langfuse( + logged: Final = langfuse_logger_to_use.log_event_on_langfuse( start_time=start_time, end_time=end_time, response_obj=None, @@ -348,6 +356,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge level="ERROR", kwargs=kwargs, ) + _remember_trace_id(litellm_call_id=kwargs.get("litellm_call_id"), logged=logged) except Exception as e: from litellm._logging import verbose_logger diff --git a/litellm/integrations/langfuse/langfuse_sdk.py b/litellm/integrations/langfuse/langfuse_sdk.py new file mode 100644 index 00000000000..66819c95ebf --- /dev/null +++ b/litellm/integrations/langfuse/langfuse_sdk.py @@ -0,0 +1,1213 @@ +from __future__ import annotations + +import logging +import os +import re +import threading +from base64 import b64encode +from collections.abc import Iterable, Mapping, Sequence +from contextvars import ContextVar +from dataclasses import dataclass, replace +from datetime import datetime +from functools import partial, reduce +from hashlib import sha256 +from importlib.metadata import version +from itertools import chain +from time import monotonic, sleep +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import quote + +import httpx +import opentelemetry.trace as otel_trace +from langfuse import LangfuseOtelSpanAttributes +from langfuse.api import LangfuseAPI, Prompt, Prompt_Chat +from langfuse.api.core.api_error import ApiError +from langfuse.api.core.request_options import RequestOptions +from langfuse.model import BasePromptClient, ChatPromptClient, PromptClient, TextPromptClient +from opentelemetry.context import Context +from opentelemetry.exporter.otlp.proto.common.trace_encoder import encode_spans +from opentelemetry.sdk.resources import Resource +from opentelemetry.sdk.trace import ReadableSpan, SpanLimits, TracerProvider +from opentelemetry.sdk.trace.export import BatchSpanProcessor, SpanExporter, SpanExportResult +from opentelemetry.sdk.trace.id_generator import RandomIdGenerator +from opentelemetry.sdk.trace.sampling import ALWAYS_ON, Decision, Sampler, SamplingResult +from opentelemetry.trace import Link, NonRecordingSpan, Span, SpanContext, SpanKind, TraceFlags, Tracer, TraceState +from opentelemetry.util.types import Attributes, AttributeValue +from pydantic import BaseModel, ConfigDict + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.langfuse.langfuse import PROMPT_CACHE_TTL_ENV, parse_langfuse_debug, whole_number +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client + +__all__ = ( + "AuthCheckFailure", + "DiscardingSpanExporter", + "LangfuseApiClient", + "LangfuseObservation", + "LangfusePromptError", + "LangfuseSpanExporter", + "LangfuseTracing", + "TraceIdHashSampler", + "acquire_langfuse_tracing", + "build_langfuse_client", + "build_langfuse_tracing", + "configured_flush_at", + "configured_max_retries", + "configured_release", + "configured_sample_rate", + "configured_timeout", + "enable_langfuse_debug_logging", + "flush_langfuse_tracing", + "observation_attributes", + "release_langfuse_tracing", + "resolve_observation_id", + "resolve_trace_id", + "start_child_span", + "start_generation", + "to_unix_nanos", + "trace_attributes", +) + +_TRACE_ID_PATTERN: Final = re.compile(r"^(?=.*[1-9a-f])[0-9a-f]{32}$") +_OBSERVATION_ID_PATTERN: Final = re.compile(r"^(?=.*[1-9a-f])[0-9a-f]{16}$") +_TRACER_NAME: Final = "langfuse-sdk" +_LANGFUSE_INGESTION_VERSION_HEADER: Final = "x-langfuse-ingestion-version" +_LANGFUSE_INGESTION_VERSION: Final = "4" +_NO_REST_RETRIES: Final = RequestOptions(max_retries=0) +_TRUNCATION_MARKER: Final = "" +_METADATA_PREFIXES: Final = (LangfuseOtelSpanAttributes.OBSERVATION_METADATA, LangfuseOtelSpanAttributes.TRACE_METADATA) +_TRUNCATION_GROUPS: Final = ( + (LangfuseOtelSpanAttributes.OBSERVATION_INPUT, LangfuseOtelSpanAttributes.TRACE_INPUT), + (LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT, LangfuseOtelSpanAttributes.TRACE_OUTPUT), + _METADATA_PREFIXES, +) +_SERVER_FLOOR_HINT: Final = ( + "; the OTLP traces route needs a self-hosted Langfuse server on 3.63.0 or newer " + "(https://langfuse.com/self-hosting/upgrade/versioning#sdk-server)" +) +_langfuse_logger: Final = logging.getLogger("langfuse") +_MAX_QUEUE_SIZE: Final = 100_000 +_DEFAULT_FLUSH_AT: Final = 512 +_CHANNEL_RETIRE_GRACE_SECONDS: Final = 60.0 +_DEFAULT_TIMEOUT_SECONDS: Final = 20.0 +_DEFAULT_MAX_RETRIES: Final = 3 +_MAX_RETRIES: Final = 1_000 +_MAX_BACKOFF_EXPONENT: Final = 6 +_DEFAULT_PROMPT_CACHE_TTL_SECONDS: Final = 60.0 +_JSON_SAFE_INT: Final = 2**53 - 1 +_COMMON_RELEASE_ENVS: Final = ( + "RENDER_GIT_COMMIT", + "CI_COMMIT_SHA", + "CIRCLE_SHA1", + "SOURCE_VERSION", + "TRAVIS_COMMIT", + "GIT_COMMIT", + "GITHUB_SHA", + "BITBUCKET_COMMIT", + "BUILD_SOURCEVERSION", + "DRONE_COMMIT_SHA", +) +_SPAN_LIMITS: Final = SpanLimits( + max_attributes=SpanLimits.UNSET, + max_events=128, + max_links=128, + max_span_attributes=SpanLimits.UNSET, + max_event_attributes=128, + max_link_attributes=128, + max_attribute_length=SpanLimits.UNSET, + max_span_attribute_length=SpanLimits.UNSET, +) + + +def to_unix_nanos(value: datetime | float | None) -> int | None: + """Langfuse v4 takes OTel timestamps, which are integer nanoseconds since the epoch. + + Guardrail entries carry unix seconds as floats rather than datetimes, so both + shapes have to convert; the v2 SDK accepted either through a pydantic model. + """ + if value is None: + return None + seconds: Final = value.timestamp() if isinstance(value, datetime) else float(value) + return int(seconds * 1_000_000_000) + + +def resolve_trace_id(trace_id: object | None) -> str: + """Map a caller's trace id onto the 32 lowercase hex characters v4 requires.""" + serialized: Final = "" if trace_id is None else str(trace_id) + normalized: Final = serialized.lower().replace("-", "") + if _TRACE_ID_PATTERN.fullmatch(normalized): + return normalized + if not serialized: + return format(RandomIdGenerator().generate_trace_id(), "032x") + return sha256(serialized.encode("utf-8")).digest()[:16].hex() + + +def resolve_observation_id(observation_id: object | None) -> str | None: + """Map a caller's parent observation id onto v4's 16 lowercase hex characters.""" + serialized: Final = "" if observation_id is None else str(observation_id) + normalized: Final = serialized.lower().replace("-", "") + if _OBSERVATION_ID_PATTERN.fullmatch(normalized): + return normalized + if not serialized: + return None + return sha256(serialized.encode("utf-8")).digest()[:8].hex() + + +def _serialize(value: object) -> str | None: + return value if value is None or isinstance(value, str) else safe_dumps(value) + + +def _string_or_none(value: object) -> str | None: + return None if value is None else str(value) + + +def _serialize_datetime(value: object) -> str | None: + """A datetime the way the SDK's ``EventSerializer`` sends one: a JSON string, naive values read as local time.""" + if isinstance(value, datetime): + return safe_dumps(value.astimezone().isoformat()) + return _serialize(value) + + +def _strings(items: Iterable[object]) -> tuple[str, ...]: + return tuple(str(item) for item in items) + + +def _string_sequence(value: object) -> Sequence[str] | None: + if value is None: + return None + if isinstance(value, (list, tuple, set, frozenset)): + return _strings(value) or None + return (str(value),) + + +def _present(entries: Iterable[tuple[str, AttributeValue | None]]) -> Mapping[str, AttributeValue]: + return MappingProxyType({key: value for key, value in entries if value is not None}) + + +def _metadata_value(value: object) -> AttributeValue | None: + """A metadata value as it survives the trip: OTLP drops ints past int64 and a JSON reader rounds ints past + 2**53, so those go as strings, which is how v2's readback showed them.""" + if isinstance(value, (str, bool)): + return value + if isinstance(value, int) and -_JSON_SAFE_INT <= value <= _JSON_SAFE_INT: + return value + return _serialize(value) + + +def _flattened_metadata(prefix: str, metadata: object) -> Mapping[str, AttributeValue]: + """Mirror the SDK's wire shape: one ``.`` attribute per key, or ```` for a non-dict.""" + if metadata is None: + return _present(()) + if not isinstance(metadata, Mapping): + return _present(((prefix, _serialize(metadata)),)) + return _present((f"{prefix}.{key}", _metadata_value(value)) for key, value in metadata.items()) + + +def trace_attributes( + *, + name: object = None, + user_id: object = None, + session_id: object = None, + version: object = None, + release: object = None, + tags: object = None, + metadata: object = None, + public: bool | None = None, + input: object = None, + output: object = None, +) -> Mapping[str, AttributeValue]: + """Trace-level fields ride on an observation's span as ``langfuse.trace.*`` style attributes in v4. + + On the root observation they define the trace; on a continuation they update it, which is + how v2's ``trace(...)`` and ``update_trace_keys`` contracts map onto the OTLP ingestion. + """ + scalar: Final[tuple[tuple[str, str | bool | None], ...]] = ( + (LangfuseOtelSpanAttributes.TRACE_NAME, _string_or_none(name)), + (LangfuseOtelSpanAttributes.TRACE_USER_ID, _string_or_none(user_id)), + (LangfuseOtelSpanAttributes.TRACE_SESSION_ID, _string_or_none(session_id)), + (LangfuseOtelSpanAttributes.VERSION, _string_or_none(version)), + (LangfuseOtelSpanAttributes.RELEASE, _string_or_none(release)), + (LangfuseOtelSpanAttributes.TRACE_PUBLIC, public), + (LangfuseOtelSpanAttributes.TRACE_INPUT, _serialize(input)), + (LangfuseOtelSpanAttributes.TRACE_OUTPUT, _serialize(output)), + ) + tags_entry: Final[tuple[str, Sequence[str] | None]] = ( + LangfuseOtelSpanAttributes.TRACE_TAGS, + _string_sequence(tags), + ) + return _present( + chain(scalar, (tags_entry,), _flattened_metadata(LangfuseOtelSpanAttributes.TRACE_METADATA, metadata).items()) + ) + + +def observation_attributes( + *, + observation_type: Literal["generation", "span"], + input: object = None, + output: object = None, + metadata: object = None, + level: object = None, + status_message: object = None, + version: object = None, + model: object = None, + model_parameters: object = None, + usage_details: object = None, + cost_details: object = None, + completion_start_time: object = None, + prompt: object = None, +) -> Mapping[str, AttributeValue]: + """The observation's own fields, serialized the way the SDK's ``create_generation_attributes`` does. + + ``prompt`` links the generation to a managed prompt only when it is a real prompt client; + v2 dropped anything else, and a fallback prompt has no server-side version to link. + """ + linked_prompt: Final = prompt if isinstance(prompt, BasePromptClient) and not prompt.is_fallback else None + scalar: Final[tuple[tuple[str, str | int | None], ...]] = ( + (LangfuseOtelSpanAttributes.OBSERVATION_TYPE, observation_type), + (LangfuseOtelSpanAttributes.OBSERVATION_LEVEL, _string_or_none(level)), + (LangfuseOtelSpanAttributes.OBSERVATION_STATUS_MESSAGE, _string_or_none(status_message)), + (LangfuseOtelSpanAttributes.VERSION, _string_or_none(version)), + (LangfuseOtelSpanAttributes.OBSERVATION_INPUT, _serialize(input)), + (LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT, _serialize(output)), + (LangfuseOtelSpanAttributes.OBSERVATION_MODEL, _string_or_none(model)), + (LangfuseOtelSpanAttributes.OBSERVATION_MODEL_PARAMETERS, _serialize(model_parameters)), + (LangfuseOtelSpanAttributes.OBSERVATION_USAGE_DETAILS, _serialize(usage_details)), + (LangfuseOtelSpanAttributes.OBSERVATION_COST_DETAILS, _serialize(cost_details)), + (LangfuseOtelSpanAttributes.OBSERVATION_COMPLETION_START_TIME, _serialize_datetime(completion_start_time)), + (LangfuseOtelSpanAttributes.OBSERVATION_PROMPT_NAME, linked_prompt.name if linked_prompt else None), + (LangfuseOtelSpanAttributes.OBSERVATION_PROMPT_VERSION, linked_prompt.version if linked_prompt else None), + ) + return _present( + chain(scalar, _flattened_metadata(LangfuseOtelSpanAttributes.OBSERVATION_METADATA, metadata).items()) + ) + + +@dataclass(frozen=True, slots=True) +class LangfuseObservation: + """A Langfuse observation as the OTel span litellm exports for it.""" + + span: Span + public: bool | None + + @property + def id(self) -> str: + return format(self.span.get_span_context().span_id, "016x") + + @property + def trace_id(self) -> str: + return format(self.span.get_span_context().trace_id, "032x") + + def end(self, end_time: datetime | float | None = None) -> None: + self.span.end(end_time=to_unix_nanos(end_time)) + + +_requested_trace_id: Final[ContextVar[int | None]] = ContextVar("litellm_langfuse_requested_trace_id", default=None) +_requested_span_id: Final[ContextVar[int | None]] = ContextVar("litellm_langfuse_requested_span_id", default=None) + + +class _RequestedIdGenerator(RandomIdGenerator): + """Hand out the ids the calling context asked for, random otherwise. + + v2 took caller trace and generation ids as plain fields; OTel derives both from + the tracer's id generator, so the request rides on a context variable instead. + """ + + def generate_trace_id(self) -> int: + requested: Final = _requested_trace_id.get() + return super().generate_trace_id() if requested is None else requested + + def generate_span_id(self) -> int: + requested: Final = _requested_span_id.get() + return super().generate_span_id() if requested is None else requested + + +def _parent_context(*, trace_id: str, parent_observation_id: str | None, existing_trace: bool) -> Context: + """Where a new observation hangs: nowhere for a fresh trace, under a remote parent when continuing one. + + ``existing_trace`` is the v2 ``existing_trace_id`` contract: the trace is appended to, never + rewritten. The server takes a root observation's name and I/O as the trace's, so a continuation + without a known parent hangs under a parent id that is never exported instead of claiming root. + An explicitly empty context also keeps the caller's own active span out of the picture. + """ + if parent_observation_id is None and not existing_trace: + return Context() + parent_span_id: Final = ( + int(parent_observation_id, 16) if parent_observation_id is not None else RandomIdGenerator().generate_span_id() + ) + remote_parent: Final = NonRecordingSpan( + SpanContext( + trace_id=int(trace_id, 16), + span_id=parent_span_id, + is_remote=True, + trace_flags=TraceFlags(TraceFlags.SAMPLED), + ) + ) + return otel_trace.set_span_in_context(remote_parent) + + +def _start_span( + tracer: Tracer, + *, + name: str, + context: Context, + start_time: datetime | float | None, + trace_id: str | None, + observation_id: str | None, + attributes: Mapping[str, AttributeValue], +) -> Span: + trace_token: Final = _requested_trace_id.set(int(trace_id, 16) if trace_id is not None else None) + span_token: Final = _requested_span_id.set(int(observation_id, 16) if observation_id is not None else None) + try: + return tracer.start_span( + name=name, context=context, start_time=to_unix_nanos(start_time), attributes=attributes + ) + finally: + _requested_span_id.reset(span_token) + _requested_trace_id.reset(trace_token) + + +def start_generation( + *, + tracing: LangfuseTracing, + trace_id: str, + parent_observation_id: str | None, + existing_trace: bool, + observation_id: str | None, + name: str, + start_time: datetime | float | None, + public: bool | None, + attributes: Mapping[str, AttributeValue], +) -> LangfuseObservation: + """Create the generation for one model call, timed from when that call began. + + ``trace_id``, ``parent_observation_id`` and ``observation_id`` are the v2 ``trace(id=...)``, + ``generation(parent_observation_id=...)`` and ``generation(id=...)`` arguments, already + normalized by ``resolve_trace_id`` and ``resolve_observation_id``. + """ + span: Final = _start_span( + tracing.tracer, + name=name, + context=_parent_context( + trace_id=trace_id, parent_observation_id=parent_observation_id, existing_trace=existing_trace + ), + start_time=start_time, + trace_id=trace_id, + observation_id=observation_id, + attributes=attributes, + ) + return LangfuseObservation(span=span, public=public) + + +def start_child_span( + *, + tracing: LangfuseTracing, + parent: LangfuseObservation, + name: str, + start_time: datetime | float | None, + attributes: Mapping[str, AttributeValue], +) -> LangfuseObservation: + """Create an observation under the generation, keeping its own time window. + + The server folds the trace's ``public`` flag across every observation, with a missing + attribute read as ``False``, so the child repeats the generation's value. + """ + public_entry: Final[tuple[str, bool | None]] = (LangfuseOtelSpanAttributes.TRACE_PUBLIC, parent.public) + span: Final = _start_span( + tracing.tracer, + name=name, + context=otel_trace.set_span_in_context(parent.span), + start_time=start_time, + trace_id=None, + observation_id=None, + attributes=_present(chain((public_entry,), attributes.items())), + ) + return LangfuseObservation(span=span, public=parent.public) + + +@dataclass(frozen=True, slots=True) +class TraceIdHashSampler(Sampler): + """Sample on a SHA-256 of the trace id rather than its low 64 bits. + + litellm trace ids are UUIDs, whose variant bits pin the top of that low word, so + ``TraceIdRatioBased`` drops every trace at rates up to 0.5 and skews above it. + """ + + rate: float + + def should_sample( + self, + parent_context: Context | None, + trace_id: int, + name: str, + kind: SpanKind | None = None, + attributes: Attributes = None, + links: Sequence[Link] | None = None, + trace_state: TraceState | None = None, + ) -> SamplingResult: + digest: Final = sha256(trace_id.to_bytes(16, "big")).digest() + sampled: Final = int.from_bytes(digest[:8], "big") < round(self.rate * 2**64) + parent: Final = otel_trace.get_current_span(parent_context).get_span_context() + return SamplingResult( + Decision.RECORD_AND_SAMPLE if sampled else Decision.DROP, + attributes if sampled else None, + parent.trace_state if parent.is_valid else None, + ) + + def get_description(self) -> str: + return f"TraceIdHashSampler{{{self.rate}}}" + + +def _parse_float(raw: str) -> float | None: + try: + return float(raw) + except ValueError: + return None + + +def _parse_sample_rate(raw: str) -> float | None: + rate: Final = _parse_float(raw) + return rate if rate is not None and 0.0 <= rate <= 1.0 else None + + +def configured_sample_rate() -> float: + """``LANGFUSE_SAMPLE_RATE`` as a fraction, exporting everything when it is unset or unusable.""" + raw: Final = os.environ.get("LANGFUSE_SAMPLE_RATE") + if raw is None: + return 1.0 + parsed: Final = _parse_sample_rate(raw) + if parsed is None: + verbose_logger.warning( + "LANGFUSE_SAMPLE_RATE=%r is not a number between 0.0 and 1.0; ignoring it and exporting every trace", raw + ) + return 1.0 + return parsed + + +def configured_timeout() -> float: + """``LANGFUSE_TIMEOUT`` in seconds for every export and REST call, the v2 SDK's 20 s when unset. + + A value that is not a number raises, as the v2 client did at construction, so a typo is not silently ignored. + """ + return float(os.environ.get("LANGFUSE_TIMEOUT", _DEFAULT_TIMEOUT_SECONDS)) + + +def configured_max_retries() -> int: + """``LANGFUSE_MAX_RETRIES`` as the number of re-sends after a failed export, the v2 SDK's knob and default. + + Capped at ``_MAX_RETRIES``: with the backoff ceiling that is already hours per batch, and the exporter holds + one delay per re-send. + """ + raw: Final = os.environ.get("LANGFUSE_MAX_RETRIES") + if raw is None: + return _DEFAULT_MAX_RETRIES + if not raw.strip().isdigit(): + verbose_logger.warning( + "LANGFUSE_MAX_RETRIES=%r is not a whole number; retrying %d times", raw, _DEFAULT_MAX_RETRIES + ) + return _DEFAULT_MAX_RETRIES + requested: Final = int(raw) + if requested > _MAX_RETRIES: + verbose_logger.warning( + "LANGFUSE_MAX_RETRIES=%d is above the ceiling; retrying %d times", requested, _MAX_RETRIES + ) + return min(requested, _MAX_RETRIES) + + +def configured_release() -> str | None: + """``LANGFUSE_RELEASE``, else the commit variable of the CI or deploy platform, as both SDK generations resolve it.""" + return os.environ.get("LANGFUSE_RELEASE") or next( + (os.environ[name] for name in _COMMON_RELEASE_ENVS if name in os.environ), None + ) + + +def configured_prompt_cache_ttl() -> float: + """``LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS`` in whole seconds as the SDK reads it, its 60 s default when unset + or unusable; ``raise_if_unusable_prompt_cache_ttl`` has already named a value that is not a whole number.""" + raw: Final = os.environ.get(PROMPT_CACHE_TTL_ENV) + if raw is None: + return _DEFAULT_PROMPT_CACHE_TTL_SECONDS + parsed: Final = whole_number(raw) + if parsed is None or parsed < 0: + verbose_logger.warning( + "%s=%r is not a whole number of seconds at or above 0; caching prompts for %.0f s", + PROMPT_CACHE_TTL_ENV, + raw, + _DEFAULT_PROMPT_CACHE_TTL_SECONDS, + ) + return _DEFAULT_PROMPT_CACHE_TTL_SECONDS + return float(parsed) + + +def configured_flush_at() -> int: + """``LANGFUSE_FLUSH_AT`` as the export batch size, the SDK's own knob, with its default when unset or unusable.""" + raw: Final = os.environ.get("LANGFUSE_FLUSH_AT") + if raw is None: + return _DEFAULT_FLUSH_AT + parsed: Final = int(raw) if raw.strip().isdigit() else None + if parsed is None or not 0 < parsed <= _MAX_QUEUE_SIZE: + verbose_logger.warning( + "LANGFUSE_FLUSH_AT=%r is not a whole number between 1 and %d; exporting batches of %d", + raw, + _MAX_QUEUE_SIZE, + _DEFAULT_FLUSH_AT, + ) + return _DEFAULT_FLUSH_AT + return parsed + + +class DiscardingSpanExporter(SpanExporter): + """Accept and drop every span, for mock mode. + + The mock intercepts the httpx client behind the REST API, but observations + travel over OTLP, so without this the "no network calls" contract silently sends + real traces to the configured host. + """ + + def export(self, spans: object) -> SpanExportResult: + return SpanExportResult.SUCCESS + + def shutdown(self) -> None: + return None + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return True + + +_ExportOutcome = Literal["delivered", "retry", "rejected", "too_large"] +_Batch = tuple[ReadableSpan, ...] + + +@dataclass(frozen=True, slots=True) +class _Halving: + """One round of a 413 split: the batches still to send and the results of the ones already settled.""" + + pending: tuple[_Batch, ...] + settled: tuple[SpanExportResult, ...] = () + + +def _smaller(batch: _Batch) -> tuple[_Batch, ...]: + """What to send after a 413: the two halves of a batch, or a single span with its largest field truncated.""" + if len(batch) != 1: + return batch[: len(batch) // 2], batch[len(batch) // 2 :] + (only,) = batch + truncated: Final = _truncated(only) + return () if truncated is None else ((truncated,),) + + +def _in_group(key: str, group: tuple[str, ...]) -> bool: + return any(key == prefix or key.startswith(prefix + ".") for prefix in group) + + +def _group_size(attributes: Mapping[str, AttributeValue], group: tuple[str, ...]) -> int: + return sum( + len(str(value)) for key, value in attributes.items() if _in_group(key, group) and value != _TRUNCATION_MARKER + ) + + +def _marker_key(prefix: str) -> str: + """Langfuse reads input and output as one string but metadata only as flattened keys, so the marker gets one.""" + return f"{prefix}.truncated" if prefix in _METADATA_PREFIXES else prefix + + +def _truncated(span: ReadableSpan) -> ReadableSpan | None: + """The span with its largest remaining input, output or metadata replaced by the marker the v2 consumer wrote + when an event went over ``LANGFUSE_MAX_EVENT_SIZE_BYTES``, or ``None`` once all three are gone.""" + attributes: Final = span.attributes or MappingProxyType({}) + largest: Final = max(_TRUNCATION_GROUPS, key=lambda group: _group_size(attributes, group)) + if _group_size(attributes, largest) == 0: + return None + kept: Final = {key: value for key, value in attributes.items() if not _in_group(key, largest)} + marked: Final = { + _marker_key(prefix): _TRUNCATION_MARKER + for prefix in largest + if any(_in_group(key, (prefix,)) for key in attributes) + } + return ReadableSpan( + name=span.name, + context=span.context, + parent=span.parent, + resource=span.resource, + attributes=MappingProxyType({**kept, **marked}), + events=span.events, + links=span.links, + kind=span.kind, + status=span.status, + start_time=span.start_time, + end_time=span.end_time, + instrumentation_scope=span.instrumentation_scope, + ) + + +def enable_langfuse_debug_logging() -> None: + """What ``Langfuse(debug=True)`` does: a root handler if none exists, and the ``langfuse`` logger at DEBUG.""" + logging.basicConfig(format="%(asctime)s - %(name)s - %(levelname)s - %(message)s") + _langfuse_logger.setLevel(logging.DEBUG) + + +def _retryable_status(status: int) -> bool: + """Any 5xx, a timeout or a rate limit: what the v2 consumer re-sent, plus the 408 the OTLP exporter retries.""" + return status in (408, 429) or 500 <= status <= 599 + + +@dataclass(frozen=True, slots=True) +class LangfuseSpanExporter(SpanExporter): + """OTLP/HTTP protobuf export through litellm's own HTTP handler. + + The handler carries litellm's TLS material (``ssl_verify``, CA bundle, client certificate) exactly + as v2's injected httpx client did. A connect or read failure and a retryable status are re-sent after + each delay, matching the v2 ingestion consumer; ``BatchSpanProcessor`` would otherwise drop the whole + batch on the first exception. A 413 splits the batch in halves until each body fits or a single span + is left; that span is re-sent with its input, output and metadata replaced by the v2 consumer's + truncation marker, largest first, and dropped only when the fully truncated span is still refused. + """ + + handler: HTTPHandler + endpoint: str + headers: Mapping[str, str] + timeout: float + delays: Sequence[float] = (1.0, 2.0, 4.0) + + def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult: + """Halving a batch of n spans settles every span within ``n.bit_length()`` rounds plus one per truncation + step, so the rounds are a fixed fold rather than a recursion.""" + rounds: Final = range(len(spans).bit_length() + 1 + len(_TRUNCATION_GROUPS)) + final: Final = reduce(lambda halving, _: self._round(halving), rounds, _Halving(pending=(tuple(spans),))) + return ( + SpanExportResult.SUCCESS + if all(result is SpanExportResult.SUCCESS for result in final.settled) + else SpanExportResult.FAILURE + ) + + def _round(self, halving: _Halving) -> _Halving: + sent: Final = tuple((batch, self._send_batch(batch)) for batch in halving.pending) + return _Halving( + pending=tuple(part for batch, outcome in sent if outcome == "too_large" for part in _smaller(batch)), + settled=halving.settled + + tuple( + SpanExportResult.SUCCESS if outcome == "delivered" else SpanExportResult.FAILURE + for _, outcome in sent + if outcome != "too_large" + ), + ) + + def _send_batch(self, batch: _Batch) -> _ExportOutcome: + """A 413 on more than one span asks for halves; on a single span it asks for a truncation, and the span is + dropped and reported once nothing is left to truncate.""" + body: Final = _encode(batch) + if body is None: + return "rejected" + outcome: Final = self._send(body) + if outcome != "too_large": + return outcome + match batch: + case (only,) if _truncated(only) is None: + verbose_logger.error( + "Langfuse rejected a single %d byte span export to %s as too large, dropping it", + len(body), + self.endpoint, + ) + return "rejected" + case (_,): + verbose_logger.warning( + "Langfuse rejected a single %d byte span export to %s as too large, resending it with its " + "largest field replaced by %r", + len(body), + self.endpoint, + _TRUNCATION_MARKER, + ) + case _: + verbose_logger.warning( + "Langfuse rejected a %d byte export of %d spans as too large, resending in halves", + len(body), + len(batch), + ) + return "too_large" + + def _send(self, body: bytes) -> _ExportOutcome: + for delay in self.delays: + outcome: _ExportOutcome = self._post(body) + if outcome != "retry": + return outcome + verbose_logger.warning("Langfuse export to %s failed, retrying in %ss", self.endpoint, delay) + sleep(delay) + last: Final = self._post(body) + if last == "retry": + verbose_logger.error("Langfuse export to %s failed after %d retries", self.endpoint, len(self.delays)) + return last + + def _post(self, body: bytes) -> _ExportOutcome: + try: + self.handler.post(self.endpoint, data=body, headers=dict(self.headers), timeout=self.timeout) + except httpx.HTTPStatusError as error: + status: Final = error.response.status_code + if _retryable_status(status): + return "retry" + if status == 413: + return "too_large" + verbose_logger.error( + "Langfuse rejected an export to %s with HTTP %d%s", + self.endpoint, + status, + _SERVER_FLOOR_HINT if status == 404 else "", + ) + return "rejected" + except (httpx.TransportError, litellm.Timeout) as error: + verbose_logger.warning("Langfuse export to %s raised %s", self.endpoint, error) + return "retry" + _langfuse_logger.debug("Exported %d bytes of spans to %s", len(body), self.endpoint) + return "delivered" + + def shutdown(self) -> None: + return None + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return True + + +def _encode(spans: Sequence[ReadableSpan]) -> bytes | None: + """The OTLP body, or ``None`` when nothing survived: a span the encoder rejects is dropped, not the whole batch.""" + try: + return encode_spans(spans).SerializeToString() + except Exception: # noqa: BLE001 # protobuf raises TypeError or ValueError depending on the field + kept: Final = tuple(span for span in spans if _encodes(span)) + verbose_logger.error("Langfuse export dropped %d span(s) the OTLP encoder rejected", len(spans) - len(kept)) + return encode_spans(kept).SerializeToString() if kept else None + + +def _encodes(span: ReadableSpan) -> bool: + try: + encode_spans((span,)) + except Exception: # noqa: BLE001 # same encoder failure modes as above + return False + return True + + +def _build_span_exporter(*, public_key: str, secret_key: str, base_url: str) -> LangfuseSpanExporter: + """Endpoint, headers and export path are the v4 SDK span processor's, so the server treats the spans as SDK + traffic; the 20 s timeout and the retry count are what the v2 consumer used. The ingestion-version header is + the one Langfuse's compatibility matrix asks a v4 producer to send.""" + export_path: Final = os.getenv("LANGFUSE_OTEL_TRACES_EXPORT_PATH") or "/api/public/otel/v1/traces" + encoded_auth: Final = b64encode(f"{public_key}:{secret_key}".encode()).decode("ascii") + return LangfuseSpanExporter( + handler=_get_httpx_client(), + endpoint=f"{base_url.rstrip('/')}/{export_path.lstrip('/')}", + headers=MappingProxyType( + { + "Authorization": "Basic " + encoded_auth, + "Content-Type": "application/x-protobuf", + "x-langfuse-sdk-name": "python", + "x-langfuse-sdk-version": version("langfuse"), + "x-langfuse-public-key": public_key, + _LANGFUSE_INGESTION_VERSION_HEADER: _LANGFUSE_INGESTION_VERSION, + } + ), + timeout=configured_timeout(), + delays=tuple(2.0 ** min(attempt, _MAX_BACKOFF_EXPONENT) for attempt in range(configured_max_retries())), + ) + + +def _resource(*, environment: str | None, release: str | None) -> Resource: + """Only litellm's own attributes: ``Resource.create`` would merge the host's ``OTEL_RESOURCE_ATTRIBUTES``.""" + return Resource( + _present( + ( + (LangfuseOtelSpanAttributes.ENVIRONMENT, environment), + (LangfuseOtelSpanAttributes.RELEASE, release), + ) + ) + ) + + +class _ExportLedger(SpanExporter): + """Counts the batches the exporter gave up on, so a flush can report delivery rather than a drained queue.""" + + def __init__(self, exporter: SpanExporter) -> None: + self.exporter: Final = exporter + self.failed_batches = 0 + + def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult: + result: Final = self.exporter.export(spans) + if result is not SpanExportResult.SUCCESS: + self.failed_batches += 1 + return result + + def shutdown(self) -> None: + self.exporter.shutdown() + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return self.exporter.force_flush(timeout_millis) + + +@dataclass(frozen=True, slots=True) +class LangfuseTracing: + """litellm's own export channel to one Langfuse project: a provider, its tracer and the exporter behind them. + + The channel is litellm's rather than the SDK's so that the process-global OTel provider stays + untouched, historical timestamps and caller ids are honoured, and no SDK internals are needed. + """ + + provider: TracerProvider + tracer: Tracer + ledger: _ExportLedger + + def flush(self, timeout_millis: int = 30_000) -> bool: + """``True`` only when the queue drained in time and every batch it held was accepted by the destination.""" + failed_before: Final = self.ledger.failed_batches + return self.provider.force_flush(timeout_millis) and self.ledger.failed_batches == failed_before + + def shutdown(self) -> None: + self.provider.shutdown() + + +@dataclass(frozen=True, slots=True) +class _TracingKey: + public_key: str + secret_key: str + base_url: str + environment: str | None + release: str | None + sample_rate: float + flush_at: int + flush_interval_millis: int + mock_mode: bool + + +@dataclass(frozen=True, slots=True) +class _Lease: + tracing: LangfuseTracing + holders: int + retire: threading.Timer | None = None + + +_TRACING_LOCK: Final = threading.Lock() +_TRACING: Final[dict[_TracingKey, _Lease]] = {} # mutable-ok: process-wide channel cache, guarded by _TRACING_LOCK + + +def acquire_langfuse_tracing( + *, + public_key: str, + secret_key: str, + base_url: str, + environment: str | None, + release: str | None, + flush_interval: float, + mock_mode: bool, +) -> LangfuseTracing: + """One export channel per credential set, shared by every logger built for it. + + A provider owns a batch export thread, so a channel lives while any logger holds it and is + retired through ``release_langfuse_tracing`` once the last holder lets go. + """ + if parse_langfuse_debug(os.getenv("LANGFUSE_DEBUG")): + enable_langfuse_debug_logging() + key: Final = _TracingKey( + public_key=public_key, + secret_key=secret_key, + base_url=base_url, + environment=environment, + release=release, + sample_rate=configured_sample_rate(), + flush_at=configured_flush_at(), + flush_interval_millis=int(flush_interval * 1000), + mock_mode=mock_mode, + ) + with _TRACING_LOCK: + cached: Final = _TRACING.get(key) + if cached is not None: + if cached.retire is not None: + cached.retire.cancel() + _TRACING[key] = replace(cached, holders=cached.holders + 1, retire=None) + return cached.tracing + created: Final = build_langfuse_tracing( + exporter=DiscardingSpanExporter() + if mock_mode + else _build_span_exporter(public_key=public_key, secret_key=secret_key, base_url=base_url), + environment=environment, + release=release, + sample_rate=key.sample_rate, + flush_at=key.flush_at, + flush_interval_millis=key.flush_interval_millis, + ) + _TRACING[key] = _Lease(tracing=created, holders=1) + return created + + +def release_langfuse_tracing(tracing: LangfuseTracing, *, grace_seconds: float = _CHANNEL_RETIRE_GRACE_SECONDS) -> None: + """Let go of one logger's hold on its channel; a channel nobody holds is retired ``grace_seconds`` later. + + The grace covers a callback that fetched its logger from the cache just before the entry expired, + and a logger rebuilt for the same credentials in the meantime picks the channel back up instead. + """ + with _TRACING_LOCK: + held: Final = next(((key, lease) for key, lease in _TRACING.items() if lease.tracing is tracing), None) + if held is None: + return + key, lease = held + if lease.holders <= 0: + return + if lease.holders > 1: + _TRACING[key] = replace(lease, holders=lease.holders - 1) + return + if grace_seconds > 0: + retire: Final = threading.Timer(grace_seconds, lambda: _retire_unless_reacquired(key, retire)) + retire.name = "langfuse-retire" + retire.daemon = True + _TRACING[key] = _Lease(tracing=tracing, holders=0, retire=retire) + retire.start() + return + del _TRACING[key] + tracing.shutdown() + + +def _retire_unless_reacquired(key: _TracingKey, timer: threading.Timer) -> None: + """Only the timer the lease still points at may retire it; a re-acquire cancels and clears the pending one.""" + with _TRACING_LOCK: + lease: Final = _TRACING.get(key) + if lease is None or lease.retire is not timer: + return + del _TRACING[key] + lease.tracing.shutdown() + + +class _FlushWorker(threading.Thread): + """Daemon, so a channel still blocked at the deadline cannot hold up interpreter exit.""" + + def __init__(self, channel: LangfuseTracing, timeout_millis: int) -> None: + super().__init__(name="langfuse-flush", daemon=True) + self.channel: Final = channel + self.timeout_millis: Final = timeout_millis + self.flushed = False + + def run(self) -> None: + self.flushed = self.channel.flush(self.timeout_millis) + + +def flush_langfuse_tracing(timeout_millis: int = 30_000) -> bool: + """Force-flush every export channel this process acquired, all within one ``timeout_millis`` deadline. + + ``True`` only when every channel flushed in time; a channel still blocked at the deadline is left to + finish in the background rather than pushing the deadline out for the channels after it. + """ + with _TRACING_LOCK: + channels: Final = tuple(lease.tracing for lease in _TRACING.values()) + workers: Final = tuple(_FlushWorker(channel, timeout_millis) for channel in channels) + deadline: Final = monotonic() + timeout_millis / 1000 + for worker in workers: + worker.start() + for worker in workers: + worker.join(max(0.0, deadline - monotonic())) + return all(not worker.is_alive() and worker.flushed for worker in workers) + + +def build_langfuse_tracing( + *, + exporter: SpanExporter, + environment: str | None, + release: str | None, + sample_rate: float, + flush_interval_millis: int, + flush_at: int = _DEFAULT_FLUSH_AT, +) -> LangfuseTracing: + """Wire the provider from litellm's own settings so a host's ``OTEL_*`` variables do not steer it. + + An unset sampler or span limit falls back to ``OTEL_TRACES_SAMPLER`` and + ``OTEL_SPAN_ATTRIBUTE_COUNT_LIMIT`` style variables, which are meant for the + host application's own tracing. ``OTEL_SDK_DISABLED`` still applies, as it does to the SDK. + + The tracer carries the SDK's scope name because Langfuse keys on it: spans from any other + scope are treated as foreign OTel traffic and get their raw attributes echoed into metadata. + """ + if os.environ.get("OTEL_SDK_DISABLED", "").strip().lower() == "true": + verbose_logger.warning("OTEL_SDK_DISABLED=true also disables the langfuse callback's export channel") + provider: Final = TracerProvider( + resource=_resource(environment=environment, release=release), + sampler=ALWAYS_ON if sample_rate >= 1 else TraceIdHashSampler(sample_rate), + id_generator=_RequestedIdGenerator(), + span_limits=_SPAN_LIMITS, + ) + ledger: Final = _ExportLedger(exporter) + provider.add_span_processor( + BatchSpanProcessor( + ledger, + max_queue_size=_MAX_QUEUE_SIZE, + max_export_batch_size=flush_at, + schedule_delay_millis=flush_interval_millis, + ) + ) + return LangfuseTracing(provider=provider, tracer=provider.get_tracer(_TRACER_NAME), ledger=ledger) + + +@dataclass(frozen=True, slots=True) +class _CachedPrompt: + prompt: PromptClient + fetched_at: float + + +_PromptKey = tuple[str, int | None, str | None] + + +def _prompt_client(prompt: Prompt) -> PromptClient: + return ChatPromptClient(prompt) if isinstance(prompt, Prompt_Chat) else TextPromptClient(prompt) + + +@dataclass(frozen=True, slots=True) +class AuthCheckFailure: + reason: str + + +def _auth_check_failure(reason: str) -> AuthCheckFailure: + verbose_logger.warning("Langfuse auth check failed: %s", reason) + return AuthCheckFailure(reason) + + +class _ApiErrorDetail(BaseModel): + """The status and body of an ``ApiError``, whose own ``str`` also dumps every response header.""" + + model_config = ConfigDict(frozen=True, from_attributes=True) + status_code: int | None + body: object + + +def _api_error_reason(error: ApiError) -> str: + detail: Final = _ApiErrorDetail.model_validate(error) + return f"status_code: {detail.status_code}, body: {detail.body}" + + +class LangfusePromptError(Exception): + """An ``ApiError`` without its ``headers``, which the proxy would otherwise forward to its own client.""" + + def __init__(self, error: ApiError) -> None: + detail: Final = _ApiErrorDetail.model_validate(error) + super().__init__(f"status_code: {detail.status_code}, body: {detail.body}") + self.status_code: Final = detail.status_code + self.body: Final = detail.body + + +def _is_server_error(error: ApiError) -> bool: + return error.status_code is not None and error.status_code >= 500 + + +class LangfuseApiClient: + """litellm's handle on one Langfuse project over its REST API: prompts, ``auth_check`` and the project id. + + The SDK's ``Langfuse`` client is deliberately not constructed. It keeps one tracing bundle per + public key and hands it to every ``Langfuse()`` a host application builds for the same key, so + litellm's exporter, host and masking would leak into that application. Observations travel + over ``LangfuseTracing``; nothing here exports spans. + + Prompts are cached for ``LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS`` (60 by default) as the SDK + does. A stale prompt is served at once and refreshed on a background thread, so the request + that finds it stale, and the event loop it runs on, never wait for the REST round trip; a + refresh that fails keeps serving the stale prompt rather than failing the request, again like the SDK. + """ + + def __init__(self, api: LangfuseAPI, *, prompt_cache_ttl_seconds: float) -> None: + self.api: Final = api + self.prompt_cache_ttl_seconds: Final = prompt_cache_ttl_seconds + # mutable-ok: per-client prompt cache, guarded by _lock + self._prompts: Final[dict[_PromptKey, _CachedPrompt]] = {} + # mutable-ok: keys with a refresh in flight, guarded by _lock + self._refreshing: Final[set[_PromptKey]] = set() + self._lock: Final = threading.Lock() + + def auth_check(self) -> AuthCheckFailure | None: + """``None`` when the keys reach a project; otherwise the reason, which is also logged. + + Mirrors the SDK's ``Langfuse.auth_check``: a 200 with no project is a failure too, and a server + error or a transport failure is reported as itself rather than as bad credentials. + """ + try: + projects: Final = self.api.projects.get(request_options=_NO_REST_RETRIES).data + except ApiError as error: + return _auth_check_failure(_api_error_reason(error)) + except Exception as error: # noqa: BLE001 # httpx transport errors or a body the response model rejects + return _auth_check_failure(str(error) or type(error).__name__) + if not projects: + return _auth_check_failure("no project found for the keys provided") + return None + + def project_id(self) -> str | None: + projects: Final = self.api.projects.get(request_options=_NO_REST_RETRIES).data + return projects[0].id if projects else None + + def get_prompt(self, name: str, *, label: str | None = None, version: int | None = None) -> PromptClient: + key: Final[_PromptKey] = (name, version, label) + with self._lock: + cached: Final = self._prompts.get(key) + if cached is None: + return self._fetch(key) + if monotonic() - cached.fetched_at >= self.prompt_cache_ttl_seconds: + self._refresh_in_background(key) + return cached.prompt + + def _fetch(self, key: _PromptKey) -> PromptClient: + fetched: Final = _prompt_client(self._request_prompt(key)) + with self._lock: + self._prompts[key] = _CachedPrompt(prompt=fetched, fetched_at=monotonic()) + return fetched + + def _request_prompt(self, key: _PromptKey) -> Prompt: + """Retried once, at once, after a 5xx or a transport failure: a cold miss runs on the caller's event + loop, so the generated client's sleeping retries stay off.""" + name, version, label = key + request: Final = partial( + self.api.prompts.get, quote(name, safe=""), version=version, label=label, request_options=_NO_REST_RETRIES + ) + try: + return request() + except ApiError as error: + if not _is_server_error(error): + raise LangfusePromptError(error) from None + verbose_logger.debug("Langfuse prompt %r fetch failed (%s), retrying once", name, _api_error_reason(error)) + except httpx.TransportError as error: + verbose_logger.debug("Langfuse prompt %r fetch failed (%s), retrying once", name, error) + try: + return request() + except ApiError as error: + raise LangfusePromptError(error) from None + + def _refresh_in_background(self, key: _PromptKey) -> None: + with self._lock: + if key in self._refreshing: + return + self._refreshing.add(key) + threading.Thread(target=self._refresh, args=(key,), name="langfuse-prompt-refresh", daemon=True).start() + + def _refresh(self, key: _PromptKey) -> None: + try: + self._fetch(key) + except Exception as error: # noqa: BLE001 # a failed refresh keeps the stale prompt in service + verbose_logger.warning("Langfuse prompt %r refresh failed, serving the cached version: %s", key[0], error) + finally: + with self._lock: + self._refreshing.discard(key) + + +def build_langfuse_client( + *, + public_key: str | None, + secret_key: str | None, + base_url: str, + httpx_client: httpx.Client | None, +) -> LangfuseApiClient: + """The REST client for prompt management, ``auth_check`` and the Slack project link. + + Missing keys are passed through as absent credentials: the server answers 401, which + ``auth_check`` reports as a failure rather than raising at construction. + """ + return LangfuseApiClient( + LangfuseAPI( + base_url=base_url, + username=public_key, + password=secret_key, + x_langfuse_sdk_name="python", + x_langfuse_sdk_version=version("langfuse"), + x_langfuse_public_key=public_key, + httpx_client=httpx_client, + timeout=configured_timeout(), + ), + prompt_cache_ttl_seconds=configured_prompt_cache_ttl(), + ) diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index b9f2359e9ea..d4642ae2aad 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -2,7 +2,11 @@ from typing import Final, cast from urllib.parse import urlparse import litellm -from litellm.constants import PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, REPLICATE_MODEL_NAME_WITH_ID_LENGTH +from litellm.constants import ( + NADIR_DEFAULT_API_BASE, + PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, + REPLICATE_MODEL_NAME_WITH_ID_LENGTH, +) from litellm.litellm_core_utils.fallback_generalizations import ( match_routing_generalization, ) @@ -277,6 +281,11 @@ def get_llm_provider( elif endpoint == "https://api.cerebras.ai/v1": custom_llm_provider = "cerebras" dynamic_api_key = get_secret_str("CEREBRAS_API_KEY") + elif endpoint == NADIR_DEFAULT_API_BASE: + custom_llm_provider = "nadir" # rebind-ok: mirrors sibling endpoint branches + dynamic_api_key = ( + get_secret_str("NADIR_API_KEY") if api_base.lower().startswith("https://") else None + ) elif endpoint == "https://inference.baseten.co/v1": custom_llm_provider = "baseten" dynamic_api_key = get_secret_str("BASETEN_API_KEY") @@ -649,6 +658,13 @@ def _get_openai_compatible_provider_info( elif custom_llm_provider == "cerebras": api_base = api_base or get_secret("CEREBRAS_API_BASE") or "https://api.cerebras.ai/v1" dynamic_api_key = api_key or get_secret_str("CEREBRAS_API_KEY") + elif custom_llm_provider == "nadir": + default_nadir_base: Final = get_secret_str("NADIR_API_BASE") or NADIR_DEFAULT_API_BASE + caller_base: Final = api_base + api_base = api_base or default_nadir_base # rebind-ok: mirrors sibling provider branches + trusted_base: Final = caller_base is None or caller_base.rstrip("/") == default_nadir_base.rstrip("/") + env_key: Final = get_secret_str("NADIR_API_KEY") if trusted_base else None + dynamic_api_key = api_key or env_key # rebind-ok: mirrors sibling provider branches elif custom_llm_provider == "baseten": # Use BasetenConfig to determine the appropriate API base URL if api_base is None: diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 680f31a797f..c635cf828eb 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -91,6 +91,8 @@ def get_supported_openai_params( return litellm.nvidiaNimEmbeddingConfig.get_supported_openai_params() elif custom_llm_provider == "cerebras": return litellm.CerebrasConfig().get_supported_openai_params(model=model) + elif custom_llm_provider == "nadir": + return litellm.NadirConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "baseten": return litellm.BasetenConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "xai": diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 8e5af4e5cd6..83ab2bc11a2 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -36,7 +36,7 @@ from litellm._logging import ( ) from litellm._uuid import uuid from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final -from litellm.caching.caching import DualCache, InMemoryCache +from litellm.caching.caching import DualCache from litellm.caching.caching_handler import LLMCachingHandler from litellm.constants import ( DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, @@ -221,6 +221,7 @@ from .initialize_dynamic_callback_params import ( initialize_standard_callback_dynamic_params as _initialize_standard_callback_dynamic_params, ) from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache +from .specialty_caches.service_trace_id_cache import in_memory_trace_id_cache if TYPE_CHECKING: from mcp.types import CallToolResult, EmbeddedResource, ImageContent, TextContent @@ -349,21 +350,6 @@ last_fetched_at_keys: Final = None #### -class ServiceTraceIDCache: - def __init__(self) -> None: - self.cache = InMemoryCache() - - def get_cache(self, litellm_call_id: str, service_name: str) -> str | None: - key_name: Final = f"{service_name}:{litellm_call_id}" - response: Final = self.cache.get_cache(key=key_name) - return response - - def set_cache(self, litellm_call_id: str, service_name: str, trace_id: str) -> None: - key_name: Final = f"{service_name}:{litellm_call_id}" - self.cache.set_cache(key=key_name, value=trace_id) - - -in_memory_trace_id_cache: Final = ServiceTraceIDCache() in_memory_dynamic_logger_cache: Final = DynamicLoggingCache() # Cached lazy import for PrometheusLogger @@ -3979,40 +3965,6 @@ class Logging(LiteLLMLoggingBaseClass): return trace_id - def _get_callback_object(self, service_name: Literal["langfuse"]) -> Any | None: - """ - Return dynamic callback object. - - Meant to solve issue when doing key-based/team-based logging - """ - global langFuseLogger - - if service_name == "langfuse": - if langFuseLogger is None or ( - ( - self.standard_callback_dynamic_params.get("langfuse_public_key") is not None - and self.standard_callback_dynamic_params.get("langfuse_public_key") != langFuseLogger.public_key - ) - or ( - self.standard_callback_dynamic_params.get("langfuse_public_key") is not None - and self.standard_callback_dynamic_params.get("langfuse_public_key") != langFuseLogger.public_key - ) - or ( - self.standard_callback_dynamic_params.get("langfuse_host") is not None - and self.standard_callback_dynamic_params.get("langfuse_host") != langFuseLogger.langfuse_host - ) - ): - return LangFuseLogger( - langfuse_public_key=self.standard_callback_dynamic_params.get("langfuse_public_key"), - langfuse_secret=self.standard_callback_dynamic_params.get("langfuse_secret") - or self.standard_callback_dynamic_params.get("langfuse_secret_key"), - langfuse_host=self.standard_callback_dynamic_params.get("langfuse_host"), - allow_env_credentials=self.standard_callback_dynamic_params.get("langfuse_host") is None, - ) - return langFuseLogger - - return None - def handle_sync_success_callbacks_for_async_calls( self, result: Any, diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 8e5d2cd0a17..14d47a15c6d 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1989,6 +1989,26 @@ def is_encrypted_reasoning_block(block: object) -> bool: return _carries_encrypted_reasoning(_encrypted_reasoning_field(mapping)) +def is_unsignable_thinking_block(block: object) -> bool: + """A thinking block Anthropic cannot accept on input. + + Anthropic verifies the thinking signature cryptographically, so a block whose + signature is null, empty, or missing (e.g. from an open-source reasoning model) + is rejected with a 400 and must be dropped rather than blanked or repaired, and + so is a block whose signature or data carries another provider's encrypted + reasoning. A `redacted_thinking` block Anthropic minted is always kept. + """ + if is_encrypted_reasoning_block(block): + return True + if not isinstance(block, Mapping): + return False + mapping: Final = cast(Mapping[str, object], block) # cast-ok: narrowed by isinstance + if mapping.get("type") != "thinking": + return False + signature: Final = mapping.get("signature") + return not (isinstance(signature, str) and len(signature) > 0) + + def strip_encrypted_reasoning_from_messages(messages: object) -> None: """Drop the bridge-tagged reasoning blocks a routed deployment cannot decrypt from Anthropic-shaped history. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 6fc319c26ae..7b12d1e939f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -7,7 +7,7 @@ import re import xml.etree.ElementTree as ET from collections.abc import Iterator, Mapping, Sequence from enum import Enum -from typing import Any, Final, TypedDict, cast, overload +from typing import Any, Final, TypeAlias, TypedDict, cast, overload from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -17,6 +17,7 @@ import litellm.types.llms from litellm import verbose_logger from litellm._uuid import uuid from litellm.constants import REDACTED_BY_LITELLM +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import anthropic_system_messages from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client from litellm.types.files import get_file_extension_from_mime_type @@ -48,8 +49,8 @@ from litellm.types.utils import GenericImageParsingChunk from .common_utils import ( convert_content_list_to_str, infer_content_type_from_url_and_content, - is_encrypted_reasoning_block, is_non_content_values_set, + is_unsignable_thinking_block, parse_tool_call_arguments, ) from .image_handling import convert_url_to_base64 @@ -2329,37 +2330,25 @@ def sanitize_messages_for_tool_calling( return sanitized_messages -def _is_unsignable_thinking_block(block: object) -> bool: - """A thinking block that Anthropic cannot accept on input. - - Anthropic verifies the thinking signature cryptographically, so a block whose - signature is null, empty, or missing (e.g. from an open-source reasoning model) - is rejected with a 400 and must be dropped rather than blanked or repaired, and - so is a block whose signature or data carries another provider's encrypted - reasoning. A `redacted_thinking` block Anthropic minted is always kept. - """ - if is_encrypted_reasoning_block(block): - return True - if not isinstance(block, dict) or block.get("type") != "thinking": - return False - signature: Final = block.get("signature") - return not (isinstance(signature, str) and len(signature) > 0) - - def _drop_unsignable_thinking_blocks( thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock], ) -> list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock]: - return [block for block in thinking_blocks if not _is_unsignable_thinking_block(block)] + return [block for block in thinking_blocks if not is_unsignable_thinking_block(block)] + + +_AnthropicMessageList: TypeAlias = list[AllAnthropicPassThroughMessageValues] def anthropic_messages_pt( messages: list[AllMessageValues], model: str, llm_provider: str, -) -> list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]: +) -> _AnthropicMessageList: """ format messages for anthropic - 1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately) + 1. Anthropic supports roles like "user" and "assistant" (system prompt sent separately). + Models flagged ``supports_mid_conversation_system`` also accept "system" inside + messages after a user turn; the caller decides placement, this keeps such messages. 2. The first message always needs to be of role "user" 3. Each message must alternate between "user" and "assistant" (this is not addressed as now by litellm) 4. final assistant content cannot end with trailing whitespace (anthropic raises an error otherwise) @@ -2384,7 +2373,7 @@ def anthropic_messages_pt( # add role=tool support to allow function call result/error submission user_message_types: Final = {"user", "tool", "function"} # reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them. - new_messages: Final[list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]] = [] + new_messages: Final[_AnthropicMessageList] = [] # mutable-ok: accumulator behind the mutable return contract if len(messages) == 0: if not litellm.modify_params: @@ -2697,7 +2686,7 @@ def anthropic_messages_pt( if ( m.get("type", "") == "thinking" and len(thinking_block) > 0 - and not _is_unsignable_thinking_block(m) + and not is_unsignable_thinking_block(m) ): # don't pass empty text blocks. anthropic api raises errors. anthropic_message: ChatCompletionThinkingBlock | AnthropicMessagesTextParam = cast( ChatCompletionThinkingBlock, m @@ -2777,6 +2766,11 @@ def anthropic_messages_pt( if assistant_content: new_messages.append({"role": "assistant", "content": assistant_content}) + ## MID-CONVERSATION SYSTEM MESSAGES (placement is the caller's job) ## + while msg_i < len(messages) and messages[msg_i]["role"] == "system": + new_messages.extend(anthropic_system_messages(messages[msg_i])) + msg_i += 1 + if msg_i == init_msg_i: # prevent infinite loops raise litellm.BadRequestError( message=BAD_MESSAGE_ERROR_STR + f"passed in {messages[msg_i]}", diff --git a/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py new file mode 100644 index 00000000000..b5e9afca86b --- /dev/null +++ b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py @@ -0,0 +1,418 @@ +"""Placement policy for ``role: "system"`` messages that appear after the first turn +of an Anthropic-shaped chat completions request. + +Only the leading run of system messages belongs in the top-level ``system`` +parameter. Hoisting a later one there rewrites the cached prefix, so the provider +re-bills the whole conversation at cache-write pricing on every reminder (#36559). + +Models flagged ``supports_mid_conversation_system`` in the cost map accept the role +inside ``messages`` under Anthropic's placement rules: the message must directly +follow a user turn, must be the last entry or be followed by an assistant turn, and +must not sit next to another system message. OpenAI-shaped clients put system +messages anywhere, so this module places each run by its neighbours alone: a run +after a user turn stays with that turn, a run after an assistant turn slides +behind the user turn that immediately follows it, and a run that ends the array +or precedes an assistant turn becomes a user turn in place. Runs that land on the +same slot merge into one system message. No later message can move an earlier +run, so a client that replays the conversation with more turns appended sends a +byte-identical prefix and preserved thinking blocks keep their binding. + +Models without the flag reject the role inside ``messages``. Their system messages +become user turns in place, prefixed with an operator note so the model can tell +the instruction apart from the user's own words. A run caught between a tool call +and its result moves to just after the result so the ``tool_result`` block stays +first in the merged user turn. + +Every transformation here is a pure function of the message sequence: turn N's +output stays a prefix of turn N+1's output, which is what keeps the provider-side +prompt cache readable across turns. Messages are handled in OpenAI format; the +Anthropic wire shape is built later by ``anthropic_messages_pt``. +""" + +from collections.abc import Iterator, Mapping, Sequence +from itertools import chain, groupby +from typing import Final, Literal, TypeAlias + +from litellm.types.llms.anthropic import AnthropicMessagesSystemMessageParam, AnthropicSystemMessageContent +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionCachedContent, + ChatCompletionSystemMessage, + ChatCompletionTextObject, + ChatCompletionUserMessage, +) + +from .common_utils import is_unsignable_thinking_block + +CONVERTED_SYSTEM_NOTE: Final = ( + "Operator note (not from the user): the following was originally a mid-conversation system-role reminder." +) + +_USER_TYPE_ROLES: Final = frozenset({"user", "tool", "function"}) +_TOOL_ROLES: Final = frozenset({"tool", "function"}) +_RENDERED_PART_TYPES: Final = frozenset({"text", "image_url", "document", "file"}) +_RENDERED_ASSISTANT_PART_TYPES: Final = frozenset({"text", "server_tool_use"}) +_THINKING_BLOCK_TYPES: Final = frozenset({"thinking", "redacted_thinking"}) + +_MessageKind: TypeAlias = Literal["system", "tool", "user", "other"] +_TextPart: TypeAlias = tuple[str, ChatCompletionCachedContent | None] + + +def _as_mapping(value: object) -> Mapping[str, object] | None: + return value if isinstance(value, Mapping) else None + + +def parts_of(value: object) -> tuple[object, ...]: + return tuple(value) if isinstance(value, Sequence) and not isinstance(value, str) else () + + +def message_field(message: object, key: str) -> object: + """A message field, whether the message is a dict or a pydantic ``Message``. + + Clients replay assistant turns straight from a response, so a history mixes + plain dicts with ``litellm.Message`` objects; every predicate reads through here. + """ + mapping: Final = _as_mapping(message) + return mapping.get(key) if mapping is not None else getattr(message, key, None) + + +def is_system_message(message: object) -> bool: + return message_field(message, "role") == "system" + + +def _is_user_type(message: object) -> bool: + return message_field(message, "role") in _USER_TYPE_ROLES + + +def _kind(message: object) -> _MessageKind: + role: Final = message_field(message, "role") + if role == "system": + return "system" + if role in _TOOL_ROLES: + return "tool" + if role == "user": + return "user" + return "other" + + +def split_leading_system_run( + messages: Sequence[AllMessageValues], +) -> tuple[tuple[AllMessageValues, ...], tuple[AllMessageValues, ...]]: + """Split ``messages`` into the leading run of system messages and everything after it.""" + leading_count: Final = next( + (index for index, message in enumerate(messages) if not is_system_message(message)), + len(messages), + ) + return tuple(messages[:leading_count]), tuple(messages[leading_count:]) + + +def _cache_control(holder: object) -> ChatCompletionCachedContent | None: + """The client's ``cache_control`` rebuilt in the only shape Anthropic accepts.""" + value: Final = _as_mapping(message_field(holder, "cache_control")) + if value is None or value.get("type") != "ephemeral": + return None + ttl: Final = value.get("ttl") + if ttl == "1h": + one_hour: Final[ChatCompletionCachedContent] = {"type": "ephemeral", "ttl": "1h"} + return one_hour + if ttl == "5m": + five_minutes: Final[ChatCompletionCachedContent] = {"type": "ephemeral", "ttl": "5m"} + return five_minutes + ephemeral: Final[ChatCompletionCachedContent] = {"type": "ephemeral"} + return ephemeral + + +def _text_parts(message: object) -> tuple[_TextPart, ...]: + """``(text, cache_control)`` for each non-empty text part of a system message. + + Anthropic rejects empty text blocks and only accepts text in system content. A + ``cache_control`` on the message itself belongs to the block built from string + content; block-level ``cache_control`` stays with its block. + """ + content: Final = message_field(message, "content") + if isinstance(content, str): + return ((content, _cache_control(message)),) if content else () + return tuple(part for part in map(_text_part, parts_of(content)) if part is not None) + + +def _text_part(part: object) -> _TextPart | None: + if message_field(part, "type") != "text": + return None + text: Final = message_field(part, "text") + return (text, _cache_control(part)) if isinstance(text, str) and text else None + + +def _openai_text_block(part: _TextPart) -> ChatCompletionTextObject: + text, cache_control = part + if cache_control is None: + plain: Final[ChatCompletionTextObject] = {"type": "text", "text": text} + return plain + cached: Final[ChatCompletionTextObject] = {"type": "text", "text": text, "cache_control": cache_control} + return cached + + +def _anthropic_text_block(part: _TextPart) -> AnthropicSystemMessageContent: + text, cache_control = part + if cache_control is None: + plain: Final[AnthropicSystemMessageContent] = {"type": "text", "text": text} + return plain + cached: Final[AnthropicSystemMessageContent] = {"type": "text", "text": text, "cache_control": cache_control} + return cached + + +def anthropic_system_messages(message: object) -> tuple[AnthropicMessagesSystemMessageParam, ...]: + """The Anthropic wire message for a system message, or nothing when it carries no text.""" + blocks: Final = tuple(_anthropic_text_block(part) for part in _text_parts(message)) + if not blocks: + return () + wire: Final[AnthropicMessagesSystemMessageParam] = { + "role": "system", + "content": list(blocks), # mutable-ok: wire payload; cache_control hooks edit content blocks in place + } + return (wire,) + + +def system_message_as_user(message: object) -> ChatCompletionUserMessage: + """A system message re-rolled as a user turn, prefixed with the operator note.""" + note: Final[ChatCompletionTextObject] = {"type": "text", "text": CONVERTED_SYSTEM_NOTE} + content: Final[list[ChatCompletionTextObject]] = [ # mutable-ok: anthropic_messages_pt only recognises list content + note, + *(_openai_text_block(part) for part in _text_parts(message)), + ] + turn: Final[ChatCompletionUserMessage] = {"role": "user", "content": content} + return turn + + +def _merged_system_message(run: Sequence[object]) -> tuple[ChatCompletionSystemMessage, ...]: + parts: Final = tuple(chain.from_iterable(_text_parts(message) for message in run)) + if not parts: + return () + content: Final[list[ChatCompletionTextObject]] = [ # mutable-ok: anthropic_messages_pt only recognises list content + _openai_text_block(part) for part in parts + ] + merged: Final[ChatCompletionSystemMessage] = {"role": "system", "content": content} + return (merged,) + + +def _converted_user_turns(run: Sequence[object]) -> tuple[ChatCompletionUserMessage, ...]: + return tuple(system_message_as_user(message) for message in run if _text_parts(message)) + + +def _runs(messages: Sequence[AllMessageValues]) -> tuple[tuple[_MessageKind, tuple[AllMessageValues, ...]], ...]: + return tuple((kind, tuple(group)) for kind, group in groupby(messages, key=_kind)) + + +def _converted_for_unflagged_model(messages: Sequence[AllMessageValues]) -> tuple[AllMessageValues, ...]: + """Convert every system message to a user turn in place. + + A system run whose follower is a tool message is emitted after that tool run: + ``tool_result`` blocks have to open the merged user turn. + """ + runs: Final = _runs(messages) + + def emit(index: int) -> tuple[AllMessageValues, ...]: + kind, run = runs[index] + follower: Final = runs[index + 1][0] if index + 1 < len(runs) else None + if kind == "system": + return () if follower == "tool" else _converted_user_turns(run) + if kind == "tool" and index > 0 and runs[index - 1][0] == "system": + return (*run, *_converted_user_turns(runs[index - 1][1])) + return run + + return tuple(chain.from_iterable(emit(index) for index in range(len(runs)))) + + +def _user_type_blocks(messages: Sequence[AllMessageValues]) -> tuple[tuple[bool, tuple[int, ...]], ...]: + """Maximal groups of consecutive non-system messages, keyed by whether they are user-type. + + Consecutive user-type messages become one user turn on the wire, so a group is + the unit a system message can validly follow. + """ + indexed: Final = tuple((index, message) for index, message in enumerate(messages) if not is_system_message(message)) + return tuple( + (is_user, tuple(index for index, _ in group)) + for is_user, group in groupby(indexed, key=lambda pair: _is_user_type(pair[1])) + ) + + +def _system_runs(messages: Sequence[AllMessageValues]) -> tuple[tuple[int, ...], ...]: + """Index runs of consecutive system messages.""" + system_indices: Final = tuple(index for index, message in enumerate(messages) if is_system_message(message)) + return tuple( + tuple(index for _, index in group) + for _, group in groupby(enumerate(system_indices), key=lambda pair: pair[1] - pair[0]) + ) + + +def _block_containing(message_index: int, blocks: Sequence[tuple[bool, tuple[int, ...]]]) -> int: + return next(index for index, (_, indices) in enumerate(blocks) if message_index in indices) + + +def _thinking_block_renders(block: object) -> bool: + """A thinking block the converter keeps: one Anthropic can verify, so never bridged encrypted reasoning.""" + return message_field(block, "type") in _THINKING_BLOCK_TYPES and not is_unsignable_thinking_block(block) + + +def _assistant_part_renders(part: object) -> bool: + """A text part always renders: the converter pads empty text with a placeholder.""" + part_type: Final = message_field(part, "type") + if part_type == "thinking": + thinking: Final = message_field(part, "thinking") + return isinstance(thinking, str) and bool(thinking) and _thinking_block_renders(part) + return part_type in _RENDERED_ASSISTANT_PART_TYPES or ( + isinstance(part_type, str) and part_type.endswith("_tool_result") + ) + + +def _separate_thinking_blocks_render(message: object, parts: Sequence[object]) -> bool: + """``thinking_blocks`` reach the wire only when no inline thinking part claims the slot. + + The converter skips the separate blocks as soon as the content list carries a + ``thinking`` or ``redacted_thinking`` part, whether or not that part itself renders. + """ + if any(message_field(part, "type") in _THINKING_BLOCK_TYPES for part in parts): + return False + return any(_thinking_block_renders(block) for block in parts_of(message_field(message, "thinking_blocks"))) + + +def _assistant_renders(message: object) -> bool: + """Whether ``anthropic_messages_pt`` puts a block on the wire for this assistant message. + + String content (the converter pads an empty one with a placeholder), a text part, + a signed thinking part, a server tool part, tool calls, a function call, a kept + thinking block and compaction blocks each render. An assistant message with none + of them, such as ``content: None`` or an empty list, vanishes from the wire. + """ + content: Final = message_field(message, "content") + if isinstance(content, str): + return True + parts: Final = parts_of(content) + return ( + any(_assistant_part_renders(part) for part in parts) + or _separate_thinking_blocks_render(message, parts) + or bool(message_field(message, "tool_calls")) + or bool(message_field(message, "function_call")) + or bool(message_field(message_field(message, "provider_specific_fields"), "compaction_blocks")) + ) + + +def _renders(message: object) -> bool: + """Whether ``anthropic_messages_pt`` puts a block on the wire for this message. + + A tool message always becomes a ``tool_result`` and a user message with string + content always becomes a text block (empty text gets a placeholder). A user list + renders only through parts of a type the converter emits; ``None``, an empty list, + and a list of other parts vanish. Assistant messages follow ``_assistant_renders``. + """ + role: Final = message_field(message, "role") + if role in _TOOL_ROLES: + return True + if role == "assistant": + return _assistant_renders(message) + content: Final = message_field(message, "content") + return isinstance(content, str) or any( + message_field(part, "type") in _RENDERED_PART_TYPES for part in parts_of(content) + ) + + +def _rendered_block( + message_index: int, + messages: Sequence[AllMessageValues], + blocks: Sequence[tuple[bool, tuple[int, ...]]], +) -> int | None: + block_index: Final = _block_containing(message_index, blocks) + _, indices = blocks[block_index] + return block_index if any(_renders(messages[index]) for index in indices) else None + + +def _system_may_follow( + block_index: int, + messages: Sequence[AllMessageValues], + blocks: Sequence[tuple[bool, tuple[int, ...]]], +) -> bool: + """Whether a system message behind this block precedes an assistant turn or ends the array on the wire. + + Blocks alternate between user-type and assistant, so the check is whether the + first later block that puts anything on the wire is an assistant block. + """ + return next( + ( + not is_user + for is_user, indices in blocks[block_index + 1 :] + if any(_renders(messages[index]) for index in indices) + ), + True, + ) + + +def _anchor_block( + run: Sequence[int], + messages: Sequence[AllMessageValues], + blocks: Sequence[tuple[bool, tuple[int, ...]]], +) -> int | None: + """The user-type block a system run must follow, or ``None`` when it converts in place. + + The run never starts at 0: the leading system run was split off before this + policy runs, so the message before a run is always a non-system message. Only + the run's neighbours decide, so a request that replays these messages with more + turns appended places the run identically. A block that puts nothing on the wire + cannot anchor a run: the system message would land first or behind an assistant + turn, so the run converts in place instead. The same happens when the assistant + turn after the anchor puts nothing on the wire and a user turn follows it: the + system message would sit directly before that user turn, which Anthropic rejects. + """ + previous: Final = run[0] - 1 + neighbour: Final = previous if _is_user_type(messages[previous]) else run[-1] + 1 + if neighbour >= len(messages) or not _is_user_type(messages[neighbour]): + return None + block_index: Final = _rendered_block(neighbour, messages, blocks) + if block_index is None or not _system_may_follow(block_index, messages, blocks): + return None + return block_index + + +def _placed_for_flagged_model(messages: Sequence[AllMessageValues]) -> tuple[AllMessageValues, ...]: + """Keep system messages as ``role: "system"`` at a placement Anthropic accepts. + + A run already sitting after a user-type message stays with that user turn. A + run after an assistant turn moves behind the user turn that immediately follows + it. A run that ends the array or is followed by an assistant turn becomes user + turns in place, so replaying the same messages with more turns appended cannot + move it. Runs that share a user turn merge into one system message. + """ + blocks: Final = _user_type_blocks(messages) + anchors: Final = tuple((run, _anchor_block(run, messages, blocks)) for run in _system_runs(messages)) + + def messages_of(run: tuple[int, ...]) -> tuple[AllMessageValues, ...]: + return tuple(messages[index] for index in run) + + def anchored_to(block_index: int) -> tuple[AllMessageValues, ...]: + anchored_runs: Final = tuple(run for run, anchor in anchors if anchor == block_index) + return tuple(chain.from_iterable(map(messages_of, anchored_runs))) + + def converted_after(message_index: int) -> tuple[ChatCompletionUserMessage, ...]: + following_runs: Final = tuple(run for run, anchor in anchors if anchor is None and run[0] == message_index + 1) + return tuple(chain.from_iterable(_converted_user_turns(messages_of(run)) for run in following_runs)) + + def emit(block_index: int) -> Iterator[AllMessageValues]: + is_user, indices = blocks[block_index] + for index in indices: + yield messages[index] + yield from converted_after(index) + if is_user: + yield from _merged_system_message(anchored_to(block_index)) + + return tuple(chain.from_iterable(emit(block_index) for block_index in range(len(blocks)))) + + +def place_mid_conversation_system( + messages: Sequence[AllMessageValues], + *, + supports_mid_conversation_system: bool, +) -> tuple[AllMessageValues, ...]: + """Apply the placement policy to the messages after the leading system run.""" + if not any(is_system_message(message) for message in messages): + return tuple(messages) + if supports_mid_conversation_system: + return _placed_for_flagged_model(messages) + return _converted_for_unflagged_model(messages) diff --git a/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py b/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py index da3ac366bfd..73aca909ce3 100644 --- a/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py +++ b/litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py @@ -1,10 +1,8 @@ """ This is a cache for LangfuseLoggers. -Langfuse Python SDK initializes a thread for each client. - This ensures we do -1. Proper cleanup of Langfuse initialized clients. +1. Release the initialized-client slot a LangfuseLogger holds when it expires. 2. Re-use created langfuse clients. """ @@ -21,45 +19,34 @@ from ...caching import InMemoryCache class LangfuseInMemoryCache(InMemoryCache): """ - Ensures we do proper cleanup of Langfuse initialized clients. + Decrements ``litellm.initialized_langfuse_clients`` when a LangFuseLogger entry expires. - Langfuse Python SDK initializes a thread for each client, we need to call Langfuse.shutdown() to properly cleanup. - - This ensures we do proper cleanup of Langfuse initialized clients. + The counter is a soft budget: loggers built concurrently for one credential set before the + first lands in the cache each take a slot, and only the cached one gives it back on expiry. + The logger's ``stop()`` below hands its shared export channel back + (https://github.com/BerriAI/litellm/issues/11169). """ def _remove_key(self, key: str) -> None: - """ - Override _remove_key in InMemoryCache to ensure we do proper cleanup of Langfuse initialized clients. - - LangfuseLoggers consume threads when initalized, this shuts them down when they are expired - - Relevant Issue: https://github.com/BerriAI/litellm/issues/11169 - """ from litellm.integrations.langfuse.langfuse import LangFuseLogger - if isinstance(self.cache_dict[key], LangFuseLogger): - _created_langfuse_logger: Final[LangFuseLogger] = self.cache_dict[key] - ######################################################### - # Clean up Langfuse initialized clients - ######################################################### + evicted: Final = self.cache_dict.pop(key, None) + self.ttl_dict.pop(key, None) + if evicted is None: + return + + if isinstance(evicted, LangFuseLogger): litellm.initialized_langfuse_clients -= 1 - _created_langfuse_logger.Langfuse.flush() - _created_langfuse_logger.Langfuse.shutdown() # Loggers with a periodic flush task (e.g. NewRelicMetricsLogger) expose # stop() so eviction actually ends the task instead of leaking it. - _evicted_stop: Final = getattr(self.cache_dict[key], "stop", None) - if callable(_evicted_stop): - try: - _evicted_stop() - except Exception: # noqa: BLE001 # a failing stop() must not block eviction - verbose_logger.debug("DynamicLoggingCache: stop() raised during eviction", exc_info=True) - - ######################################################### - # Call parent class to remove key from cache - ######################################################### - return super()._remove_key(key) + _evicted_stop: Final = getattr(evicted, "stop", None) + if not callable(_evicted_stop): + return + try: + _evicted_stop() + except Exception: # noqa: BLE001 # a failing stop() must not block eviction + verbose_logger.debug("DynamicLoggingCache: stop() raised during eviction", exc_info=True) class DynamicLoggingCache: diff --git a/litellm/litellm_core_utils/specialty_caches/service_trace_id_cache.py b/litellm/litellm_core_utils/specialty_caches/service_trace_id_cache.py new file mode 100644 index 00000000000..f1f60d3e7b8 --- /dev/null +++ b/litellm/litellm_core_utils/specialty_caches/service_trace_id_cache.py @@ -0,0 +1,20 @@ +from typing import Final + +from ...caching import InMemoryCache + + +class ServiceTraceIDCache: + def __init__(self) -> None: + self.cache = InMemoryCache() + + def get_cache(self, litellm_call_id: str, service_name: str) -> str | None: + key_name: Final = f"{service_name}:{litellm_call_id}" + response: Final = self.cache.get_cache(key=key_name) + return response + + def set_cache(self, litellm_call_id: str, service_name: str, trace_id: str) -> None: + key_name: Final = f"{service_name}:{litellm_call_id}" + self.cache.set_cache(key=key_name, value=trace_id) + + +in_memory_trace_id_cache: Final = ServiceTraceIDCache() diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index b7c2ce3c568..3bffee48d6a 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -31,13 +31,17 @@ from litellm.litellm_core_utils.prompt_templates.image_handling import ( async_inline_remote_media, inline_remote_image_urls, ) +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import ( + place_mid_conversation_system, + split_leading_system_run, +) from litellm.llms.base_llm.base_utils import type_to_response_format_param from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.types.llms.anthropic import ( ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_BETA_HEADER_VALUES, ANTHROPIC_HOSTED_TOOLS, - AllAnthropicMessageValues, + AllAnthropicPassThroughMessageValues, AllAnthropicToolsValues, AnthropicCodeExecutionTool, AnthropicComputerTool, @@ -87,6 +91,7 @@ from litellm.utils import ( get_max_tokens, has_tool_call_blocks, last_assistant_with_tool_calls_has_no_thinking_blocks, + supports_mid_conversation_system, supports_reasoning, token_counter, ) @@ -1743,10 +1748,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def add_code_execution_tool( self, - messages: list[AllAnthropicMessageValues], + messages: list[AllAnthropicPassThroughMessageValues], tools: list[AllAnthropicToolsValues | dict], ) -> list[AllAnthropicToolsValues | dict]: - """if 'container_upload' in messages, add code_execution tool""" + """if 'container_upload' in messages, add code_execution tool + + Takes the pass-through union because the translator emits ``role: "system"`` + in ``messages`` for models that accept it; only ``content`` is read here.""" add_code_execution_tool = False for message in messages: message_content = message.get("content", None) @@ -1966,16 +1974,27 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if _name_reverse_map and isinstance(litellm_params, dict): litellm_params[ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY] = _name_reverse_map - # Separate system prompt from rest of message - anthropic_system_message_list: Final = self.translate_system_message(messages=messages) + # Only the leading system run becomes the top-level system prompt. A later + # system message stays in the conversation: hoisting it rewrites the cached + # prefix and re-bills the whole history at cache-write pricing (#36559). + leading_system_run, later_messages = split_leading_system_run(messages) + anthropic_system_message_list: Final = self.translate_system_message( + messages=list(leading_system_run) # mutable-ok: translate_system_message pops from the list it is given + ) # Handling anthropic API Prompt Caching if len(anthropic_system_message_list) > 0: optional_params["system"] = anthropic_system_message_list + conversation: Final = place_mid_conversation_system( + later_messages, + supports_mid_conversation_system=supports_mid_conversation_system( + model=model, custom_llm_provider=self.custom_llm_provider + ), + ) # Format rest of message according to anthropic guidelines try: anthropic_messages = anthropic_messages_pt( model=model, - messages=messages, + messages=list(conversation), # mutable-ok: anthropic_messages_pt rewrites entries in place llm_provider=self._resolved_provider, ) except Exception as e: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py b/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py index ddefec6bac9..f588133812c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py @@ -2,9 +2,7 @@ from collections.abc import Mapping, Sequence from itertools import groupby from typing import Final -CONVERTED_SYSTEM_NOTE: Final = ( - "Operator note (not from the user): the following was originally a mid-conversation system-role reminder." -) +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import CONVERTED_SYSTEM_NOTE def as_system_content_blocks(value: object) -> list[object]: diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 497020c2836..2a2f3052b2a 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -7,7 +7,8 @@ import json import re import time import types -from collections.abc import Mapping +from collections.abc import Mapping, Sequence +from itertools import chain from typing import TYPE_CHECKING, Final, Literal, cast, overload import httpx @@ -34,6 +35,12 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( _bedrock_tools_pt, make_valid_bedrock_tool_name, ) +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import ( + CONVERTED_SYSTEM_NOTE, + is_system_message, + message_field, + parts_of, +) from litellm.llms.anthropic.chat.transformation import ( DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING, @@ -55,9 +62,11 @@ from litellm.types.llms.openai import ( ChatCompletionAnnotation, ChatCompletionAssistantMessage, ChatCompletionAssistantToolCall, + ChatCompletionCachedContent, ChatCompletionRedactedThinkingBlock, ChatCompletionResponseMessage, ChatCompletionSystemMessage, + ChatCompletionTextObject, ChatCompletionThinkingBlock, ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, @@ -1343,30 +1352,157 @@ class AmazonConverseConfig(BaseConfig): cache_point["ttl"] = ttl return cache_point + @staticmethod + def _assistant_has_tool_calls(message: object) -> bool: + return message_field(message, "role") == "assistant" and bool(message_field(message, "tool_calls")) + + @staticmethod + def _opens_with_tool_result(message: object) -> bool: + """Whether the message starts a tool-result turn on Converse. + + ``_bedrock_converse_messages_pt`` builds ``toolResult`` blocks from ``tool`` + messages only, so a ``function`` message never opens one.""" + role: Final = message_field(message, "role") + if role == "tool": + return True + if role != "user": + return False + first_part: Final = next(iter(parts_of(message_field(message, "content"))), None) + return message_field(first_part, "type") == "tool_result" + + def _system_run_before(self, messages: Sequence[AllMessageValues], index: int) -> Sequence[AllMessageValues]: + start: Final = next( + (j + 1 for j in range(index - 1, -1, -1) if not is_system_message(messages[j])), + 0, + ) + return messages[start:index] + + def _system_run_end(self, messages: Sequence[AllMessageValues], index: int) -> int: + return next( + (j for j in range(index, len(messages)) if not is_system_message(messages[j])), + len(messages), + ) + + def _reordered_around_tool_results( + self, messages: Sequence[AllMessageValues], index: int + ) -> tuple[AllMessageValues, ...]: + """Move a system run wedged between an assistant tool-call turn and its + tool-result turn(s) to after the tool results. + + A converted system entry becomes a user turn, and a user turn between + a tool call and its result would split them. Everything else stays in + place so the cached prefix stays byte-identical.""" + message: Final = messages[index] + if self._opens_with_tool_result(message): + if index + 1 < len(messages) and self._opens_with_tool_result(messages[index + 1]): + return (message,) + tool_run_start: Final = next( + (j + 1 for j in range(index, -1, -1) if not self._opens_with_tool_result(messages[j])), + 0, + ) + run: Final = self._system_run_before(messages, tool_run_start) + prev_idx: Final = tool_run_start - len(run) - 1 + if run and prev_idx >= 0 and self._assistant_has_tool_calls(messages[prev_idx]): + return (message, *run) + return (message,) + if not is_system_message(message): + return (message,) + run_start: Final = next( + (j + 1 for j in range(index - 1, -1, -1) if not is_system_message(messages[j])), + 0, + ) + run_end: Final = self._system_run_end(messages, index) + follower: Final = messages[run_end] if run_end < len(messages) else None + if ( + follower is not None + and self._opens_with_tool_result(follower) + and run_start > 0 + and self._assistant_has_tool_calls(messages[run_start - 1]) + ): + return () + return (message,) + + def _system_role_message_as_user(self, message: ChatCompletionSystemMessage) -> ChatCompletionUserMessage | None: + """Convert a mid-conversation system entry to a user turn, in place. + + The Converse API only accepts user/assistant roles in ``messages``, + so keeping the role is not an option. Hoisting it to the top-level + ``system`` block would mutate the system prefix and collapse implicit + prompt caching; converting in place keeps everything before the entry + byte-identical. An entry that carries no text becomes ``None``.""" + text_blocks: Final = self._converted_text_blocks(message) + if not text_blocks: + return None + note: Final = ChatCompletionTextObject(type="text", text=CONVERTED_SYSTEM_NOTE) + body: Final = [ # mutable-ok: _bedrock_converse_messages_pt narrows content with isinstance(list) + note, + *text_blocks, + ] + return ChatCompletionUserMessage(role="user", content=body) + + def _converted_or_kept(self, message: AllMessageValues) -> AllMessageValues | None: + if not is_system_message(message): + return message + return self._system_role_message_as_user( + cast(ChatCompletionSystemMessage, message) # cast-ok: the role is checked on the line above + ) + + def _converted_text_blocks(self, message: ChatCompletionSystemMessage) -> tuple[ChatCompletionTextObject, ...]: + content: Final = message["content"] + if isinstance(content, str): + return (self._converted_text_block(content, message.get("cache_control")),) if content else () + parts: Final[Sequence[object]] = content or () + return tuple( + self._converted_text_block(part["text"], part.get("cache_control")) + for part in map(self._text_part, parts) + if part is not None + ) + + @staticmethod + def _text_part(part: object) -> ChatCompletionTextObject | None: + if not isinstance(part, dict) or part.get("type") != "text" or not part.get("text"): + return None + return cast(ChatCompletionTextObject, part) # cast-ok: the shape is checked on the line above + + @staticmethod + def _converted_text_block(text: str, cache_control: ChatCompletionCachedContent | None) -> ChatCompletionTextObject: + if cache_control is None: + return ChatCompletionTextObject(type="text", text=text) + return ChatCompletionTextObject(type="text", text=text, cache_control=cache_control) + def _transform_system_message( self, messages: list[AllMessageValues], model: str | None = None ) -> tuple[list[AllMessageValues], list[SystemContentBlock]]: - system_prompt_indices: Final = [] + leading_count: Final = next( + (i for i, m in enumerate(messages) if not is_system_message(m)), + len(messages), + ) + hoisted: Final = messages[:leading_count] + remaining: Final = messages[leading_count:] system_content_blocks: Final[list[SystemContentBlock]] = [] - for idx, message in enumerate(messages): - if message["role"] == "system": - system_prompt_indices.append(idx) - if isinstance(message["content"], str) and message["content"]: - system_content_blocks.append(SystemContentBlock(text=message["content"])) - cache_block = self.get_cache_point_block(message, block_type="system", model=model) - if cache_block: - system_content_blocks.append(cache_block) - elif isinstance(message["content"], list): - for m in message["content"]: - if m.get("type") == "text" and m.get("text"): - system_content_blocks.append(SystemContentBlock(text=m["text"])) - cache_block = self.get_cache_point_block(m, block_type="system", model=model) - if cache_block: - system_content_blocks.append(cache_block) - if len(system_prompt_indices) > 0: - for idx in reversed(system_prompt_indices): - messages.pop(idx) - return messages, system_content_blocks + for message in hoisted: + if message["role"] != "system": + continue + if isinstance(message["content"], str) and message["content"]: + system_content_blocks.append(SystemContentBlock(text=message["content"])) + cache_block = self.get_cache_point_block(message, block_type="system", model=model) + if cache_block: + system_content_blocks.append(cache_block) + elif isinstance(message["content"], list): + for m in message["content"]: + if m.get("type") == "text" and m.get("text"): + system_content_blocks.append(SystemContentBlock(text=m["text"])) + cache_block = self.get_cache_point_block(m, block_type="system", model=model) + if cache_block: + system_content_blocks.append(cache_block) + reordered: Final = tuple( + chain.from_iterable( + self._reordered_around_tool_results(remaining, index) for index in range(len(remaining)) + ) + ) + converted: Final = tuple(self._converted_or_kept(message) for message in reordered) + kept: Final = [message for message in converted if message is not None] # mutable-ok: converse pt takes a list + return kept, system_content_blocks def _transform_inference_params(self, inference_params: dict) -> InferenceConfig: if "top_k" in inference_params: diff --git a/litellm/llms/fal_ai/cost_calculator.py b/litellm/llms/fal_ai/cost_calculator.py index 519a11de13c..aa3d7942174 100644 --- a/litellm/llms/fal_ai/cost_calculator.py +++ b/litellm/llms/fal_ai/cost_calculator.py @@ -142,14 +142,35 @@ def _resolution_key(resolution: object) -> str | None: return str(resolution) +def _resolution_cost_per_image(entry: Mapping[str, object] | None, resolution: object) -> float | None: + resolution_key: Final = _resolution_key(resolution) + if entry is None or resolution_key is None: + return None + cost: Final = entry.get(f"output_cost_per_image_{resolution_key}") + return float(cost) if isinstance(cost, (int, float)) else None + + +def _requested_image_count(request_body: Mapping[str, object]) -> int: + num_images: Final = request_body.get("num_images") + return num_images if type(num_images) is int and num_images > 0 else 1 + + +def _passthrough_cost_per_image(entry: Mapping[str, object], request_body: Mapping[str, object]) -> float | None: + resolution_cost: Final = _resolution_cost_per_image(entry, request_body.get("resolution")) + if resolution_cost is not None: + return resolution_cost + cost: Final = entry.get("output_cost_per_image") + return float(cost) if isinstance(cost, (int, float)) else None + + def fal_ai_passthrough_cost(model: str, request_body: Mapping[str, object]) -> float | None: entry: Final = _entry(f"{litellm.LlmProviders.FAL_AI.value}/{model}") if entry is None: return None - resolution: Final = _resolution_key(request_body.get("resolution")) - keyed_cost: Final = entry.get(f"output_cost_per_image_{resolution}") if resolution is not None else None - cost: Final = keyed_cost if isinstance(keyed_cost, (int, float)) else entry.get("output_cost_per_image") - return float(cost) if isinstance(cost, (int, float)) else None + cost_per_image: Final = _passthrough_cost_per_image(entry, request_body) + if cost_per_image is None: + return None + return cost_per_image * _requested_image_count(request_body) def cost_calculator( @@ -172,6 +193,11 @@ def cost_calculator( if deployment_cost_per_image is not None: return deployment_cost_per_image * len(images) params: Final[Mapping[str, object]] = optional_params or MappingProxyType({}) + resolution_cost_per_image: Final = _resolution_cost_per_image( + _entry(f"{litellm.LlmProviders.FAL_AI.value}/{normalized_model}"), params.get("resolution") + ) + if resolution_cost_per_image is not None: + return resolution_cost_per_image * len(images) keyed_costs: Final = tuple( _keyed_cost_per_image( model=normalized_model, diff --git a/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py b/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py index bb104a0793f..35df7a72fe0 100644 --- a/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py +++ b/litellm/llms/fal_ai/image_generation/nano_banana_transformation.py @@ -8,12 +8,17 @@ from .transformation import FalAIBaseConfig class FalAINanoBananaConfig(FalAIBaseConfig): """ - Configuration for Fal AI's Nano Banana / Gemini 2.5 Flash Image models. + Configuration for Fal AI's Nano Banana family (Gemini Flash / Pro Image models). - Serves the imagen4 deprecation migration path. The same underlying model is - exposed under two endpoints that share an identical schema: + Serves the imagen4 deprecation migration path. Every endpoint shares the same + request schema, so one config covers all of them: - fal-ai/nano-banana - fal-ai/gemini-25-flash-image + - fal-ai/nano-banana-2 + - fal-ai/nano-banana-pro + + Provider-specific params such as ``resolution`` ("0.5K", "1K", "2K", "4K") are + forwarded as-is and drive the per-resolution price in the cost map. Documentation: https://fal.ai/models/fal-ai/nano-banana """ diff --git a/litellm/llms/nadir/chat/transformation.py b/litellm/llms/nadir/chat/transformation.py new file mode 100644 index 00000000000..306df1208b9 --- /dev/null +++ b/litellm/llms/nadir/chat/transformation.py @@ -0,0 +1,68 @@ +import math +from typing import Final + +import httpx + +from litellm.litellm_core_utils.core_helpers import set_response_cost_in_hidden_params +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import ModelResponse + +_SUPPORTED_OPENAI_PARAMS: Final = ( + "extra_headers", + "frequency_penalty", + "max_retries", + "max_tokens", + "presence_penalty", + "response_format", + "stream", + "temperature", + "top_p", +) + + +def _reported_cost_usd(raw_response: httpx.Response) -> float | None: + try: + cost: Final = raw_response.json()["nadir_metadata"]["cost"]["total_cost_usd"] + except (ValueError, KeyError, TypeError): + return None + if isinstance(cost, bool) or not isinstance(cost, (int, float)): + return None + if not math.isfinite(cost) or cost < 0: + return None + return float(cost) + + +class NadirConfig(OpenAIGPTConfig): + def get_supported_openai_params(self, model: str) -> list: # mutable-ok: return type fixed by the base interface + return list(_SUPPORTED_OPENAI_PARAMS) # mutable-ok: the base interface returns a list + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: object, + request_data: dict, # mutable-ok: signature fixed by the base interface + messages: list[AllMessageValues], # mutable-ok: signature fixed by the base interface + optional_params: dict, # mutable-ok: signature fixed by the base interface + litellm_params: dict, # mutable-ok: signature fixed by the base interface + encoding: object, + api_key: str | None = None, + json_mode: bool | None = None, + ) -> ModelResponse: + transformed: Final = super().transform_response( + model=model, + raw_response=raw_response, + model_response=model_response, + logging_obj=logging_obj, + request_data=request_data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + api_key=api_key, + json_mode=json_mode, + ) + set_response_cost_in_hidden_params(transformed, _reported_cost_usd(raw_response)) + return transformed diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 14aebcaabaf..6d050d5a856 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -25,6 +25,21 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages +VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset( + { + "audio", + "max_retries", + "modalities", + "prediction", + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + } +) + class VertexAILyriaModelInfo(TypedDict): vertex_ai_audio_api: ReadOnly[Literal["lyria_predict", "lyria_interactions"]] @@ -370,6 +385,27 @@ def get_vertex_base_model_name(model: str) -> str: return model +def vertex_model_garden_model_id_in_json_body(model: str) -> bool: + """ + Vertex catalog / publisher models are addressed as publisher/model (e.g. + xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. + + Deployed Model Garden endpoints are typically a single segment (often numeric) + and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. + """ + return "/" in model + + +def is_vertex_self_deployed_openai_compatible_endpoint(model: str) -> bool: + local_model: Final = model.removeprefix("vertex_ai/") + route: Final = get_vertex_ai_model_route(local_model) + if route == VertexAIModelRoute.GEMMA: + return True + return route == VertexAIModelRoute.MODEL_GARDEN and not vertex_model_garden_model_id_in_json_body( + get_vertex_base_model_name(local_model) + ) + + def get_vertex_ai_fine_tuned_endpoint_id(model: str) -> str | None: """ Fine-tuned Gemini deployments are addressed by a numeric endpoint id, diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 0bcf16ee06f..aa2cc8575f9 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -1,4 +1,5 @@ -from collections.abc import Mapping +from collections.abc import Iterable, Mapping, Sequence +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Protocol import httpx @@ -57,23 +58,57 @@ class VertexSearchSnippet(TypedDict, total=False): htmlSnippet: ReadOnly[str] +class VertexSearchExtractiveContent(TypedDict, total=False): + """One ``extractive_answers`` or ``extractive_segments`` entry (opt-in via ``extractiveContentSpec``).""" + + content: ReadOnly[str] + pageNumber: ReadOnly[str] + + class VertexSearchDerivedStructData(TypedDict, total=False): - """The ``derivedStructData`` blob Discovery Engine attaches to each search hit.""" + """The ``derivedStructData`` blob Discovery Engine attaches to each document hit.""" title: ReadOnly[str] link: ReadOnly[str] displayLink: ReadOnly[str] formattedUrl: ReadOnly[str] snippets: ReadOnly[list[VertexSearchSnippet]] + extractive_answers: ReadOnly[list[VertexSearchExtractiveContent]] + extractive_segments: ReadOnly[list[VertexSearchExtractiveContent]] class VertexSearchDocument(TypedDict, total=False): + id: ReadOnly[str] + structData: ReadOnly[Mapping[str, object]] derivedStructData: ReadOnly[VertexSearchDerivedStructData] +class VertexSearchChunkDocumentMetadata(TypedDict, total=False): + uri: ReadOnly[str] + title: ReadOnly[str] + structData: ReadOnly[Mapping[str, object]] + + +class VertexSearchChunkPageSpan(TypedDict, total=False): + pageStart: ReadOnly[int] + pageEnd: ReadOnly[int] + + +class VertexSearchChunk(TypedDict, total=False): + """A hit when ``searchResultMode`` is ``CHUNKS``; such hits carry no ``document`` and no top-level ``id``.""" + + id: ReadOnly[str] + name: ReadOnly[str] + content: ReadOnly[str] + documentMetadata: ReadOnly[VertexSearchChunkDocumentMetadata] + pageSpan: ReadOnly[VertexSearchChunkPageSpan] + relevanceScore: ReadOnly[float] + + class VertexSearchHit(TypedDict, total=False): id: ReadOnly[str] document: ReadOnly[VertexSearchDocument] + chunk: ReadOnly[VertexSearchChunk] class VertexSearchApiResponse(TypedDict, total=False): @@ -98,6 +133,97 @@ def _vertex_search_payload(response: _VertexSearchApiSource) -> VertexSearchApiR return response.json() +_UNKNOWN_DOCUMENT: Final = "Unknown Document" +_EMPTY_DOCUMENT: Final[VertexSearchDocument] = {} +_EMPTY_DERIVED_STRUCT_DATA: Final[VertexSearchDerivedStructData] = {} +_EMPTY_CHUNK_DOCUMENT_METADATA: Final[VertexSearchChunkDocumentMetadata] = {} + + +def _joined_content(entries: Sequence[VertexSearchExtractiveContent]) -> str: + return "\n\n".join(content for entry in entries if (content := entry.get("content"))) + + +def _snippet_text(snippets: Sequence[VertexSearchSnippet]) -> str: + return " ".join(snippet.get("snippet", snippet.get("htmlSnippet", "")) for snippet in snippets) + + +def _document_text(derived: VertexSearchDerivedStructData) -> str: + candidates: Final = ( + _joined_content(derived.get("extractive_segments", ())), + _joined_content(derived.get("extractive_answers", ())), + _snippet_text(derived.get("snippets", ())), + derived.get("title", ""), + ) + return next((text for text in candidates if text), "") + + +def _document_id_from_chunk_name(name: str) -> str: + return name.partition("/documents/")[2].partition("/")[0] + + +def _non_empty_attributes(pairs: Iterable[tuple[str, object]]) -> Mapping[str, object]: + return MappingProxyType({key: value for key, value in pairs if value}) + + +def _chunk_result(chunk: VertexSearchChunk, positional_score: float) -> VectorStoreSearchResult: + metadata: Final = chunk.get("documentMetadata", _EMPTY_CHUNK_DOCUMENT_METADATA) + uri: Final = metadata.get("uri", "") + title: Final = metadata.get("title", "") + document_id: Final = _document_id_from_chunk_name(chunk.get("name", "")) + return VectorStoreSearchResult( + score=chunk.get("relevanceScore", positional_score), + content=[VectorStoreResultContent(text=chunk.get("content", ""), type="text")], + file_id=uri or document_id, + filename=title or _UNKNOWN_DOCUMENT, + attributes={ + "document_id": document_id, + **_non_empty_attributes( + ( + ("chunk_id", chunk.get("id", "")), + ("link", uri), + ("title", title), + ("structData", metadata.get("structData")), + ("pageSpan", chunk.get("pageSpan")), + ) + ), + }, + ) + + +def _document_result(hit: VertexSearchHit, score: float) -> VectorStoreSearchResult: + document: Final = hit.get("document", _EMPTY_DOCUMENT) + derived: Final = document.get("derivedStructData", _EMPTY_DERIVED_STRUCT_DATA) + link: Final = derived.get("link", "") + title: Final = derived.get("title", "") + document_id: Final = hit.get("id", "") + return VectorStoreSearchResult( + score=score, + content=[VectorStoreResultContent(text=_document_text(derived), type="text")], + file_id=link or document_id, + filename=title or _UNKNOWN_DOCUMENT, + attributes={ + "document_id": document_id, + **_non_empty_attributes( + ( + ("link", link), + ("title", title), + ("displayLink", derived.get("displayLink", "")), + ("formattedUrl", derived.get("formattedUrl", "")), + ("structData", document.get("structData")), + ) + ), + }, + ) + + +def _search_result(hit: VertexSearchHit, position: int) -> VectorStoreSearchResult: + score: Final = 1.0 / (position + 1) + chunk: Final = hit.get("chunk") + if chunk is not None: + return _chunk_result(chunk, score) + return _document_result(hit, score) + + class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): """ Configuration for Vertex AI Search API Vector Store @@ -285,98 +411,19 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj ) -> VectorStoreSearchResponse: """ - Transform Vertex AI Search API response to standard vector store search response + Transform a Discovery Engine ``:search`` response into the standard vector store search response. - Handles the format from Discovery Engine Search API which returns: - { - "results": [ - { - "id": "...", - "document": { - "derivedStructData": { - "title": "...", - "link": "...", - "snippets": [...] - } - } - } - ] - } + Document hits (``results[].document``) take their text from ``derivedStructData`` in a fixed order: + ``extractive_segments``, then ``extractive_answers``, then ``snippets``, then ``title``; ``structData`` + and the link metadata land in ``attributes``. Chunk hits (``results[].chunk``, returned when the + caller sets ``contentSearchSpec.searchResultMode`` to ``CHUNKS`` via ``extra_body``) take their text + from ``chunk.content`` and their file id and name from ``chunk.documentMetadata``. """ try: response_json: Final = _vertex_search_payload(response) - - # Extract results from Vertex AI Search API response - results: Final = response_json.get("results", []) - - # Transform results to standard format - search_results: Final[list[VectorStoreSearchResult]] = [] - for result in results: - document: VertexSearchDocument = result.get("document", {}) - derived_data: VertexSearchDerivedStructData = document.get("derivedStructData", {}) - - # Extract text content from snippets - snippets = derived_data.get("snippets", []) - text_content = "" - - if snippets: - # Combine all snippets into one text - text_parts = [snippet.get("snippet", snippet.get("htmlSnippet", "")) for snippet in snippets] - text_content = " ".join(text_parts) - - # If no snippets, use title as fallback - if not text_content: - text_content = derived_data.get("title", "") - - content = [ - VectorStoreResultContent( - text=text_content, - type="text", - ) - ] - - # Extract file/document information - document_link = derived_data.get("link", "") - document_title = derived_data.get("title", "") - document_id = result.get("id", "") - - # Use link as file_id if available, otherwise use document ID - file_id = document_link if document_link else document_id - filename = document_title if document_title else "Unknown Document" - - # Build attributes with available metadata - attributes = { - "document_id": document_id, - } - - if document_link: - attributes["link"] = document_link - if document_title: - attributes["title"] = document_title - - # Add display link if available - display_link = derived_data.get("displayLink", "") - if display_link: - attributes["displayLink"] = display_link - - # Add formatted URL if available - formatted_url = derived_data.get("formattedUrl", "") - if formatted_url: - attributes["formattedUrl"] = formatted_url - - # Note: Search API doesn't provide explicit scores in the response - # You can use the position/rank as an implicit score - score = 1.0 / (float(search_results.__len__() + 1)) # Decreasing score based on position - - result_obj = VectorStoreSearchResult( - score=score, - content=content, - file_id=file_id, - filename=filename, - attributes=attributes, - ) - search_results.append(result_obj) - + search_results: Final = [ + _search_result(hit, position) for position, hit in enumerate(response_json.get("results", ())) + ] query_view: Final[_SearchQueryView] = {"query": litellm_logging_obj.model_call_details.get("query", "")} return VectorStoreSearchResponse( object="vector_store.search_results.page", diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index f2d2c0896d2..ca0bcb74906 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -18,7 +18,11 @@ from litellm.types.utils import ( Usage, ) -from ...common_utils import VertexAIError +from ...common_utils import ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS, + VertexAIError, + is_vertex_self_deployed_openai_compatible_endpoint, +) if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -66,13 +70,15 @@ class VertexAILlama3Config(OpenAIGPTConfig): and v is not None } - def get_supported_openai_params(self, model: str): - supported_params: Final = super().get_supported_openai_params(model=model) - try: - supported_params.remove("max_retries") - except KeyError: - pass - return supported_params + def get_supported_openai_params(self, model: str) -> list[str]: + unsupported_params: Final = ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + if is_vertex_self_deployed_openai_compatible_endpoint(model) + else frozenset({"max_retries"}) + ) + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param for param in super().get_supported_openai_params(model=model) if param not in unsupported_params + ] def map_openai_params( self, diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index 2ae8b4cd188..ea97f0a0a9a 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -12,6 +12,7 @@ from collections.abc import Callable from typing import TYPE_CHECKING, Any, Final, cast import httpx +from pydantic import ValidationError from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -20,7 +21,9 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, ) from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.llms.vertex_ai.common_utils import VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.vertex_ai_gemma import VertexGemmaContainerError from litellm.types.utils import ModelResponse if TYPE_CHECKING: @@ -29,6 +32,13 @@ if TYPE_CHECKING: from litellm.llms.base_llm.base_model_iterator import MockResponseIterator +def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None: + try: + return VertexGemmaContainerError.model_validate(predictions) + except ValidationError: + return None + + class VertexGemmaConfig(OpenAIGPTConfig): """ Configuration and transformation class for Vertex AI Gemma models @@ -40,6 +50,13 @@ class VertexGemmaConfig(OpenAIGPTConfig): def __init__(self) -> None: super().__init__() + def get_supported_openai_params(self, model: str) -> list[str]: + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param + for param in super().get_supported_openai_params(model=model) + if param not in VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + ] + def should_fake_stream( self, model: str | None, @@ -123,7 +140,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): Unwrap the Vertex Gemma predictions format to OpenAI format. Vertex Gemma wraps the OpenAI-compatible response in a 'predictions' field. - This method extracts it so the parent class can process it normally. + This method extracts it so the parent class can process it normally. A serving + container can also answer with its own OpenAI-shaped error object inside that + field, still under HTTP 200, which is raised with its own status and message. """ if "predictions" not in response_json: raise BaseLLMException( @@ -131,7 +150,11 @@ class VertexGemmaConfig(OpenAIGPTConfig): message="Invalid response format: missing 'predictions' field", ) - return response_json["predictions"] + predictions: Final = response_json["predictions"] + container_error: Final = parse_vertex_gemma_container_error(predictions) + if container_error is None: + return predictions + raise BaseLLMException(status_code=container_error.code, message=container_error.message) @staticmethod def _sync_post( diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index f5c9ac623a1..84907f01685 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -24,21 +24,14 @@ import httpx from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.utils import ModelResponse -from ..common_utils import VertexAIError, get_vertex_base_model_name +from ..common_utils import ( + VertexAIError, + get_vertex_base_model_name, + vertex_model_garden_model_id_in_json_body, +) from ..vertex_llm_base import VertexBase -def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: - """ - Vertex catalog / publisher models are addressed as publisher/model (e.g. - xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. - - Deployed Model Garden endpoints are typically a single segment (often numeric) - and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. - """ - return "/" in model - - def create_vertex_url( vertex_location: str, vertex_project: str, @@ -48,7 +41,7 @@ def create_vertex_url( ) -> str: """Return the api base for vertex model garden (without /chat/completions).""" base_url: Final = get_vertex_base_url(vertex_location) - if _vertex_model_garden_model_id_in_json_body(model): + if vertex_model_garden_model_id_in_json_body(model): return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi" return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}" @@ -124,7 +117,7 @@ class VertexAIModelGardenModels(VertexBase): ) # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. - if not _vertex_model_garden_model_id_in_json_body(model): + if not vertex_model_garden_model_id_in_json_body(model): model = "" return openai_like_chat_completions.completion( model=model, diff --git a/litellm/main.py b/litellm/main.py index 98bb5126a90..12854db15d0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -64,6 +64,7 @@ from litellm.constants import ( AZURE_OPENAI_AUDIO_PROVIDERS, DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT, + NADIR_DEFAULT_API_BASE, OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS, ) from litellm.exceptions import LiteLLMUnknownProvider @@ -126,6 +127,7 @@ from litellm.types.completion import ( _CompletionDispatchContext, _CompletionDispatchResult, ) +from litellm.types.litellm_params import RetryStrategy from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -3493,6 +3495,33 @@ def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatch return response +def _complete_nadir(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + api_base: Final = ctx.api_base or litellm.api_base or get_secret_str("NADIR_API_BASE") or NADIR_DEFAULT_API_BASE + api_key: Final = ctx.api_key + + response: Final = base_llm_http_handler.completion( + model=ctx.model, + stream=ctx.stream, + messages=ctx.messages, + acompletion=ctx.acompletion, + api_base=api_base, + model_response=ctx.model_response, + optional_params=ctx.optional_params, + litellm_params=ctx.litellm_params, + shared_session=ctx.shared_session, + custom_llm_provider="nadir", + timeout=ctx.timeout, + headers=ctx.headers or litellm.headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=ctx.logging, + client=ctx.client, + ) + ctx.logging.post_call(input=ctx.messages, api_key=api_key, original_response=response) + + return response + + def _complete_vercel_ai_gateway( ctx: _CompletionDispatchContext, ) -> _CompletionDispatchResult: @@ -5922,6 +5951,8 @@ def completion( response = _complete_datarobot(_dispatch_ctx) elif custom_llm_provider == "openrouter": response = _complete_openrouter(_dispatch_ctx) + elif custom_llm_provider == "nadir": + response = _complete_nadir(_dispatch_ctx) # rebind-ok: mirrors sibling provider branches elif custom_llm_provider == "vercel_ai_gateway": response = _complete_vercel_ai_gateway(_dispatch_ctx) elif custom_llm_provider == "palm": @@ -6026,9 +6057,7 @@ def completion_with_retries(*args, **kwargs): # reset retries in .completion() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6054,7 +6083,7 @@ async def acompletion_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( @@ -6082,9 +6111,7 @@ def responses_with_retries(*args, **kwargs): # reset retries in .responses() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", responses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6111,7 +6138,7 @@ async def aresponses_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", aresponses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index a0c791a136c..13a030e7ebe 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -3,6 +3,8 @@ from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Itera from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable +from litellm.exceptions import BadRequestError +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.llms.anthropic.experimental_pass_through.messages import handler as main from litellm.rust_bridge.catalog import Delivery, Route, RouteContext from litellm.rust_bridge.dispatch import PublicDispatch, call_hook @@ -71,10 +73,17 @@ def _public_request( ) +def _resolved_provider(request: LiteLLMMessagesRequest) -> str | None: + try: + return get_llm_provider(request.model, request.custom_llm_provider)[1] + except BadRequestError: + return request.custom_llm_provider + + def _context(request: LiteLLMMessagesRequest) -> RouteContext: return RouteContext( Route.MESSAGES, - provider=request.custom_llm_provider, + provider=_resolved_provider(request), model=request.model, delivery=Delivery.STREAMING if request.stream else Delivery.COMPLETED, ) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index aa79cadae2a..5a207dc4c02 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3301,12 +3301,14 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "azure/ada": { + "deprecation_date": "2028-02-09", "input_cost_per_token": 1e-07, "litellm_provider": "azure", "max_input_tokens": 8191, "max_tokens": 8191, "mode": "embedding", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule" }, "azure/codex-mini": { "cache_read_input_token_cost": 3.75e-07, @@ -3340,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3347,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5151,6 +5155,7 @@ "supports_tool_choice": true }, "azure/gpt-35-turbo-16k": { + "deprecation_date": "2025-04-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 16385, @@ -5158,9 +5163,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-35-turbo-16k-0613": { + "deprecation_date": "2025-04-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 16385, @@ -5168,6 +5175,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5188,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5195,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5211,6 +5221,7 @@ "supports_tool_choice": true }, "azure/gpt-4-0613": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5218,6 +5229,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5234,6 +5246,7 @@ "supports_tool_choice": true }, "azure/gpt-4-32k": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 6e-05, "litellm_provider": "azure", "max_input_tokens": 32768, @@ -5241,9 +5254,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.00012, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-4-32k-0613": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 6e-05, "litellm_provider": "azure", "max_input_tokens": 32768, @@ -5251,6 +5266,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.00012, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-4-turbo": { @@ -5511,6 +5527,7 @@ }, "azure/gpt-4.5-preview": { "cache_read_input_token_cost": 3.75e-05, + "deprecation_date": "2025-07-14", "input_cost_per_token": 7.5e-05, "input_cost_per_token_batches": 3.75e-05, "litellm_provider": "azure", @@ -5520,6 +5537,7 @@ "mode": "chat", "output_cost_per_token": 0.00015, "output_cost_per_token_batches": 7.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -15797,6 +15815,7 @@ "supports_tool_choice": true }, "computer-use-preview": { + "deprecation_date": "2026-07-23", "input_cost_per_token": 3e-06, "litellm_provider": "openai", "max_input_tokens": 8192, @@ -22838,6 +22857,37 @@ "/v1/images/generations" ] }, + "fal_ai/fal-ai/nano-banana-2": { + "litellm_provider": "fal_ai", + "metadata": { + "comment": "priced by the request's resolution field (0.5K, 1K default, 2K, 4K); the web search and high thinking surcharges are not modeled" + }, + "mode": "image_generation", + "output_cost_per_image": 0.08, + "output_cost_per_image_0.5K": 0.06, + "output_cost_per_image_1K": 0.08, + "output_cost_per_image_2K": 0.12, + "output_cost_per_image_4K": 0.16, + "source": "https://fal.ai/models/fal-ai/nano-banana-2", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, + "fal_ai/fal-ai/nano-banana-pro": { + "litellm_provider": "fal_ai", + "metadata": { + "comment": "priced by the request's resolution field (1K default, 2K, 4K); the web search surcharge is not modeled" + }, + "mode": "image_generation", + "output_cost_per_image": 0.15, + "output_cost_per_image_1K": 0.15, + "output_cost_per_image_2K": 0.15, + "output_cost_per_image_4K": 0.3, + "source": "https://fal.ai/models/fal-ai/nano-banana-pro", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "fal_ai/openai/gpt-image-2": { "litellm_provider": "fal_ai", "metadata": { @@ -28748,6 +28798,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 1e-05, "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -29823,6 +29874,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "rpm": 10000, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_modalities": [ @@ -32113,6 +32165,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-1.5": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -32133,6 +32186,7 @@ "supports_pdf_input": true }, "gpt-image-1.5-2025-12-16": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -32153,6 +32207,7 @@ "supports_pdf_input": true }, "gpt-image-2": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.25e-07, "input_cost_per_token": 5e-06, @@ -32171,12 +32226,14 @@ "supports_pdf_input": true }, "gpt-image-2-2026-04-21": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", "input_cost_per_image_token": 8e-06, "output_cost_per_image_token": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32185,6 +32242,7 @@ "supports_pdf_input": true }, "gpt-image-2.5-flare": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32200,6 +32258,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-flare-2026-09-08": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32215,6 +32274,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32230,6 +32290,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst-2026-09-08": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32657,6 +32718,7 @@ }, "gpt-5-codex": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 272000, @@ -32740,6 +32802,7 @@ }, "gpt-5.1-codex-mini": { "cache_read_input_token_cost": 2.5e-08, + "deprecation_date": "2026-07-23", "input_cost_per_token": 2.5e-07, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32770,6 +32833,7 @@ }, "gpt-5.1-codex-max": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32801,6 +32865,7 @@ }, "gpt-5.1-codex": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32832,6 +32897,7 @@ }, "gpt-5.1-chat-latest": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -32968,6 +33034,7 @@ }, "gpt-5.2-codex": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 272000, @@ -32999,6 +33066,7 @@ }, "gpt-5.2-chat-latest": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-08-10", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -34822,6 +34890,7 @@ }, "gpt-5.3-chat-latest": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-08-10", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -35054,6 +35123,7 @@ "supports_minimal_reasoning_effort": true }, "gpt-image-1": { + "cache_read_input_image_token_cost": 2.5e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-10-23", @@ -35071,6 +35141,7 @@ ] }, "gpt-image-1-mini": { + "cache_read_input_image_token_cost": 2.5e-07, "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_batches": 1e-07, "deprecation_date": "2026-12-01", @@ -35090,6 +35161,7 @@ "gpt-realtime": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 3.2e-05, @@ -35125,6 +35197,7 @@ "gpt-realtime-1.5": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35159,6 +35232,7 @@ "gpt-realtime-2": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35193,6 +35267,7 @@ "gpt-realtime-2.1": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35229,6 +35304,7 @@ "gpt-realtime-2.1-mini": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, @@ -35265,6 +35341,7 @@ "gpt-realtime-mini": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1e-05, @@ -35300,6 +35377,7 @@ "gpt-realtime-2025-08-28": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 3.2e-05, @@ -55976,6 +56054,7 @@ "gpt-realtime-mini-2025-12-15": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, @@ -56068,6 +56147,7 @@ ] }, "chatgpt-image-latest": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -56451,6 +56531,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -56519,6 +56600,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 1e-05, "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -68835,6 +68917,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/NousResearch/Nous-Hermes-2-Mixtral-8x7B-DPO": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68844,6 +68927,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/QwQ-32B": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68853,6 +68937,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2-72B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 9e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68862,6 +68947,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2-VL-72B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68871,6 +68957,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-72B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68880,6 +68967,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-Coder-32B-Instruct": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 16384, @@ -68889,6 +68977,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-VL-72B-Instruct": { + "deprecation_date": "2026-01-05", "input_cost_per_token": 1.95e-06, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68898,6 +68987,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "deprecation_date": "2026-06-04", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68907,6 +68997,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-Next-FP8": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68916,6 +69007,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68925,6 +69017,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68934,6 +69027,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-32B-Instruct": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68943,6 +69037,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-8B-Instruct": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68953,6 +69048,7 @@ }, "together_ai/Qwen/Qwen3.5-397B-A17B": { "cache_read_input_token_cost": 3.5e-07, + "deprecation_date": "2026-06-29", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68962,6 +69058,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Llama-70B": { + "deprecation_date": "2025-12-23", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68971,6 +69068,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68980,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-14B": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 1.6e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68989,6 +69088,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-V3.1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68998,6 +69098,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/deepseek-coder-33b-instruct": { + "deprecation_date": "2024-08-22", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 16384, @@ -69007,6 +69108,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-2-27b-it": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69016,6 +69118,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69025,6 +69128,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Llama-3-8b-chat-hf": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69034,6 +69138,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Llama-4-Scout-17B-16E-Instruct": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -69043,6 +69148,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3-70B-Instruct-Turbo": { + "deprecation_date": "2025-12-23", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69052,6 +69158,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3-8B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69061,6 +69168,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69070,6 +69178,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69079,6 +69188,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-7B-Instruct-v0.1": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69088,6 +69198,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69097,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69107,6 +69219,7 @@ }, "together_ai/moonshotai/Kimi-K2.6": { "cache_read_input_token_cost": 2e-07, + "deprecation_date": "2026-08-19", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69117,6 +69230,7 @@ }, "together_ai/moonshotai/Kimi-K2.7-Code": { "cache_read_input_token_cost": 1.9e-07, + "deprecation_date": "2026-08-27", "input_cost_per_token": 9.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69126,6 +69240,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/nvidia/Llama-3.1-Nemotron-70B-Instruct-HF": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69145,6 +69260,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69154,6 +69270,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-4.5-Air-FP8": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69163,6 +69280,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-4.7": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 202752, @@ -69172,6 +69290,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5": { + "deprecation_date": "2026-06-22", "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", "max_input_tokens": 202752, @@ -69182,6 +69301,7 @@ }, "together_ai/zai-org/GLM-5.1": { "cache_read_input_token_cost": 2.6e-07, + "deprecation_date": "2026-07-10", "input_cost_per_token": 1.4e-06, "litellm_provider": "together_ai", "max_input_tokens": 202752, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index cdf52e6dc8d..a93ffaeac9f 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -13,6 +13,7 @@ from typing_extensions import assert_never import litellm from litellm._logging import verbose_logger +from litellm.constants import MCP_ALL_TOOLS_WILDCARD from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_passthrough_resource_metadata_url, get_passthrough_www_authenticate, @@ -2136,7 +2137,11 @@ class MCPRequestHandler: via_toolsets: Sequence[str] | None, ) -> Sequence[str] | None: """Union of one level's direct tool grants and its toolset-granted tools on one server, - ``None`` when neither source restricts (allow-all from this level).""" + ``None`` when neither source restricts (allow-all from this level). A direct grant + containing ``MCP_ALL_TOOLS_WILDCARD`` makes the level unrestricted, so it returns + ``None`` whatever the toolsets name.""" + if direct is not None and MCP_ALL_TOOLS_WILDCARD in direct: + return None if direct is None and via_toolsets is None: return None return tuple({*(direct or ()), *(via_toolsets or ())}) @@ -2251,11 +2256,7 @@ class MCPRequestHandler: else None ) - key_tools: Final = ( - list(set(key_direct_tools or []) | set(key_toolset_tools or [])) - if key_direct_tools is not None or key_toolset_tools is not None - else None - ) + key_tools: Final = _as_list(MCPRequestHandler._union_tool_grants(key_direct_tools, key_toolset_tools)) team_direct_tools: Final = ( global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id) if team_obj_perm diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index 1879e285789..a88d400282c 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -5,6 +5,7 @@ from datetime import datetime from types import MappingProxyType from typing import Final, Protocol +from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -26,6 +27,7 @@ class OperationContext: raw_headers: Mapping[str, str] | None = field(default=None, repr=False) client_ip: str | None = None mcp_proxy_mode: bool = False + wire_compat: WireCompat = WireCompat.LEGACY def __post_init__(self) -> None: object.__setattr__(self, "_caller", copy_caller(self._caller)) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 312dcb27d89..24cae976174 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -27,7 +27,8 @@ from collections.abc import ( from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache -from itertools import chain +from itertools import chain, groupby +from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -43,6 +44,7 @@ from mcp.types import ( CallToolResult, GetPromptRequestParams, GetPromptResult, + InputRequiredResult, Prompt, ResourceTemplate, ) @@ -132,6 +134,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ServerSpec, TokenExchangeConfig, ) +from litellm.proxy._experimental.mcp_server.result_conversion import ( + WireCompat, + complete_call_tool_result, + handler_outcome, + to_gateway_tool, +) from litellm.proxy._experimental.mcp_server.sampling_handler import ( MCP_SAMPLING_AVAILABLE, ) @@ -5360,16 +5368,9 @@ class MCPServerManager: prefix: Final = get_server_prefix(server) for tool in tools: - tool_copy = tool.model_copy(deep=True) - - original_name = tool_copy.name + original_name = tool.name prefixed_name = add_server_prefix_to_name(original_name, prefix) - - name_to_use = prefixed_name if add_prefix else original_name - - # Preserve all tool fields including metadata/_meta by avoiding mutation - tool_copy.name = name_to_use - prefixed_tools.append(tool_copy) + prefixed_tools.append(to_gateway_tool(tool, prefixed_name if add_prefix else original_name)) # Register every known prefix form (alias, server_name, server_id, # short ID) so call_tool can resolve regardless of which form a @@ -5546,6 +5547,7 @@ class MCPServerManager: server: MCPServer, tool_name: str, arguments: _ToolArguments, + wire_compat: WireCompat = WireCompat.LEGACY, ) -> CallToolResult: """ Call an OpenAPI tool handler directly. @@ -5585,14 +5587,7 @@ class MCPServerManager: # Call the tool handler with the arguments # The handler is an async function that makes the HTTP request handler_result: Final = await tool.handler(**arguments) - - # Convert the handler result (string response) to CallToolResult format - result: Final = CallToolResult( - content=[TextContent(type="text", text=str(handler_result))], - is_error=False, - ) - - return result + return complete_call_tool_result(handler_outcome(handler_result), wire_compat) except MCPUpstreamAuthError: # The caller must re-authenticate upstream, so this keeps its type all the way to the @@ -5819,7 +5814,8 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, - ) -> CallToolResult: + allow_input_required: bool = False, + ) -> CallToolResult | InputRequiredResult: """Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry. The exchanged token is baked into the client at build time, so the retry invalidates the @@ -5829,7 +5825,10 @@ class MCPServerManager: """ try: return await client.call_tool( - call_tool_params, host_progress_callback=host_progress_callback, raise_on_error=True + call_tool_params, + host_progress_callback=host_progress_callback, + raise_on_error=True, + allow_input_required=allow_input_required, ) except Exception as exc: if _extract_upstream_auth_failure(exc) is None: @@ -5847,7 +5846,11 @@ class MCPServerManager: raw_headers=raw_headers, client_ip=client_ip, ) - return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback) + return await retry_client.call_tool( + call_tool_params, + host_progress_callback=host_progress_callback, + allow_input_required=allow_input_required, + ) async def _call_regular_mcp_tool( self, @@ -5864,7 +5867,8 @@ class MCPServerManager: hook_extra_headers: dict[str, str] | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, client_ip: str | None = None, - ) -> CallToolResult: + allow_input_required: bool = False, + ) -> CallToolResult | InputRequiredResult: """ Call a regular MCP tool using the MCP client. @@ -6035,6 +6039,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + allow_input_required=allow_input_required, ) tool_call_coro = _obo_call_tool_limited() @@ -6048,7 +6053,11 @@ class MCPServerManager: async def _call_tool_via_client(client, params): async with self._limit_outbound_concurrency(mcp_server): if not relays_upstream_auth: - return await client.call_tool(params, host_progress_callback=host_progress_callback) + return await client.call_tool( + params, + host_progress_callback=host_progress_callback, + allow_input_required=allow_input_required, + ) # The client-forwarded modes carry the caller's own upstream token, so an upstream # 401 (expired/invalid token) is the caller's to resolve: relay it as # MCPUpstreamAuthError so single-server REST callers turn it into a 401 + @@ -6060,7 +6069,10 @@ class MCPServerManager: # the same isError degradation the default path produces. try: return await client.call_tool( - params, host_progress_callback=host_progress_callback, raise_on_error=True + params, + host_progress_callback=host_progress_callback, + raise_on_error=True, + allow_input_required=allow_input_required, ) except Exception as e: auth_info: Final = _extract_upstream_auth_failure(e) @@ -6113,7 +6125,7 @@ class MCPServerManager: result: Final = mcp_responses[result_index] self._remember_upstream_initialize_instructions(mcp_server, client) - return cast(CallToolResult, result) + return cast("CallToolResult | InputRequiredResult", result) def _resolve_mcp_server_for_tool_call( self, @@ -6317,7 +6329,8 @@ class MCPServerManager: litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, - ) -> CallToolResult: + wire_compat: WireCompat = WireCompat.LEGACY, + ) -> CallToolResult | InputRequiredResult: """ Call a tool with the given name and arguments @@ -6426,7 +6439,7 @@ class MCPServerManager: resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) try: async with self._limit_outbound_concurrency(mcp_server): - return await self._call_openapi_tool_handler(mcp_server, name, arguments) + return await self._call_openapi_tool_handler(mcp_server, name, arguments, wire_compat) finally: _request_auth_header.reset(auth_token) _request_extra_headers.reset(extra_token) @@ -6448,6 +6461,7 @@ class MCPServerManager: host_progress_callback=host_progress_callback, hook_extra_headers=hook_result.get("extra_headers"), user_api_key_auth=user_api_key_auth, + allow_input_required=wire_compat is WireCompat.MODERN, ) return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) @@ -6812,9 +6826,11 @@ class MCPServerManager: """ Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so every key is a concrete server_id where possible. Tool lists from - keys that point at the same server are unioned, matching the - "duplicate names grant access to all matches" semantics of - ``expand_permission_list``. + keys that point at the same server are unioned and deduplicated + first-seen, matching the "duplicate names grant access to all + matches" semantics of ``expand_permission_list``; the + ``MCP_ALL_TOOLS_WILDCARD`` entry is preserved as an ordinary list + entry for the caller to interpret. Required so name-based keys don't silently drop their tool restrictions when the lookup uses the resolved server_id. Unresolved @@ -6823,11 +6839,15 @@ class MCPServerManager: """ if not tool_permissions: return {} - result: Final[dict[str, list[str]]] = {} - for key, tools in tool_permissions.items(): - for server_id in self.expand_permission_list([key]): - result.setdefault(server_id, []).extend(tools or []) - return result + expanded: Final = tuple( + (server_id, tuple(tools or ())) + for key, tools in tool_permissions.items() + for server_id in self.expand_permission_list([key]) + ) + return { + server_id: list(dict.fromkeys(tool for _, tools in group for tool in tools)) + for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) + } def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: """ diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 1247ff1ac28..5b23695d06d 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -52,6 +52,7 @@ from litellm.llms.custom_httpx.http_handler import ( header_value, httpxSpecialProvider, ) +from litellm.proxy._experimental.mcp_server.tool_outcome import JsonResult, TextResult, parse_http_body from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -497,7 +498,7 @@ def create_tool_function( path_params, query_params, body_params = extract_parameters(operation) original_method: Final = method.lower() - async def tool_function(**kwargs: object) -> str: + async def tool_function(**kwargs: object) -> TextResult | JsonResult: """ Dynamically generated tool function. @@ -531,7 +532,7 @@ def create_tool_function( # Sanitize and encode path parameter to prevent traversal attacks safe_value = _sanitize_path_parameter_value(param_value, param_name) except ValueError as exc: - return "Invalid path parameter: " + str(exc) + return TextResult("Invalid path parameter: " + str(exc)) # Replace {param_name} or {{param_name}} in URL url = url.replace("{" + param_name + "}", safe_value) url = url.replace("{{" + param_name + "}}", safe_value) @@ -580,7 +581,7 @@ def create_tool_function( elif original_method == "patch": response = await client.patch(url, params=params, json=json_body, headers=effective_headers) else: - return f"Unsupported HTTP method: {original_method}" + return TextResult(f"Unsupported HTTP method: {original_method}") except MaskedHTTPStatusError as e: _raise_for_upstream_failure(e.response, upstream, relays_upstream_auth) raise @@ -588,7 +589,7 @@ def create_tool_function( _request_upstream_url.reset(url_token) _raise_for_upstream_failure(response, upstream, relays_upstream_auth) - return response.text + return parse_http_body(response.text) return tool_function diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index dcab43bdc76..ebd26e4bf87 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -17,6 +17,7 @@ from mcp.types import ( GetPromptRequest, GetPromptRequestParams, GetPromptResult, + InputRequiredResult, ListPromptsRequest, ListPromptsResult, ListResourcesRequest, @@ -85,6 +86,12 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_extra_headers, _request_resolved_auth_headers, ) +from litellm.proxy._experimental.mcp_server.result_conversion import ( + WireCompat, + complete_call_tool_result, + handler_outcome, + to_call_tool_result, +) from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -1804,8 +1811,9 @@ async def execute_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + wire_compat: WireCompat = WireCompat.LEGACY, **kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract -) -> CallToolResult: +) -> CallToolResult | InputRequiredResult: context: Final = prepare_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -1813,6 +1821,7 @@ async def execute_mcp_tool( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=client_ip, + wire_compat=wire_compat, ) operation: Final = AuthorizedToolCall( name=name, @@ -1839,8 +1848,9 @@ async def _execute_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + wire_compat: WireCompat = WireCompat.LEGACY, **kwargs: Any, -) -> CallToolResult: +) -> CallToolResult | InputRequiredResult: """ Execute MCP tool. @@ -2088,7 +2098,7 @@ async def _execute_mcp_tool( _extra_token: Final = _request_extra_headers.set(forwarded_headers) _resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) try: - response = await _handle_local_mcp_tool(name, arguments) + response = await _handle_local_mcp_tool(name, arguments, wire_compat) finally: _request_auth_header.reset(_auth_token) _request_extra_headers.reset(_extra_token) @@ -2112,6 +2122,7 @@ async def _execute_mcp_tool( litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, host_progress_callback=host_progress_callback, + wire_compat=wire_compat, ) # Fall back to local tool registry with original name (legacy support) @@ -2169,10 +2180,13 @@ async def _execute_mcp_tool( if "arguments" in hook_result: arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args - response = await _handle_local_mcp_tool(original_tool_name, arguments) + response = await _handle_local_mcp_tool(original_tool_name, arguments, wire_compat) + converted: Final = to_call_tool_result(response, wire_compat) + if isinstance(converted, InputRequiredResult): + return converted return await _run_post_mcp_call_guardrails( - result=response, + result=converted, litellm_logging_obj=litellm_logging_obj, user_api_key_auth=user_api_key_auth, request_data=kwargs, @@ -2206,6 +2220,13 @@ async def _run_post_mcp_call_guardrails( ) +def suppress_completed_success_logging(logging_obj: LiteLLMLoggingObj) -> None: + """An interim ``InputRequiredResult`` is not a completed call, so the ``@client`` wrapper + on ``call_mcp_tool`` must not run the success handlers for it when the coroutine returns.""" + logging_obj.has_run_logging(event_type="sync_success") + logging_obj.has_run_logging(event_type="async_success") + + async def _fire_mcp_tool_call_logging( logging_obj: LiteLLMLoggingObj, result: CallToolResult, @@ -2322,10 +2343,14 @@ async def call_mcp_tool( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + wire_compat: WireCompat = WireCompat.LEGACY, **kwargs: Any, -) -> CallToolResult: +) -> CallToolResult | InputRequiredResult: """ Call a specific tool with the provided arguments (handles prefixed tool names). + + A modern ``InputRequiredResult`` is an interim answer, so it is returned as is and skips the + completed-call logging below. """ start_time: Final = datetime.now() litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) @@ -2376,12 +2401,17 @@ async def call_mcp_tool( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=client_ip, + wire_compat=wire_compat, **kwargs, ) except Exception as e: await fire_mcp_tool_call_failure_logging(litellm_logging_obj, e, start_time, user_api_key_auth, kwargs) raise + if isinstance(response, InputRequiredResult): + if litellm_logging_obj: + suppress_completed_success_logging(litellm_logging_obj) + return response if litellm_logging_obj: response = await _fire_mcp_tool_call_logging( logging_obj=litellm_logging_obj, @@ -2547,7 +2577,8 @@ async def _handle_managed_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, -) -> CallToolResult: + wire_compat: WireCompat = WireCompat.LEGACY, +) -> CallToolResult | InputRequiredResult: """Handle tool execution for managed server tools""" # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj @@ -2566,12 +2597,15 @@ async def _handle_managed_mcp_tool( host_progress_callback=host_progress_callback, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + wire_compat=wire_compat, ) verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result) return call_tool_result -async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> CallToolResult: +async def _handle_local_mcp_tool( + name: str, arguments: dict[str, object], wire_compat: WireCompat = WireCompat.LEGACY +) -> CallToolResult: """Execute a local-registry tool and report whether it succeeded. Returns the result rather than bare content because the verdict is part of it: the content @@ -2604,10 +2638,7 @@ async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> Cal content=[TextContent(text=f"Error: {e}", type="text")], # mutable-ok: MCP result content is_error=True, ) - return CallToolResult( - content=[TextContent(text=str(result), type="text")], # mutable-ok: MCP result content - is_error=False, - ) + return complete_call_tool_result(handler_outcome(result), wire_compat) _MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset( @@ -2694,7 +2725,7 @@ async def _execute_handle_list_tools( async def _execute_mcp_server_tool_call( context: OperationContext, params: CallToolRequestParams, host_progress_callback: ProgressCallback | None = None -) -> CallToolResult: +) -> CallToolResult | InputRequiredResult: from mcp.types import CallToolResult from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException @@ -2778,6 +2809,7 @@ async def _execute_mcp_server_tool_call( raw_headers=raw_headers, client_ip=_client_ip, host_progress_callback=host_progress_callback, + wire_compat=context.wire_compat, **data, # for logging ) except MCPMissingUserEnvVarsError as e: @@ -3032,6 +3064,7 @@ def prepare_context( raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + wire_compat: WireCompat = WireCompat.LEGACY, ) -> OperationContext: return OperationContext( _caller=user_api_key_auth, @@ -3042,6 +3075,7 @@ def prepare_context( raw_headers=raw_headers, client_ip=client_ip, mcp_proxy_mode=mcp_proxy_mode, + wire_compat=wire_compat, ) @@ -3058,6 +3092,7 @@ GatewayOperation: TypeAlias = ( GatewayResult: TypeAlias = ( ListToolsResult | CallToolResult + | InputRequiredResult | ListPromptsResult | GetPromptResult | ListResourcesResult @@ -3071,13 +3106,17 @@ class GatewayOperations: self._host_progress_callback = host_progress_callback @overload - async def execute(self, operation: AuthorizedToolCall, context: OperationContext) -> CallToolResult: ... + async def execute( + self, operation: AuthorizedToolCall, context: OperationContext + ) -> CallToolResult | InputRequiredResult: ... @overload async def execute(self, operation: ListToolsRequest, context: OperationContext) -> ListToolsResult: ... @overload - async def execute(self, operation: CallToolRequest, context: OperationContext) -> CallToolResult: ... + async def execute( + self, operation: CallToolRequest, context: OperationContext + ) -> CallToolResult | InputRequiredResult: ... @overload async def execute(self, operation: ListPromptsRequest, context: OperationContext) -> ListPromptsResult: ... @@ -3113,6 +3152,7 @@ class GatewayOperations: client_ip=_client_ip, host_progress_callback=operation.host_progress_callback, guardrail_context=operation.guardrail_context, + wire_compat=context.wire_compat, **operation.logging_data, ) case ListToolsRequest(params=params): diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 5922285f643..7f519e2c0d9 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -36,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( ) from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url +from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result from litellm.proxy._experimental.mcp_server.ui_session_utils import ( acting_user_auth, build_effective_auth_contexts, @@ -1197,7 +1198,7 @@ if MCP_AVAILABLE: # Call execute_mcp_tool directly (permission checks already done) _tool_start_time: Final = datetime.now() - result: Final = await execute_mcp_tool( + executed: Final = await execute_mcp_tool( name=tool_name, arguments=tool_arguments, allowed_mcp_servers=allowed_mcp_servers, @@ -1212,6 +1213,7 @@ if MCP_AVAILABLE: guardrail_context=MCPRequestContext.resolve_guardrail_context(data), requested_server_id=canonical_server_id, ) + result: Final = complete_call_tool_result(executed, WireCompat.LEGACY) except Exception as e: request_data: Final = proxy_base_llm_response_processor.data await _safe_fire_mcp_tool_call_failure_logging( diff --git a/litellm/proxy/_experimental/mcp_server/result_conversion.py b/litellm/proxy/_experimental/mcp_server/result_conversion.py new file mode 100644 index 00000000000..52931fae116 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/result_conversion.py @@ -0,0 +1,120 @@ +"""Compatibility-aware conversion of upstream outcomes into MCP SDK results. + +Every gateway surface that turns a tool outcome (text, JSON, an SDK result, an +interim result, an exception) into the ``CallToolResult`` it sends downstream +goes through ``to_call_tool_result`` so the per-revision wire rules live in one +place. SDK 2.x serializes ``structuredContent`` as object-only on the handshake +revisions (``2024-11-05`` .. ``2025-11-25``) and admits any JSON value, plus +``input_required`` interim results, only on ``2026-07-28``. +""" + +from __future__ import annotations + +import json +from typing import Final, TypeAlias + +from mcp.types import CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool +from typing_extensions import ReadOnly, TypedDict, assert_never + +from litellm.proxy._experimental.mcp_server.tool_outcome import ( + JsonResult, + TextResult, + WireCompat, + handler_outcome, + parse_http_body, + wire_compat_for, +) + +__all__ = ( + "INPUT_REQUIRED_UNSUPPORTED_MESSAGE", + "JsonResult", + "TextResult", + "ToolOutcome", + "WireCompat", + "complete_call_tool_result", + "error_text_result", + "handler_outcome", + "parse_http_body", + "to_call_tool_result", + "to_gateway_tool", + "wire_compat_for", +) + +ToolOutcome: TypeAlias = TextResult | JsonResult | CallToolResult | InputRequiredResult | Exception + + +class _Downgraded(TypedDict): + structured_content: ReadOnly[None] + content: ReadOnly[list[ContentBlock]] # mutable-ok: SDK list field + + +class _Renamed(TypedDict): + name: ReadOnly[str] + + +INPUT_REQUIRED_UNSUPPORTED_MESSAGE: Final = ( + "Error: upstream tool returned an input_required interim result, which this MCP protocol revision cannot carry" +) + + +def error_text_result(exc: Exception) -> CallToolResult: + return CallToolResult( + content=[TextContent(type="text", text=f"{type(exc).__name__}: {exc}")], # mutable-ok: SDK list field + is_error=True, + ) + + +def to_call_tool_result(outcome: ToolOutcome, compat: WireCompat) -> CallToolResult | InputRequiredResult: + match outcome: + case TextResult(): + return CallToolResult( + content=[TextContent(type="text", text=outcome.text)], # mutable-ok: SDK list field + is_error=False, + ) + case JsonResult(): + keep_structured: Final = compat is WireCompat.MODERN or isinstance(outcome.value, dict) + return CallToolResult( + content=[TextContent(type="text", text=outcome.original_text)], # mutable-ok: SDK list field + is_error=False, + structured_content=outcome.value if keep_structured else None, + ) + case CallToolResult(): + return _downgrade_structured_content(outcome) if compat is WireCompat.LEGACY else outcome + case InputRequiredResult(): + if compat is WireCompat.MODERN: + return outcome + return CallToolResult( + content=[TextContent(type="text", text=INPUT_REQUIRED_UNSUPPORTED_MESSAGE)], # mutable-ok: SDK + is_error=True, + ) + case Exception(): + return error_text_result(outcome) + return assert_never(outcome) + + +def complete_call_tool_result(outcome: ToolOutcome, compat: WireCompat) -> CallToolResult: + """``to_call_tool_result`` for callers that can never carry an interim result.""" + converted: Final = to_call_tool_result(outcome, compat) + if isinstance(converted, InputRequiredResult): + return CallToolResult( + content=[TextContent(type="text", text=INPUT_REQUIRED_UNSUPPORTED_MESSAGE)], # mutable-ok: SDK + is_error=True, + ) + return converted + + +def _downgrade_structured_content(result: CallToolResult) -> CallToolResult: + structured: Final = result.structured_content + if structured is None or isinstance(structured, dict): + return result + fallback: Final = TextContent(type="text", text=json.dumps(structured)) + update: Final[_Downgraded] = { + "structured_content": None, + "content": [*result.content, fallback], # mutable-ok: SDK list field + } + return result.model_copy(update=update) + + +def to_gateway_tool(tool: Tool, name: str) -> Tool: + update: Final[_Renamed] = {"name": name} + return tool.model_copy(deep=True, update=update) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 6261f36983d..1bd31d971b0 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -145,6 +145,7 @@ try: from mcp import ReadResourceResult, Resource from mcp.server import Server + from mcp.server.runner import serve_loop from mcp.server.session import ServerSession as _McpServerSession from mcp.types import ( BlobResourceContents, @@ -504,6 +505,7 @@ if MCP_AVAILABLE: _invalidate_byok_cred_cache, _mcp_session_id_from_headers, ) + from litellm.proxy._experimental.mcp_server.result_conversion import wire_compat_for try: from mcp.server.streamable_http_manager import StreamableHTTPSessionManager @@ -516,6 +518,7 @@ if MCP_AVAILABLE: GetPromptRequestParams, Implementation, InitializeRequest, + InputRequiredResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult, @@ -818,7 +821,15 @@ if MCP_AVAILABLE: client_ip, ) = await get_or_extract_auth_context() yield operations.prepare_context( - auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get() + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + _mcp_proxy_mode.get(), + wire_compat_for(ctx.protocol_version), ) async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: @@ -875,7 +886,9 @@ if MCP_AVAILABLE: _dispatch_virtual_mcp_tool, ) - async def mcp_server_tool_call(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + async def mcp_server_tool_call( + ctx: ServerRequestContext, params: CallToolRequestParams + ) -> CallToolResult | InputRequiredResult: async with _legacy_operation_context(ctx, trace=True) as context: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( CallToolRequest(params=params), context @@ -2384,8 +2397,17 @@ if MCP_AVAILABLE: scoped_server_endpoint=scoped_server_endpoint, is_initialize=scope.get("method") == "GET", ): - async with sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream): - await server.run(read_stream, write_stream, server.create_initialization_options()) + async with ( + sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream), + server.lifespan(server) as lifespan_state, + ): + await serve_loop( + server, + read_stream, + write_stream, + lifespan_state=lifespan_state, + init_options=server.create_initialization_options(), + ) except MCPUpstreamAuthError as e: # Upstream delegated auth returned 401; surface it to the client so # standards-compliant MCP clients trigger the upstream OAuth flow. diff --git a/litellm/proxy/_experimental/mcp_server/tool_outcome.py b/litellm/proxy/_experimental/mcp_server/tool_outcome.py new file mode 100644 index 00000000000..ac712241b5b --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/tool_outcome.py @@ -0,0 +1,56 @@ +"""SDK-free half of the result conversion boundary. + +``openapi_to_mcp_generator`` and ``contracts`` must import without the ``mcp`` +package installed, so the compatibility enum and the tagged handler outcomes +live here; ``result_conversion`` turns them into SDK results. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum +from typing import Final + +from mcp_types.version import MODERN_PROTOCOL_VERSIONS +from pydantic import JsonValue, TypeAdapter, ValidationError + +_JSON_VALUE: Final = TypeAdapter(JsonValue) + + +class WireCompat(str, Enum): + LEGACY = "legacy" + MODERN = "modern" + + +def wire_compat_for(protocol_version: str) -> WireCompat: + return WireCompat.MODERN if protocol_version in MODERN_PROTOCOL_VERSIONS else WireCompat.LEGACY + + +@dataclass(frozen=True, slots=True) +class TextResult: + text: str + + +@dataclass(frozen=True, slots=True) +class JsonResult: + value: JsonValue + original_text: str + + +def parse_http_body(body: str) -> TextResult | JsonResult: + if not body.strip(): + return TextResult(body) + try: + value: Final = _JSON_VALUE.validate_json(body) + except ValidationError: + return TextResult(body) + if value is None: + return TextResult(body) + return JsonResult(value=value, original_text=body) + + +def handler_outcome(value: object) -> TextResult | JsonResult: + """Normalize what a registered tool handler returned; OpenAPI handlers already return a tagged outcome.""" + if isinstance(value, (TextResult, JsonResult)): + return value + return TextResult(str(value)) diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 71e46f8df25..9d117a1a1fa 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -13,6 +13,7 @@ from typing_extensions import ReadOnly, Required, assert_never import litellm from litellm.llms.litellm_proxy.skills.skill_search import DEFAULT_SKILL_SEARCH_TOP_K +from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result from litellm.proxy.agent_endpoints.agent_search import DEFAULT_AGENT_SEARCH_TOP_K from litellm.proxy.common_utils.semantic_text_index import ( Embedder, @@ -634,7 +635,7 @@ async def handle_mcp_tool_call( raise HTTPException(status_code=403, detail="User not allowed to call this tool.") - return await execute_mcp_tool( + result: Final = await execute_mcp_tool( name=tool_name, arguments=arguments, allowed_mcp_servers=allowed_mcp_servers, @@ -649,3 +650,4 @@ async def handle_mcp_tool_call( requested_server_id=requested_server_id, guardrail_context=guardrail_context, ) + return complete_call_tool_result(result, WireCompat.LEGACY) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 67950e603c0..f19a8055ae6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, ) +from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, @@ -722,6 +723,7 @@ async def _run_project_checks( model=_model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if not skip_budget_checks: @@ -1018,6 +1020,7 @@ async def common_checks( team_object=team_object, llm_router=llm_router, team_model_aliases=(valid_token.team_model_aliases if valid_token else None), + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -1027,6 +1030,7 @@ async def common_checks( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -1043,6 +1047,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent @@ -1081,6 +1086,7 @@ async def common_checks( model=_model, llm_router=llm_router, user_object=user_object, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget) @@ -4349,6 +4355,7 @@ def _can_object_call_model( models: list[str], team_model_aliases: dict[str, str] | None = None, team_id: str | None = None, + key_model_aliases: Mapping[str, str] | None = None, object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -4378,6 +4385,7 @@ def _can_object_call_model( models=models, team_model_aliases=team_model_aliases, team_id=team_id, + key_model_aliases=key_model_aliases, object_type=object_type, fallback_depth=fallback_depth + 1, ) @@ -4386,13 +4394,32 @@ def _can_object_call_model( from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent compaction_parent: Final = native_compaction_parent(model) - potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model] - if model in litellm.model_alias_map: - potential_models.append(litellm.model_alias_map[model]) - elif llm_router and model in llm_router.model_group_alias: - _model: Final = llm_router._get_model_from_alias(model) - if _model: - potential_models.append(_model) + global_or_router_alias_target: Final = ( + litellm.model_alias_map[model] + if model in litellm.model_alias_map + else ( + llm_router._get_model_from_alias(model) + if llm_router is not None and model in llm_router.model_group_alias + else None + ) + ) + after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model + after_key_alias: Final = ( + key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias + ) + after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias) + dispatched_model: Final = ( + key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias + ) + key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias + potential_models: Final = ( + (dispatched_model,) + if key_alias_applied + else ( + *((model, compaction_parent) if compaction_parent is not None else (model,)), + *((global_or_router_alias_target,) if global_or_router_alias_target else ()), + ) + ) ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: @@ -4418,6 +4445,35 @@ def _can_object_call_model( ) +def _resolve_team_alias( + model: str | list[str], + team_model_aliases: dict[str, str] | None, + team_id: str | None, + llm_router: Router | None, +) -> str | list[str]: + if not team_model_aliases: + return model + if isinstance(model, str): + return _live_team_alias_target(model, team_model_aliases, team_id, llm_router) + return [ # mutable-ok: _can_object_call_model takes list[str] + _live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model + ] + + +def _live_team_alias_target( + model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None +) -> str: + target: Final = team_model_aliases.get(model) + if target is None: + return model + deleted_team_deployment: Final = ( + llm_router is not None + and target.startswith(f"model_name_{team_id}_") + and target not in llm_router.model_name_to_deployment_indices + ) + return model if deleted_team_deployment else target + + async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, @@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) + dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) return _can_object_call_model( - model=model, + model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), team_id=valid_token.team_id, object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access( if caller_auth is None: return caller_team: Final = await load_team(valid_token) + caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token) if caller_team is not None: await can_team_access_model( model=model, team_object=caller_team, llm_router=llm_router, prisma_client=prisma_client, + key_model_aliases=caller_key_model_aliases, ) await _check_team_member_model_access( model=model, @@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=caller_key_model_aliases, ) return caller_user: Final = await load_user(valid_token) if caller_user is None: return - await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user) + await can_user_call_model( + model=model, + llm_router=llm_router, + user_object=caller_user, + key_model_aliases=caller_key_model_aliases, + ) def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool: @@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None return False +def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None: + return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None + + def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]: """ Expand key model sentinels before auth checks. @@ -4831,6 +4901,7 @@ async def can_key_call_model( models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) except ProxyException: @@ -4848,6 +4919,7 @@ async def can_key_call_model( models=models_from_groups, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) raise @@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model( team_object=team_object, llm_router=llm_router, team_model_aliases=valid_token.team_model_aliases, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if valid_token.project_id is not None: @@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model( model=model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4968,6 +5044,7 @@ async def can_team_access_model( team_object: LiteLLM_TeamTable | None, llm_router: Router | None, team_model_aliases: dict[str, str] | None = None, + key_model_aliases: Mapping[str, str] | None = None, prisma_client: DatabaseClient | None = None, ) -> Literal[True]: """ @@ -4983,6 +5060,7 @@ async def can_team_access_model( models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) except ProxyException: @@ -5000,6 +5078,7 @@ async def can_team_access_model( models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) raise @@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> bool: """ Returns True if the key's `access_group_ids` expand to models that grant @@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model( models=authorized_models, team_model_aliases=valid_token.team_model_aliases if valid_token else None, team_id=valid_token.team_id if valid_token else None, + key_model_aliases=key_model_aliases, object_type="key", ) return True @@ -5089,6 +5170,7 @@ def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: """ Returns True if the project can access a specific model. @@ -5099,6 +5181,7 @@ def can_project_access_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], + key_model_aliases=key_model_aliases, object_type="project", ) @@ -5107,6 +5190,7 @@ async def can_user_call_model( model: str | list[str], llm_router: Router | None, user_object: LiteLLM_UserTable | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: if user_object is None: return True @@ -5128,6 +5212,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + key_model_aliases=key_model_aliases, object_type="user", ) @@ -5682,6 +5767,7 @@ async def _check_team_member_model_access( proxy_logging_obj: ProxyLogging, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, + key_model_aliases: Mapping[str, str] | None = None, ) -> None: """ Check if a team member's per-member model scope allows access to the requested model. @@ -5717,6 +5803,7 @@ async def _check_team_member_model_access( models=member_allowed_models, object_type="team", team_id=team_object.team_id, + key_model_aliases=key_model_aliases, ) except ProxyException: internal_message: Final = ( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ae95e94dd2d..22c3a248b9d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import ( get_user_object, is_valid_fallback_model, jwt_key_mapping_cache_key, + key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, ) @@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback( models=valid_token.team_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="team", ) except ProxyException: diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 11cb66d1a7f..3e09ad7157f 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -74,12 +74,6 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None: async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None: try: client: Final = _pubsub_capable_client(redis_cache) - if client is None: - verbose_proxy_logger.debug( - "auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support", - cache_key, - ) - return async with _in_flight_publishes: await client.publish(auth_cache_invalidation_channel(redis_cache), message) except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors @@ -185,12 +179,6 @@ class AuthCacheInvalidationSubscriber: while True: try: client = _pubsub_capable_client(self._redis_cache) - if client is None: - verbose_proxy_logger.warning( - "auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; " - "cross-worker eviction falls back to the local cache TTL" - ) - return pubsub = client.pubsub() try: await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index b20c0d9c9a5..b4ebb5fa876 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -82,22 +82,13 @@ def config_sync_channel(redis_cache: "RedisCache") -> str: return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}" -def _raw_async_client(redis_cache: "RedisCache") -> object: - return cast( # cast-ok: redis-py generics leave the client type partially unknown - object, - redis_cache.init_async_client(), # pyright: ignore[reportUnknownMemberType] # redis generics +def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient: + return cast( # cast-ok: protocol view of the pub/sub-capable async redis client + _ConfigSyncPubSubClient, + redis_cache.init_pubsub_client(), # pyright: ignore[reportUnknownMemberType] # redis generics ) -def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient | None: - from redis.asyncio import Redis - - client: Final = _raw_async_client(redis_cache) - if isinstance(client, Redis): - return cast(_ConfigSyncPubSubClient, client) # cast-ok: protocol view of the standalone redis client - return None - - @dataclass(frozen=True, slots=True) class _ConfigChangeMessage: object_type: str @@ -112,12 +103,6 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s return try: client: Final = _pubsub_capable_client(redis_cache) - if client is None: - verbose_proxy_logger.debug( - "config sync publish for %s skipped: cluster redis client has no pub/sub support", - object_type, - ) - return await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e) @@ -238,12 +223,6 @@ class ConfigSyncSubscriber: while True: try: client = _pubsub_capable_client(self._redis_cache) - if client is None: - verbose_proxy_logger.warning( - "config sync subscriber disabled: cluster redis client has no pub/sub support; " - "interval polling remains the only sync mechanism" - ) - return pubsub = client.pubsub() try: await pubsub.subscribe(config_sync_channel(self._redis_cache)) diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 3c6555662e6..8958fb20918 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -180,7 +180,7 @@ def caller_alias_maps( return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases)) -def _alias_map(aliases: object) -> Mapping[str, str]: +def alias_map(aliases: object) -> Mapping[str, str]: try: entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True) except ValidationError: @@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] = already `listed` keeps its own row, so it is never rewritten.""" if model_id in listed: return None - return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)) + return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite)) def alias_listing_entries( @@ -213,8 +213,8 @@ def alias_listing_entries( ) -> tuple[tuple[str, str], ...]: """`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is listed. An alias colliding with a listed id keeps the listed entry.""" - maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite) - own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own) + maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite) + own: Final = tuple(alias_map(raw) for raw in aliases.own) lookup_by_response: Final = MappingProxyType(dict(entries)) lookup_ids: Final = frozenset(lookup_by_response.values()) targets: Final = MappingProxyType( diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 85e19fa8a32..c6f52bf074b 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -37,6 +37,7 @@ StopReason: TypeAlias = Literal["exhausted", "budget_exhausted", "batch_cap_reac class TableCleanupResult: """Outcome of pruning one table, so the caller can report why a run ended.""" + table_name: str rows_deleted: int stop_reason: StopReason @@ -472,11 +473,11 @@ class SpendLogCleanup: from the last run that finished inside its budget. """ if time.monotonic() >= deadline: - return TableCleanupResult(rows_deleted=rows_deleted, stop_reason=stop_reason) + return TableCleanupResult(table_name=table_name, rows_deleted=rows_deleted, stop_reason=stop_reason) remaining: Final = await self._count_remaining(prisma_client, cutoff_date, table_name, time_column, deadline) if remaining is not None: SpendLogCleanupMetrics.set_rows_remaining(table_name, remaining) - return TableCleanupResult(rows_deleted=rows_deleted, stop_reason=stop_reason) + return TableCleanupResult(table_name=table_name, rows_deleted=rows_deleted, stop_reason=stop_reason) async def _delete_old_logs( self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float @@ -571,7 +572,9 @@ class SpendLogCleanup: ) verbose_proxy_logger.info("Dropped %d expired spend-log partitions: %s", len(dropped), dropped) - logs_result: Final = await self._delete_old_logs(prisma_client, cutoff_date, deadline) + logs_result: Final = await self._delete_old_logs( + prisma_client, cutoff_date, self._group_deadline(deadline, groups_remaining=2) + ) verbose_proxy_logger.info("Deleted %s logs", logs_result.rows_deleted) index_result: Final = await self._delete_old_tool_index_rows(prisma_client, cutoff_date, deadline) @@ -638,6 +641,17 @@ class SpendLogCleanup: return "batch_cap_reached" return "completed" + @staticmethod + def _log_run_summary(outcome: RunOutcome, results: tuple[TableCleanupResult, ...], elapsed_seconds: float) -> None: + per_table: Final = ", ".join( + f"{result.table_name}: deleted={result.rows_deleted} stop_reason={result.stop_reason}" for result in results + ) + message: Final = "Spend log cleanup run finished: outcome=%s elapsed=%.1fs [%s]" + if outcome == "completed": + verbose_proxy_logger.info(message, outcome, elapsed_seconds, per_table) + return + verbose_proxy_logger.warning(message, outcome, elapsed_seconds, per_table) + async def cleanup_old_spend_logs(self, prisma_client: PrismaClient) -> None: """ Main cleanup function. Deletes old spend logs in batches. @@ -724,9 +738,10 @@ class SpendLogCleanup: else () ) - SpendLogCleanupMetrics.record_run( - self._run_outcome(spend_log_results + session_results + health_check_results) - ) + results: Final = spend_log_results + session_results + health_check_results + outcome: Final = self._run_outcome(results) + SpendLogCleanupMetrics.record_run(outcome) + self._log_run_summary(outcome, results, time.monotonic() - run_started_at) except asyncio.CancelledError: verbose_proxy_logger.error( diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index 28220f09f00..26dd2a95dc4 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -7,6 +7,7 @@ hook scan such a stream, and re-emit it when the guardrail rewrote the response. from __future__ import annotations +import codecs import json from collections.abc import Mapping, Sequence from typing import Final @@ -38,7 +39,7 @@ def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None: if isinstance(chunk, (str, bytes)) ) try: - return raw.decode("utf-8") + return codecs.getincrementaldecoder("utf-8")().decode(raw, final=False) except UnicodeDecodeError: return None diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 794bf08729e..f5e24c501f1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -10,9 +10,11 @@ import asyncio import json +import re import threading -from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Sequence +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Iterator, Sequence from contextlib import asynccontextmanager +from dataclasses import dataclass from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast @@ -39,6 +41,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames from litellm.proxy.guardrails.anthropic_sse import ( anthropic_sse_chunks_from_response, assemble_anthropic_sse_stream, @@ -97,16 +100,47 @@ def _json_escaped_len(text: str) -> int: _MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024 -def _holds_complete_sse_frame(raw: bytes) -> bool: - """Whether ``raw`` holds one blank-line terminated SSE event, or is too large to keep joining.""" - return b"\n\n" in raw or b"\r\n\r\n" in raw or len(raw) >= _MAX_FIRST_SSE_FRAME_BYTES +@dataclass(frozen=True, slots=True) +class _SsePreface: + """Complete leading SSE frames with no ``data:`` line, relayed verbatim before the stream shape is decided.""" + + raw: bytes + + +_SSE_FRAME_END: Final = re.compile(rb"\r\n\r\n|\n\n|\r\r") + + +def _split_sse_preface(complete_frames: bytes) -> tuple[bytes, bytes]: + """Split complete frames into ``(frames before the first data-bearing frame, that frame and everything after)``.""" + start = 0 + for end in _SSE_FRAME_END.finditer(complete_frames): + frame = complete_frames[start : end.end()] + if any(line.startswith(b"data:") for line in frame.splitlines()): + return complete_frames[:start], complete_frames[start:] + start = end.end() + return complete_frames, b"" + + +def _flush_unmaskable_buffer(all_chunks: list[ModelResponseStream]) -> Iterator[ModelResponseStream]: + """Buffered chunks flushed unmasked when a mixed stream shape makes reconstruction impossible.""" + if not all_chunks: + return + verbose_proxy_logger.warning( + "Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). " + "Flushing %d buffered chunks without PII masking and switching to transparent passthrough.", + len(all_chunks), + ) + yield from all_chunks async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGenerator[object, None]: """ - Join leading raw ``bytes`` chunks until they hold one complete SSE event, so - the stream shape is decided on a whole frame rather than a transport fragment. - Everything after that first frame is forwarded untouched. + Relay leading data-less SSE frames (comment keepalives, events without a + ``data:`` line) as they complete, and join raw ``bytes`` chunks until they + hold one complete SSE event with a data line, so the stream shape is + decided on a whole frame rather than a transport fragment. Everything + after that first frame is forwarded untouched. The byte cap can only be + reached by a single unterminated frame. """ pending = b"" try: @@ -115,7 +149,12 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener yield chunk continue pending += chunk - if _holds_complete_sse_frame(pending): + complete_frames, tail = split_complete_sse_frames(pending) + preface, classifiable = _split_sse_preface(complete_frames) + if preface: + yield _SsePreface(preface) + pending = classifiable + tail + if classifiable or len(pending) >= _MAX_FIRST_SSE_FRAME_BYTES: break else: if pending: @@ -1400,6 +1439,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield chunk else: all_chunks.append(chunk) + elif isinstance(chunk, _SsePreface): + yield chunk.raw elif isinstance(chunk, bytes): first_frame_is_anthropic = ( not passthrough_due_to_unknown_stream_shape @@ -1416,18 +1457,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield masked_chunk return else: - if all_chunks: - # Flush buffered chunks and switch to transparent passthrough for this stream shape. - # NOTE: these buffered chunks are emitted unmasked because this - # stream mixed chunk types and cannot be safely reconstructed. - verbose_proxy_logger.warning( - "Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). " - "Flushing %d buffered chunks without PII masking and switching to transparent passthrough.", - len(all_chunks), - ) - for buffered_chunk in all_chunks: - yield buffered_chunk - all_chunks = [] + for buffered_chunk in _flush_unmaskable_buffer(all_chunks): + yield buffered_chunk + all_chunks = [] passthrough_due_to_unknown_stream_shape = True yield chunk if passthrough_due_to_unknown_stream_shape: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index f8801e65c82..fbd4d57bf77 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -395,7 +395,9 @@ async def health_services_endpoint( from litellm.integrations.langfuse.langfuse import LangFuseLogger langfuse_logger: Final = LangFuseLogger() - langfuse_logger.Langfuse.auth_check() + auth_failure: Final = langfuse_logger.api_client.auth_check() + if auth_failure is not None: + raise ValueError(f"langfuse auth_check failed: {auth_failure.reason}") _ = litellm.completion( model="openai/litellm-mock-response-model", messages=[{"role": "user", "content": "Hey, how's it going?"}], diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index cfa54ae01a2..7ce50bf5ead 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -1,3 +1,4 @@ +import asyncio import json import time from collections.abc import Iterable, Mapping, Sequence @@ -297,6 +298,9 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): def __init__(self, dual_cache: DualCache): self.dual_cache = dual_cache self.redis_increment_operation_queue = [] + self._redis_increment_queue_lock = asyncio.Lock() + self._redis_increment_flush_lock = asyncio.Lock() + self._detached_increment_operations = None self.deployment_budget_config = None async def is_key_within_model_budget( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2f1dc270cf6..7e159ec90e7 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -26,6 +26,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeV import fastapi import yaml from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status +from pydantic import TypeAdapter from typing_extensions import ReadOnly, TypedDict import litellm @@ -99,6 +100,11 @@ from litellm.proxy.management_endpoints.model_management_endpoints import ( _add_model_to_db, ) from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights +from litellm.proxy.management_endpoints.team_admin_field_permissions import ( + team_admin_key_edit_verdict, + team_admin_key_request_or_raise, + team_admin_may_edit_member_key_budgets, +) from litellm.proxy.management_helpers.access_group_key_sync import ( sync_key_access_group_membership, sync_key_regeneration_access_group_membership, @@ -3008,6 +3014,55 @@ async def _validate_end_user_budget_id_change( raise HTTPException(status_code=400, detail=missing_detail) +_GENERAL_SETTINGS: Final = TypeAdapter(dict[str, object]) + + +def _general_settings() -> Mapping[str, object]: + from litellm.proxy.proxy_server import ( + general_settings, # pyright: ignore[reportUnknownVariableType] # untyped module-level dict in proxy_server + ) + + return _GENERAL_SETTINGS.validate_python(general_settings) + + +async def _acting_as_team_admin_for_key_update( + data: UpdateKeyRequest, + existing_key_row: LiteLLM_VerificationToken, + user_api_key_dict: UserAPIKeyAuth, + checked_prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + is_proxy_admin: bool, +) -> bool: + """Whether the caller acts as a team admin on another member's team key. + + Raises 403 when the caller administers the key's team but the request edits fields + outside the member_key_budgets permission (or that permission is disabled). + """ + if ( + is_proxy_admin + or existing_key_row.team_id is None + or existing_key_row.user_id is None + or existing_key_row.user_id == user_api_key_dict.user_id + ): + return False + team_for_grant: Final = await get_team_object( + team_id=existing_key_row.team_id, + prisma_client=checked_prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, + ) + if not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_for_grant): + return False + team_admin_key_request_or_raise( + team_admin_key_edit_verdict( + data=data, + existing=existing_key_row, + enabled=team_admin_may_edit_member_key_budgets(_general_settings()), + ) + ) + return True + + async def _validate_update_key_data( data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken, @@ -3058,10 +3113,19 @@ async def _validate_update_key_data( ) is_project_change: Final = "project_id" in data.model_fields_set and data.project_id != existing_key_row.project_id + acting_as_team_admin: Final = await _acting_as_team_admin_for_key_update( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=user_api_key_dict, + checked_prisma_client=checked_prisma_client, + user_api_key_cache=user_api_key_cache, + is_proxy_admin=_is_proxy_admin, + ) + common_key_access_checks( user_api_key_dict=user_api_key_dict, data=data, - user_id=existing_key_row.user_id, + user_id=user_api_key_dict.user_id if acting_as_team_admin else existing_key_row.user_id, llm_router=llm_router, premium_user=premium_user, ) diff --git a/litellm/proxy/management_endpoints/team_admin_field_permissions.py b/litellm/proxy/management_endpoints/team_admin_field_permissions.py index 6038775d96b..5146f5e0979 100644 --- a/litellm/proxy/management_endpoints/team_admin_field_permissions.py +++ b/litellm/proxy/management_endpoints/team_admin_field_permissions.py @@ -1,5 +1,6 @@ """Proxy-wide allow-list of what a team admin may do on the teams they administer: team-settings fields on -/team/update, plus the ``projects`` permission for /project/new and /project/update.""" +/team/update, the ``projects`` permission for /project/new and /project/update, and the +``member_key_budgets`` permission for budget fields on other members' keys via /key/update.""" from collections.abc import Mapping from dataclasses import dataclass @@ -12,9 +13,11 @@ from typing_extensions import assert_never from litellm._logging import verbose_proxy_logger from litellm.models.team import LiteLLM_TeamTable +from litellm.models.verification_token import LiteLLM_VerificationToken from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, LiteLLM_ManagementEndpoint_MetadataFields_Premium, + UpdateKeyRequest, UpdateTeamRequest, ) @@ -23,12 +26,20 @@ TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING: Final = "team_admin_editable_team_field # TODO(LIT-5722): add the remaining team settings one per PR, each with its value-diff tests and dashboard field SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS: Final[frozenset[str]] = frozenset({"tpm_limit", "rpm_limit", "max_budget"}) TEAM_ADMIN_PROJECTS_PERMISSION: Final = "projects" +TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION: Final = "member_key_budgets" SUPPORTED_TEAM_ADMIN_PERMISSIONS: Final[frozenset[str]] = SUPPORTED_TEAM_ADMIN_EDITABLE_TEAM_FIELDS | { - TEAM_ADMIN_PROJECTS_PERMISSION + TEAM_ADMIN_PROJECTS_PERMISSION, + TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION, } +# spend is deliberately excluded: the stored row lags the live cross-pod counter, so a value-diff gate +# would let a team admin overwrite real usage. +KEY_BUDGET_FIELDS: Final[frozenset[str]] = frozenset({"max_budget", "budget_duration", "soft_budget", "budget_limits"}) +_KEY_REQUEST_IDENTITY: Final[frozenset[str]] = frozenset({"key", "token", "metadata"}) + _FIELD_LIST: Final = TypeAdapter(list[str]) _JSON_OBJECT: Final = TypeAdapter(dict[str, object]) +_WINDOW_LIST: Final = TypeAdapter(list[dict[str, object]]) _EMPTY: Final[Mapping[str, object]] = MappingProxyType({}) _METADATA_FOLDED_FIELDS: Final[frozenset[str]] = frozenset( (*LiteLLM_ManagementEndpoint_MetadataFields, *LiteLLM_ManagementEndpoint_MetadataFields_Premium) @@ -89,6 +100,12 @@ def team_admin_may_manage_projects(general_settings: Mapping[str, object]) -> bo ) +def team_admin_may_edit_member_key_budgets(general_settings: Mapping[str, object]) -> bool: + return TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION in resolve_team_admin_editable_fields( + general_settings, frozenset({TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION}) + ) + + def _as_object(value: object) -> Mapping[str, object]: try: return _JSON_OBJECT.validate_json(value) if isinstance(value, str) else _JSON_OBJECT.validate_python(value) @@ -101,7 +118,7 @@ def _stored_metadata(existing: Mapping[str, object]) -> Mapping[str, object]: def _submitted_metadata( - data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object] + data: UpdateTeamRequest | UpdateKeyRequest, submitted: Mapping[str, object], existing: Mapping[str, object] ) -> Mapping[str, object]: """Metadata as it would be stored: the caller's dict (or the stored one) with top-level folded fields laid over.""" base: Final = ( @@ -112,7 +129,7 @@ def _submitted_metadata( def _metadata_changes( - data: UpdateTeamRequest, submitted: Mapping[str, object], existing: Mapping[str, object] + data: UpdateTeamRequest | UpdateKeyRequest, submitted: Mapping[str, object], existing: Mapping[str, object] ) -> frozenset[str]: merged: Final = _submitted_metadata(data, submitted, existing) stored: Final = _stored_metadata(existing) @@ -200,3 +217,109 @@ def team_admin_request_or_raise(verdict: TeamAdminEditVerdict) -> UpdateTeamRequ ) case _: assert_never(verdict) + + +def _budget_windows(value: object) -> frozenset[tuple[object, object]] | None: + """(budget_duration, max_budget) pairs for a stored or submitted budget_limits value. + + Stored windows carry server-added keys like ``reset_at``; only the caller-owned pair matters. + ``None`` means the value is not a list of windows and needs a plain comparison. + """ + if value is None: + return frozenset() + if not isinstance(value, list): + return None + try: + windows_input: Final = _WINDOW_LIST.validate_python(value) + except ValidationError: + return None + windows: Final = frozenset((window.get("budget_duration"), window.get("max_budget")) for window in windows_input) + if len(windows) != len(windows_input): + return None + return windows + + +def _key_column_changed(field: str, submitted: Mapping[str, object], existing: Mapping[str, object]) -> bool: + if field == "budget_limits": + sent: Final = _budget_windows(submitted.get(field)) + stored: Final = _budget_windows(existing.get(field)) + if sent is not None and stored is not None: + return sent != stored + if field in LiteLLM_VerificationToken.model_fields: + return submitted.get(field) != existing.get(field) + return True + + +def changed_key_fields(data: UpdateKeyRequest, existing_row: LiteLLM_VerificationToken) -> frozenset[str]: + """Logical field names whose stored value the key-update request would change. + + Same JSON-value comparison as :func:`changed_team_fields`: columns compare against the stored row, + fields the key endpoint folds into ``metadata`` compare against ``existing_row.metadata``, other + ``metadata`` keys are attributed to ``metadata``, and fields with no stored counterpart count as + changed whenever they are sent. ``budget_limits`` compares (budget_duration, max_budget) pairs so + order and server-computed ``reset_at`` values do not read as edits. + """ + submitted: Final = _JSON_OBJECT.validate_json(data.model_dump_json(exclude_unset=True)) + existing: Final = _JSON_OBJECT.validate_json(existing_row.model_dump_json()) + column_fields: Final = frozenset(data.model_fields_set) - _KEY_REQUEST_IDENTITY - _METADATA_FOLDED_FIELDS + column_changes: Final = frozenset( + field for field in column_fields if _key_column_changed(field, submitted, existing) + ) + return column_changes | _metadata_changes(data, submitted, existing) + + +@dataclass(frozen=True, slots=True) +class TeamAdminKeyEditAllowed: + changed: frozenset[str] + kind: Literal["allowed"] = "allowed" + + +@dataclass(frozen=True, slots=True) +class TeamAdminMemberKeyEditingDisabled: + kind: Literal["disabled"] = "disabled" + + +TeamAdminKeyEditVerdict: TypeAlias = ( + TeamAdminKeyEditAllowed | TeamAdminMemberKeyEditingDisabled | TeamAdminFieldNotPermitted +) + + +def team_admin_key_edit_verdict( + data: UpdateKeyRequest, + existing: LiteLLM_VerificationToken, + enabled: bool, +) -> TeamAdminKeyEditVerdict: + if not enabled: + return TeamAdminMemberKeyEditingDisabled() + changed: Final = changed_key_fields(data, existing) + blocked: Final = sorted( + (changed | (frozenset({"spend"}) if "spend" in data.model_fields_set else frozenset())) - KEY_BUDGET_FIELDS + ) + if blocked: + return TeamAdminFieldNotPermitted(field=blocked[0]) + return TeamAdminKeyEditAllowed(changed=changed) + + +def team_admin_key_request_or_raise(verdict: TeamAdminKeyEditVerdict) -> None: + match verdict: + case TeamAdminKeyEditAllowed(): + return + case TeamAdminMemberKeyEditingDisabled(): + raise HTTPException( + status_code=403, + detail=( + "Team admins on this proxy cannot update budgets on other members' keys. " + f"Ask a proxy admin to enable '{TEAM_ADMIN_MEMBER_KEY_BUDGETS_PERMISSION}' " + f"under {_SETTINGS_LOCATION}." + ), + ) + case TeamAdminFieldNotPermitted(field=field): + raise HTTPException( + status_code=403, + detail=( + "Team admins on this proxy may only update budget fields on other members' keys, " + f"not '{field}'. Ask a proxy admin to add it under {_SETTINGS_LOCATION}." + ), + ) + case _: + assert_never(verdict) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 4aaa77f8d45..b12a689429d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -565,17 +565,35 @@ async def _get_team_allowed_mcp_servers( """ Get the full set of MCP server IDs a team allows. - If team has no object_permission or no MCP config, returns empty set - (meaning only allow_all_keys servers are permitted). + Combines servers granted via the team's object_permission with servers + granted via the team's unified access groups (access_group_ids). If the + team grants neither, returns empty set (meaning only allow_all_keys + servers are permitted). """ if team_obj is None: return set() + from litellm.proxy.auth.auth_checks import ( + _get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same resolver runtime MCP auth calls + ) + + access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + access_group_ids=team_obj.access_group_ids or [], + prisma_client=prisma_client, + ) + resolved_access_group_servers: Final = await _resolve_mcp_server_identifiers_to_ids( + identifiers=set(access_group_servers), + prisma_client=prisma_client, + ) + unified_servers: Final = _flatten_resolved_mcp_server_ids(resolved_access_group_servers) | { + server for server in access_group_servers if not resolved_access_group_servers.get(server) + } + team_object_permission: Final = team_obj.object_permission if team_object_permission is None: - return set() + return unified_servers - return await _resolve_team_allowed_mcp_servers( + return unified_servers | await _resolve_team_allowed_mcp_servers( team_object_permission=team_object_permission, prisma_client=prisma_client, ) @@ -650,14 +668,17 @@ async def validate_key_mcp_servers_against_team( Rules: - If key is in a team: key's mcp_servers must be a subset of - (team's allowed servers + allow_all_keys servers) + (team's allowed servers + allow_all_keys servers), where the team's + allowed servers include servers granted via the team's unified + access groups - If key is NOT in a team and the caller is a proxy admin: any server or access group may be assigned. A proxy admin can already reach every MCP server, and runtime access is granted directly from the key's own object_permission, so the key is scoped to exactly what the admin selected - If key is NOT in a team and the caller is not a proxy admin: key's mcp_servers must only contain allow_all_keys servers - - If team has no MCP config: key can only use allow_all_keys servers + - If team has no MCP config (no object_permission and no unified + access groups): key can only use allow_all_keys servers Raises HTTPException(403) if validation fails. """ diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 7d6db30e3e3..e0a4184291e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -105,6 +105,8 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository from litellm.secret_managers.main import get_secret_str +from litellm.types import utils as types_utils +from litellm.types.litellm_params import ProxyRequestState, wire_names from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, @@ -133,6 +135,9 @@ router: Final = APIRouter() pass_through_endpoint_logging: Final = PassThroughEndpointLogging() +_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata")) +_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState)) + # Global registry to track registered pass-through routes and prevent memory leaks _registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {} @@ -578,21 +583,21 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): """ Filter out litellm params from the request body """ - from litellm.types.utils import all_litellm_params - _parsed_body = _parsed_body or {} - litellm_params_in_body: Final = {} - for k in all_litellm_params: - if k in _parsed_body: - litellm_params_in_body[k] = _parsed_body.pop(k, None) + litellm_keys_in_body: Final = MappingProxyType( + {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body} + ) + litellm_params_in_body: Final = MappingProxyType( + {k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS} + ) _metadata = dict( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) - litellm_metadata: Final = litellm_params_in_body.pop("litellm_metadata", None) - metadata: Final = litellm_params_in_body.pop("metadata", None) + litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata") + metadata: Final = litellm_keys_in_body.get("metadata") if litellm_metadata: _metadata.update(litellm_metadata) if metadata: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a5a6e9f6a9a..f8d9c9c7be9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -68,6 +68,7 @@ from litellm.constants import ( DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL, DEFAULT_SHARED_HEALTH_CHECK_TTL, DEFAULT_SLACK_ALERTING_THRESHOLD, + LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS, LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS, LITELLM_SETTINGS_SAFE_DB_OVERRIDES, LITELLM_UI_ALLOW_HEADERS, @@ -1122,17 +1123,21 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N if shutdown_billing_metrics_recorder is not None: shutdown_billing_metrics_recorder() - # flush remaining langfuse logs - if "langfuse" in litellm.success_callback: + if "litellm.integrations.langfuse.langfuse_sdk" in sys.modules: try: - # flush langfuse logs on shutdow - from litellm.utils import langFuseLogger + from litellm.integrations.langfuse.langfuse_sdk import flush_langfuse_tracing - if langFuseLogger is not None: - langFuseLogger.Langfuse.flush() - except Exception: - # [DO NOT BLOCK shutdown events for this] - pass + flushed: Final = await asyncio.to_thread(flush_langfuse_tracing, LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS) + if flushed: + verbose_proxy_logger.info("Langfuse export channels flushed") + else: + verbose_proxy_logger.warning( + "Langfuse shutdown flush incomplete: a channel did not finish within %dms or a batch was rejected " + "(see the export errors above); remaining spans are left to the background exporter", + LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS, + ) + except Exception as e: # noqa: BLE001 # shutdown must continue even if the flush fails + verbose_proxy_logger.exception("Error flushing Langfuse export channels on shutdown: %s", e) ## RESET CUSTOM VARIABLES ## cleanup_router_config_variables() diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 0ca08cb7992..87d38606aba 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -2639,6 +2639,24 @@ ], "default_model_placeholder": "gpt-3.5-turbo" }, + { + "provider": "Nadir", + "provider_display_name": "Nadir", + "litellm_provider": "nadir", + "credential_fields": [ + { + "key": "api_key", + "label": "API Key", + "placeholder": null, + "tooltip": null, + "required": true, + "field_type": "password", + "options": null, + "default_value": null + } + ], + "default_model_placeholder": "nadir/auto" + }, { "provider": "Oracle", "provider_display_name": "Oracle Cloud Infrastructure (OCI)", diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index d520177965c..c91b1afd64a 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -328,6 +328,7 @@ class UISettings(BaseModel): description=( "Team settings fields a team admin may change on the teams they administer. " "Include 'projects' to let team admins create and update projects for those teams. " + "Include 'member_key_budgets' to let team admins update budget fields on keys owned by other members of those teams. " "Empty means team admins cannot edit team settings or manage projects at all. " "Proxy admins and org admins are not affected." ), diff --git a/litellm/router.py b/litellm/router.py index dabbfcb6094..4a1491c9364 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import ( validate_routing_strategy, ) from litellm.scheduler import FlowItem, Scheduler +from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolParam, @@ -814,15 +815,7 @@ class Router: allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure disable_cooldowns: bool | None = None, - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - "cost-based-routing", - "usage-based-routing-v2", - "lar1", - ] = "simple-shuffle", + routing_strategy: RoutingStrategyName = "simple-shuffle", optional_pre_call_checks: OptionalPreCallChecks | None = None, routing_strategy_args: dict = {}, # just for latency-based routing_groups: list[RoutingGroup | dict] | None = None, @@ -2958,7 +2951,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) fallback_response = await self.async_function_with_fallbacks_common_utils( e=e, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3402,7 +3395,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3493,8 +3486,9 @@ class Router: for item in model_response: yield item except MidStreamFallbackError as e: - if not e.is_pre_first_chunk and ( - e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) + if fallbacks_disabled_for_request(initial_kwargs) or ( + not e.is_pre_first_chunk + and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)) ): if e.original_exception is not None: raise e.original_exception from e @@ -5629,7 +5623,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 3e094df7ac8..4e84bded9de 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -21,8 +21,10 @@ anthropic: import asyncio import builtins import logging -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone +from itertools import groupby +from types import MappingProxyType from typing import Any, Final import litellm @@ -93,11 +95,12 @@ class _LiteLLMParamsDictView: return dict(self._params) -async def _push_increments_to_redis(redis_cache: RedisCache, queued: list[RedisPipelineIncrementOperation]) -> None: - try: - await redis_cache.async_increment_pipeline(increment_list=queued) - except Exception as e: - log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e) +def _sum_increments_by_key(operations: Sequence[RedisPipelineIncrementOperation]) -> Mapping[str, float]: + by_key: Final = groupby( + sorted(operations, key=lambda operation: operation["key"]), + key=lambda operation: operation["key"], + ) + return MappingProxyType({key: sum(operation["increment_value"] for operation in group) for key, group in by_key}) class RouterBudgetLimiting(CustomLogger): @@ -109,6 +112,9 @@ class RouterBudgetLimiting(CustomLogger): ): self.dual_cache = dual_cache self.redis_increment_operation_queue: list[RedisPipelineIncrementOperation] = [] + self._redis_increment_queue_lock = asyncio.Lock() + self._redis_increment_flush_lock = asyncio.Lock() + self._detached_increment_operations: tuple[RedisPipelineIncrementOperation, ...] | None = None asyncio.create_task(self.periodic_sync_in_memory_spend_with_redis()) self.provider_budget_config: GenericBudgetConfigType | None = provider_budget_config self.deployment_budget_config: GenericBudgetConfigType | None = None @@ -392,17 +398,97 @@ class RouterBudgetLimiting(CustomLogger): - Increments the spend in memory cache (so spend instantly updated in memory) - Queues the increment operation to Redis Pipeline (using batched pipeline to optimize performance. Using Redis for multi instance environment of LiteLLM) """ - await self.dual_cache.in_memory_cache.async_increment( - key=spend_key, - value=response_cost, - ttl=ttl, - ) increment_op: Final = RedisPipelineIncrementOperation( key=spend_key, increment_value=response_cost, ttl=ttl, ) - self.redis_increment_operation_queue.append(increment_op) + async with self._get_redis_increment_queue_lock(): + await self.dual_cache.in_memory_cache.async_increment( + key=spend_key, + value=response_cost, + ttl=ttl, + ) + self.redis_increment_operation_queue.append(increment_op) + + def _get_redis_increment_queue_lock(self) -> asyncio.Lock: + return self._redis_increment_queue_lock + + async def _detach_queued_increment_operations(self) -> tuple[RedisPipelineIncrementOperation, ...]: + async with self._get_redis_increment_queue_lock(): + if self._detached_increment_operations is not None: + return self._detached_increment_operations + increment_operations_to_flush: Final = tuple(self.redis_increment_operation_queue) + if not increment_operations_to_flush: + return increment_operations_to_flush + self.redis_increment_operation_queue = [] # mutable-ok: emptied queue must stay appendable + self._detached_increment_operations = increment_operations_to_flush + return increment_operations_to_flush + + async def _clear_detached_increment_operations(self) -> None: + async with self._get_redis_increment_queue_lock(): + self._detached_increment_operations = None + + async def _requeue_detached_increment_operations(self) -> None: + async with self._get_redis_increment_queue_lock(): + detached_increment_operations: Final = self._detached_increment_operations + if detached_increment_operations is None: + return + operations: Final = (*detached_increment_operations, *self.redis_increment_operation_queue) + grouped_operations: Final = ( + (key, tuple(group)) + for key, group in groupby( + sorted(operations, key=lambda operation: operation["key"]), + key=lambda operation: operation["key"], + ) + ) + self.redis_increment_operation_queue = [ + RedisPipelineIncrementOperation( + key=key, + increment_value=sum(operation["increment_value"] for operation in group), + ttl=group[-1]["ttl"], + ) + for key, group in grouped_operations + ] + self._detached_increment_operations = None + + async def _flush_queued_increment_operations(self, redis_cache: RedisCache) -> bool: + flush_task: Final = asyncio.create_task(self._write_queued_increment_operations(redis_cache)) + return await self._await_flush_task(flush_task) + + async def _await_flush_task(self, flush_task: asyncio.Task[bool]) -> bool: + try: + return await asyncio.shield(flush_task) + except asyncio.CancelledError: + while not flush_task.done(): + try: + await asyncio.shield(flush_task) + except asyncio.CancelledError: + continue + flush_task.result() + raise + + async def _write_queued_increment_operations(self, redis_cache: RedisCache) -> bool: + increment_operations_to_flush: Final = await self._detach_queued_increment_operations() + if len(increment_operations_to_flush) == 0: + await self._clear_detached_increment_operations() + return True + + verbose_router_logger.debug( + "Pushing Redis Increment Pipeline for queue: %s", + increment_operations_to_flush, + ) + increment_list: Final = list( # mutable-ok: Redis pipeline contract requires a list + increment_operations_to_flush + ) + try: + await redis_cache.async_increment_pipeline(increment_list=increment_list) + except Exception as error: + log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", error) + await self._requeue_detached_increment_operations() + return False + await self._clear_detached_increment_operations() + return True async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): """Original method now uses helper functions""" @@ -528,29 +614,25 @@ class RouterBudgetLimiting(CustomLogger): DEFAULT_REDIS_SYNC_INTERVAL ) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying - async def _push_in_memory_increments_to_redis(self): + async def _push_in_memory_increments_to_redis(self) -> bool: """ How this works: - async_log_success_event collects all provider spend increments in `redis_increment_operation_queue` - This function pushes all increments to Redis in a batched pipeline to optimize performance - Only runs if Redis is initialized + Only runs if Redis is initialized. Returns False when the detached batch could not be + written, so callers must not treat Redis as up to date. """ - try: - if not self.dual_cache.redis_cache: - return # Redis is not initialized + redis_cache: Final = self.dual_cache.redis_cache + if redis_cache is None: + return True - verbose_router_logger.debug( - "Pushing Redis Increment Pipeline for queue: %s", - self.redis_increment_operation_queue, - ) - queued: Final = self.redis_increment_operation_queue - self.redis_increment_operation_queue = [] - if queued: - asyncio.create_task(_push_increments_to_redis(self.dual_cache.redis_cache, queued)) + flush_task: Final = asyncio.create_task(self._flush_queued_increments_with_lock(redis_cache)) + return await self._await_flush_task(flush_task) - except Exception as e: - log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e) + async def _flush_queued_increments_with_lock(self, redis_cache: RedisCache) -> bool: + async with self._redis_increment_flush_lock: + return await self._flush_queued_increment_operations(redis_cache) async def _sync_in_memory_spend_with_redis(self): """ @@ -569,44 +651,51 @@ class RouterBudgetLimiting(CustomLogger): # No need to sync if Redis cache is not initialized if self.dual_cache.redis_cache is None: return - - # 1. Push all provider spend increments to Redis - await self._push_in_memory_increments_to_redis() - - # 2. Fetch all current provider spend from Redis to update in-memory cache - cache_keys: Final = [] - - if self.provider_budget_config is not None: - for provider, config in self.provider_budget_config.items(): - if config is None: - continue - cache_keys.append(f"provider_spend:{provider}:{config.budget_duration}") - - if self.deployment_budget_config is not None: - for model_id, config in self.deployment_budget_config.items(): - if config is None: - continue - cache_keys.append(f"deployment_spend:{model_id}:{config.budget_duration}") - - if self.tag_budget_config is not None: - for tag, config in self.tag_budget_config.items(): - if config is None: - continue - cache_keys.append(f"tag_spend:{tag}:{config.budget_duration}") - - # Batch fetch current spend values from Redis - redis_values: Final = await self.dual_cache.redis_cache.async_batch_get_cache(key_list=cache_keys) - - # Update in-memory cache with Redis values - if isinstance(redis_values, dict): # Check if redis_values is a dictionary - for key, value in redis_values.items(): - if value is not None: - await self.dual_cache.in_memory_cache.async_set_cache(key=key, value=float(value)) - verbose_router_logger.debug("Updated in-memory cache for %s: %s", key, value) - + async with self._redis_increment_flush_lock: + await self._flush_increments_then_copy_redis_spend() except Exception as e: log_redis_failure(verbose_router_logger, logging.ERROR, "Error syncing in-memory cache with Redis", e) + async def _flush_increments_then_copy_redis_spend(self) -> None: + redis_cache: Final = self.dual_cache.redis_cache + if redis_cache is None or not await self._flush_queued_increment_operations(redis_cache): + return + + cache_keys: Final = [] + + if self.provider_budget_config is not None: + for provider, config in self.provider_budget_config.items(): + if config is None: + continue + cache_keys.append(f"provider_spend:{provider}:{config.budget_duration}") + + if self.deployment_budget_config is not None: + for model_id, config in self.deployment_budget_config.items(): + if config is None: + continue + cache_keys.append(f"deployment_spend:{model_id}:{config.budget_duration}") + + if self.tag_budget_config is not None: + for tag, config in self.tag_budget_config.items(): + if config is None: + continue + cache_keys.append(f"tag_spend:{tag}:{config.budget_duration}") + + redis_values: Final = await redis_cache.async_batch_get_cache(key_list=cache_keys) + + if not isinstance(redis_values, dict): + return + async with self._get_redis_increment_queue_lock(): + pending_spend_by_key: Final = _sum_increments_by_key(self.redis_increment_operation_queue) + updated_spend_by_key: Final = tuple( + (key, float(value) + pending_spend_by_key.get(key, 0.0)) + for key, value in redis_values.items() + if value is not None + ) + for key, updated_spend in updated_spend_by_key: + await self.dual_cache.in_memory_cache.async_set_cache(key=key, value=updated_spend) + verbose_router_logger.debug("Updated in-memory cache for %s: %s", key, updated_spend) + def _get_budget_config_for_deployment( self, model_id: str, diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index d9834adc7e8..6e455817194 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -109,6 +109,7 @@ RULES: Final[Rules] = ( RouteRule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY), RouteRule(Route.EMBEDDINGS, Rollout.PYTHON_ONLY), RouteRule(Route.OCR, Rollout.RUST_REQUIRED), + RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY, providers=frozenset({"anthropic"})), RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY), RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY), RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY), diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index 2f243e8c212..7a6485a5f2c 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import AsyncIterator, Awaitable, Iterator +from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping from dataclasses import dataclass from typing import Final, Protocol @@ -17,7 +17,7 @@ class Complete: @dataclass(frozen=True, slots=True) class Open: - value: None + value: Mapping[str, object] | None @dataclass(frozen=True, slots=True) @@ -68,7 +68,7 @@ async def drive(execution: Execution) -> object: step: Final = await _settle(execution, execution.start()) if isinstance(step, Open): handed_off = True - return Stream(execution) + return Stream(execution, step.value) return step.value finally: if not handed_off: @@ -78,10 +78,10 @@ async def drive(execution: Execution) -> object: class Stream(AsyncIterator[object]): """A streamed native call: each read resumes the execution until its next chunk.""" - def __init__(self, execution: Execution) -> None: + def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None: self._execution: Final = execution self._done = False - self._hidden_params: dict[str, object] = {} # mutable-ok: header writers mutate _hidden_params in place + self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it def __aiter__(self) -> Stream: return self @@ -115,10 +115,10 @@ class Stream(AsyncIterator[object]): class SyncStream(Iterator[object]): """The sync form of `Stream`; its execution never suspends on an awaitable.""" - def __init__(self, execution: Execution) -> None: + def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None: self._execution: Final = execution self._done = False - self._hidden_params: dict[str, object] = {} # mutable-ok: header writers mutate _hidden_params in place + self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it def __iter__(self) -> SyncStream: return self diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index d49d7b75a6f..0a23989a59c 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -4,6 +4,7 @@ from collections.abc import Mapping, Sequence from dataclasses import asdict, dataclass from typing import Final, cast # noqa: TID251 # narrows the normalized native payload to the public TypedDict +import httpx from pydantic import TypeAdapter, ValidationError import litellm @@ -53,6 +54,14 @@ def response(value: Mapping[str, object]) -> AnthropicMessagesResponse: ) +def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, object]: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + anthropic_messages_stream_hidden_params, + ) + + return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers))) + + def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]: return request.kwargs diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py index 5de58a20242..9a9f3ae34ce 100644 --- a/litellm/types/integrations/custom_logger.py +++ b/litellm/types/integrations/custom_logger.py @@ -3,8 +3,10 @@ from typing import Any, Final from pydantic import BaseModel, Field -CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions" -RESPONSES_AGENTIC_SURFACE: Final = "responses" +from litellm.types.litellm_params import AgenticSurface + +CHAT_COMPLETION_AGENTIC_SURFACE: Final[AgenticSurface] = "chat_completions" +RESPONSES_AGENTIC_SURFACE: Final[AgenticSurface] = "responses" CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception" HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception" HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream" diff --git a/litellm/types/integrations/langfuse.py b/litellm/types/integrations/langfuse.py index 6742aefea39..fe070a3dd18 100644 --- a/litellm/types/integrations/langfuse.py +++ b/litellm/types/integrations/langfuse.py @@ -14,3 +14,8 @@ class LangfuseUsageDetails(TypedDict): total: int | None cache_creation_input_tokens: int | None cache_read_input_tokens: int | None + + +class LangfuseLoggedEvent(TypedDict): + trace_id: ReadOnly[str | None] + generation_id: ReadOnly[str | None] diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py new file mode 100644 index 00000000000..83a42c235f9 --- /dev/null +++ b/litellm/types/litellm_params.py @@ -0,0 +1,364 @@ +"""LiteLLM-owned request kwargs declared as typed fields; types/utils.py splices these with the callback and pricing +models and KWARG_ARTIFACTS into all_litellm_params.""" + +from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence +from dataclasses import dataclass, field, fields, is_dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeAlias + +if TYPE_CHECKING: + import httpx + from aiohttp import ClientSession + from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.router_strategy.complexity_router.context_compaction import CompactionState + from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets + from litellm.types.caching import DynamicCacheControl + from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage + from litellm.types.proxy.litellm_pre_call_utils import SecretFields + from litellm.types.router import ConfigurableClientsideParamsCustomAuth, DeploymentTypedDict, RetryPolicy + from litellm.types.router_weights import RouterWeights + from litellm.types.utils import ModelResponse, ModelResponseStream, ProviderSpecificHeader + + ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient + ) + MockResponse: TypeAlias = ( + str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + ) + +RetryStrategy: TypeAlias = Literal["constant_retry", "exponential_backoff_retry"] +AgenticSurface: TypeAlias = Literal["chat_completions", "responses"] +RoutingStrategyName: TypeAlias = Literal[ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", +] + +TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" +ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" + +WIRE_NAME: Final = "wire_name" + + +def wire(name: str) -> Mapping[str, str]: + return MappingProxyType({WIRE_NAME: name}) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProviderConnection: + api_key: str | None = None + api_base: str | None = None + api_version: str | None = None + region_name: str | None = None + headers: Mapping[str, str] | None = None + provider_specific_header: "ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None" = None + client: "ProviderClient | None" = None + shared_session: "ClientSession | None" = None + ssl_verify: bool | str | None = None + request_timeout: float | None = None + force_timeout: float | None = None + stream_timeout: float | str | None = None + max_retries: int | None = None + tenant_id: str | None = None + client_id: str | None = None + client_secret: str | None = None + azure_username: str | None = None + azure_password: str | None = None + azure_scope: str | None = None + azure_ad_token_provider: Callable[[], str] | None = None + litellm_credential_name: str | None = None + configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None + use_xai_oauth: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class BedrockBatchConnection: + # Bedrock rejects these names in request bodies, so register them as LiteLLM-owned + aws_batch_role_arn: str | None = None + s3_bucket_name: str | None = None + s3_region_name: str | None = None + s3_endpoint_url: str | None = None + s3_output_bucket_name: str | None = None + s3_bucket_owner: str | None = None + s3_access_key_id: str | None = None + s3_secret_access_key: str | None = None + s3_encryption_key_id: str | None = None + bedrock_tags: Sequence[Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ConnectionSettings: + provider: ProviderConnection + bedrock_batch: BedrockBatchConnection + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DispatchOptions: + custom_llm_provider: str | None = None + azure: bool | None = None + use_litellm_proxy: bool | None = None + use_chat_completions_api: bool | None = None + use_in_pass_through: bool | None = None + allowed_openai_params: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RoutingOptions: + fallbacks: Sequence[str | Mapping[str, object]] | None = None + context_window_fallback_dict: Mapping[str, str] | None = None + num_retries: int | None = None + retry_policy: "RetryPolicy | Mapping[str, object] | None" = None + retry_strategy: RetryStrategy | None = None + routing_strategy: RoutingStrategyName | None = None + cooldown_time: float | None = None + allowed_model_region: str | None = None + enable_tag_filtering: bool | None = None + fastest_response: bool | None = None + provider_affinity_header: str | None = None + search_tool_name: str | None = None + model_list: "Sequence[DeploymentTypedDict] | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DeploymentOptions: + model_info: Mapping[str, object] | None = None + rpm: int | None = None + tpm: int | None = None + itpm: int | None = None + otpm: int | None = None + default_api_key_rpm_limit: int | None = None + default_api_key_tpm_limit: int | None = None + max_parallel_requests: int | None = None + weight: int | None = None + order: int | None = None + tag_regex: Sequence[str] | None = None + max_file_size_mb: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class SpecializedRouterOptions: + auto_router_config_path: str | None = None + auto_router_config: str | None = None + auto_router_default_model: str | None = None + auto_router_embedding_model: str | None = None + auto_router_max_input_chars: int | None = None + auto_router_routing_compression: str | None = None + auto_router_model_compression: str | None = None + complexity_router_config: Mapping[str, object] | None = None + complexity_router_default_model: str | None = None + adaptive_router_config: Mapping[str, object] | None = None + adaptive_router_default_model: str | None = None + quality_router_config: Mapping[str, object] | None = None + quality_router_default_model: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CachingOptions: + caching: bool | None = None + cache: "DynamicCacheControl | None" = None + ttl: float | None = None + enable_prompt_caching: bool | None = None + caching_groups: Sequence[Sequence[str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CostOptions: + cost_per_query: float | None = None + base_model: str | None = None + max_budget: float | None = None + budget_duration: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ObservabilityOptions: + id: str | None = None + metadata: MutableMapping[str, object] | None = None # mutable-ok: the router and logging write keys into it + litellm_metadata: MutableMapping[str, object] | None = None # mutable-ok: the proxy writes keys into it + tags: Sequence[str] | None = None + litellm_trace_id: str | None = None + litellm_session_id: str | None = None + litellm_request_debug: bool | None = None + logger_fn: Callable[[Mapping[str, object]], None] | None = None + verbose: bool | None = None + no_log: bool | None = field(default=None, metadata=wire("no-log")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopOptions: + max_agentic_loops: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GuardrailOptions: + guardrails: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PromptOptions: + prompt_id: str | None = None + prompt_variables: Mapping[str, object] | None = None + prompt_version: str | None = None + prompt_environment: str | None = None + prompt_label: str | None = None + litellm_system_prompt: str | None = None + custom_prompt_dict: Mapping[str, object] | None = None + roles: Mapping[str, object] | None = None + final_prompt_value: str | None = None + bos_token: str | None = None + eos_token: str | None = None + hf_model_name: str | None = None + supports_system_message: bool | None = None + ensure_alternating_roles: bool | None = None + user_continue_message: "ChatCompletionUserMessage | None" = None + assistant_continue_message: "ChatCompletionAssistantMessage | None" = None + disable_add_transform_inline_image_block: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ResponseOptions: + merge_reasoning_content_in_choices: bool | None = None + enable_json_schema_validation: bool | None = None + complete_response: bool | None = None + stream_chunk_size: int | None = None + keepalive_seconds: float | None = None + allow_client_keepalive_override: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class MockOptions: + mock_response: "MockResponse | None" = None + mock_timeout: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class LiteLLMOptions: + dispatch: DispatchOptions + routing: RoutingOptions + deployment: DeploymentOptions + specialized_routers: SpecializedRouterOptions + caching: CachingOptions + cost: CostOptions + observability: ObservabilityOptions + agentic_loop: AgenticLoopOptions + guardrails: GuardrailOptions + prompt: PromptOptions + response: ResponseOptions + mock: MockOptions + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CallState: + litellm_call_id: str | None = None + completion_call_id: str | None = None + model_alias_map: Mapping[str, str] | None = None + data_residency: str | None = None + litellm_logging_obj: "Logging | None" = None + preset_cache_key: str | None = None + cache_key: str | None = None + stream_response: "Mapping[str, ModelResponse] | None" = None + context_compaction_state: "CompactionState | None" = field(default=None, metadata=wire("_context_compaction_state")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopState: + depth: int | None = field(default=None, metadata=wire("_agentic_loop_depth")) + fingerprints: Sequence[str] | None = field(default=None, metadata=wire("_agentic_loop_fingerprints")) + api_surface: Literal["chat_completions", "responses"] | None = field( + default=None, metadata=wire("_agentic_loop_api_surface") + ) + code_interpreter_active: bool | None = field(default=None, metadata=wire("_code_interpreter_interception_active")) + code_interpreter_sandbox_key: str | None = field( + default=None, metadata=wire("_code_interpreter_interception_sandbox_key") + ) + code_interpreter_session_scoped: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_session_scoped") + ) + code_interpreter_converted_stream: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_converted_stream") + ) + websearch_emit_native_blocks: bool | None = field( + default=None, metadata=wire("_websearch_interception_emit_native_blocks") + ) + websearch_converted_stream: bool | None = field( + default=None, metadata=wire("_websearch_interception_converted_stream") + ) + headroom_converted_stream: bool | None = field( + default=None, metadata=wire("_headroom_interception_converted_stream") + ) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RouterState: + weights: "RouterWeights | None" = field(default=None, metadata=wire("_router_weights")) + fallback_depth: int | None = None + max_fallbacks: int | None = None + attempted_targets: "AttemptedFallbackTargets | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProxyRequestState: + proxy_server_request: Mapping[str, object] | None = None + secret_fields: "SecretFields | None" = None + trusted_callback_vars: Mapping[str, str] | None = field(default=None, metadata=wire(TRUSTED_CALLBACK_VARS_FIELD)) + addressed_response_id: str | None = field(default=None, metadata=wire(ADDRESSED_RESPONSE_ID_FIELD)) + strip_stream_usage: bool | None = field(default=None, metadata=wire("_litellm_strip_stream_usage")) + client_side_timeout: bool | None = None + model_file_id_mapping: Mapping[str, Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class EntrypointState: + acompletion: bool | None = None + aembedding: bool | None = None + aimg_generation: bool | None = None + atext_completion: bool | None = None + text_completion: bool | None = None + allm_passthrough_route: bool | None = None + async_call: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InternalState: + call: CallState + agentic_loop: AgenticLoopState + router: RouterState + proxy: ProxyRequestState + entrypoint: EntrypointState + + +KWARG_ARTIFACTS: Final[tuple[str, ...]] = ("self", "use_client", "model_config", "rust") + +LITELLM_OWNED_ROOTS: Final = (ConnectionSettings, LiteLLMOptions, InternalState) + + +def wire_names(owner: type) -> tuple[str, ...]: + return tuple(owned.metadata.get(WIRE_NAME, owned.name) for owned in fields(owner)) + + +def owned_wire_names(root: type) -> tuple[str, ...]: + def names() -> Iterator[str]: + for leaf in fields(root): + if not is_dataclass(leaf.type): + raise TypeError(f"{root.__name__}.{leaf.name} is not a dataclass leaf") + yield from wire_names(leaf.type) # pyright: ignore[reportArgumentType] # Field.type admits str + + return tuple(names()) + + +OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) +AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions)) +BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection) diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index 042df6f37fa..a818daf554d 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -393,7 +393,8 @@ class AnthropicMessagesSystemMessageParam(TypedDict, total=False): AllAnthropicMessageValues = AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam -# System is not a native Anthropic message role; only pass-through adapters use this union. +# role=system inside messages is accepted after a user turn on models flagged +# supports_mid_conversation_system; pass-through adapters and the chat translator both emit it. AllAnthropicPassThroughMessageValues: TypeAlias = ( AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam | AnthropicMessagesSystemMessageParam ) diff --git a/litellm/types/llms/vertex_ai_gemma.py b/litellm/types/llms/vertex_ai_gemma.py new file mode 100644 index 00000000000..f64f4d1d4fc --- /dev/null +++ b/litellm/types/llms/vertex_ai_gemma.py @@ -0,0 +1,10 @@ +from typing import Annotated, Literal + +from pydantic import BaseModel, ConfigDict, Field + + +class VertexGemmaContainerError(BaseModel): + model_config = ConfigDict(frozen=True) + object: Literal["error"] + message: str + code: Annotated[int, Field(ge=400, le=599)] diff --git a/litellm/types/router.py b/litellm/types/router.py index c0f724584fd..b72809f625f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -24,6 +24,7 @@ if TYPE_CHECKING: from .completion import CompletionRequest from .embedding import EmbeddingRequest +from .litellm_params import RoutingStrategyName from .llms.bedrock import AwsSessionTag from .llms.openai import OpenAIFileObject from .search import SearchProvider @@ -104,12 +105,7 @@ class RouterConfig(BaseModel): context_window_fallbacks: list | None = [] model_group_alias: dict[str, list[str]] | None = {} retry_after: int | None = 0 - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - ] = "simple-shuffle" + routing_strategy: RoutingStrategyName = "simple-shuffle" routing_groups: list[RoutingGroup] | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index caf88e5d517..3e306b48887 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -56,8 +56,15 @@ from litellm.types.llms.base import ( from litellm.types.mcp import MCPServerCostInfo from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers +from . import litellm_params as _litellm_params from .agents import LiteLLMSendMessageResponse from .guardrails import GuardrailEventHooks +from .litellm_params import ( + AGENTIC_LOOP_KWARG_NAMES, + BEDROCK_BATCH_KWARG_NAMES, + KWARG_ARTIFACTS, + OWNED_KWARG_NAMES, +) from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from .llms.base import HiddenParams from .llms.openai import ( @@ -3901,205 +3908,20 @@ def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]: ) -# Server-controlled fields that bound or drive an interceptor's agentic loop -# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed -# in all_litellm_params so they are treated as LiteLLM-level and excluded from -# get_non_default_completion_params; otherwise the OpenAI param builder sweeps -# any unrecognized top-level key into extra_body and leaks them to the provider. -# This is what lets the loop carry state across rerun calls without a provider -# scrubber. -agentic_loop_internal_litellm_params: Final = [ - "_agentic_loop_depth", - "_agentic_loop_fingerprints", - "_agentic_loop_api_surface", - "max_agentic_loops", - "_code_interpreter_interception_active", - "_code_interpreter_interception_sandbox_key", - "_code_interpreter_interception_session_scoped", - "_code_interpreter_interception_converted_stream", - "_websearch_interception_emit_native_blocks", - "_websearch_interception_converted_stream", - "_headroom_interception_converted_stream", +agentic_loop_internal_litellm_params: Final = list(AGENTIC_LOOP_KWARG_NAMES) # mutable-ok: public type stays a list + +bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES + +TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD +ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD + +all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it # mutable-ok: callers concat + *OWNED_KWARG_NAMES, + *KWARG_ARTIFACTS, + *StandardCallbackDynamicParams.__annotations__, + *CustomPricingLiteLLMParams.model_fields, ] -# Proxy-owned callback credentials, stamped from admin-configured team/key callback -# settings. Listed in all_litellm_params for the same reason as the agentic-loop -# fields above: an unrecognized top-level key is swept into extra_body and sent to -# the provider. -TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" - -ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" - -# Bedrock managed-batch deployment config, read from litellm_params by the batch and -# files transformations. Listed for the same reason as the fields above: these sit on -# a deployment that also serves chat, so leaking them into extra_body makes Bedrock -# reject every non-batch request to that deployment. -bedrock_batch_litellm_params: Final = ( - "aws_batch_role_arn", - "s3_bucket_name", - "s3_region_name", - "s3_endpoint_url", - "s3_output_bucket_name", - "s3_bucket_owner", - "s3_access_key_id", - "s3_secret_access_key", - "s3_encryption_key_id", - "bedrock_tags", -) - -all_litellm_params = ( - agentic_loop_internal_litellm_params - + [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params] - + [ - "_context_compaction_state", - "metadata", - "litellm_metadata", - "keepalive_seconds", - "allow_client_keepalive_override", - "litellm_trace_id", - "litellm_request_debug", - "guardrails", - "tags", - "acompletion", - "aimg_generation", - "atext_completion", - "text_completion", - "caching", - "mock_response", - "mock_timeout", - "disable_add_transform_inline_image_block", - "api_key", - "api_version", - "prompt_id", - "prompt_variables", - "litellm_system_prompt", - "provider_specific_header", - "prompt_version", - "prompt_environment", - "api_base", - "force_timeout", - "logger_fn", - "verbose", - "custom_llm_provider", - "model_file_id_mapping", - "litellm_logging_obj", - "litellm_call_id", - "completion_call_id", - "model_alias_map", - "custom_prompt_dict", - "stream_response", - "cost_per_query", - "ssl_verify", - "data_residency", - "async_call", - "aembedding", - "allm_passthrough_route", - "_litellm_strip_stream_usage", - "use_client", - "id", - "fallbacks", - "routing_strategy", - "_router_weights", - "azure", - "headers", - "model_list", - "num_retries", - "context_window_fallback_dict", - "retry_policy", - "retry_strategy", - "roles", - "final_prompt_value", - "bos_token", - "eos_token", - "request_timeout", - "client_side_timeout", - "complete_response", - "self", - "client", - "rpm", - "tpm", - "default_api_key_rpm_limit", - "default_api_key_tpm_limit", - "itpm", - "otpm", - "max_parallel_requests", - "input_cost_per_token", - "output_cost_per_token", - "input_cost_per_second", - "output_cost_per_second", - "hf_model_name", - "model_info", - "proxy_server_request", - "secret_fields", - "preset_cache_key", - "caching_groups", - "ttl", - "cache", - "enable_prompt_caching", - "no-log", - "base_model", - "stream_timeout", - "stream_chunk_size", - "supports_system_message", - "region_name", - "allowed_model_region", - "model_config", - "fastest_response", - "cooldown_time", - "cache_key", - "max_retries", - "azure_ad_token_provider", - "tenant_id", - "client_id", - "azure_username", - "azure_password", - "azure_scope", - "client_secret", - "user_continue_message", - "configurable_clientside_auth_params", - "weight", - "ensure_alternating_roles", - "assistant_continue_message", - "user_continue_message", - "fallback_depth", - "max_fallbacks", - "attempted_targets", - "max_budget", - "budget_duration", - "use_in_pass_through", - "merge_reasoning_content_in_choices", - "litellm_credential_name", - "allowed_openai_params", - "litellm_session_id", - "provider_affinity_header", - "use_litellm_proxy", - "use_chat_completions_api", - "rust", - "prompt_label", - "shared_session", - "search_tool_name", - "order", - "enable_tag_filtering", - "enable_json_schema_validation", - "use_xai_oauth", - "auto_router_config_path", - "auto_router_config", - "auto_router_default_model", - "auto_router_embedding_model", - "auto_router_max_input_chars", - "auto_router_routing_compression", - "auto_router_model_compression", - "complexity_router_config", - "complexity_router_default_model", - "adaptive_router_config", - "adaptive_router_default_model", - "quality_router_config", - "quality_router_default_model", - ] - + list(StandardCallbackDynamicParams.__annotations__.keys()) - + list(CustomPricingLiteLLMParams.model_fields.keys()) -) - class KeyGenerationConfig(TypedDict, total=False): required_params: list[str] # specify params that must be present in the key generation request @@ -4194,6 +4016,7 @@ class LlmProviders(str, Enum): NVIDIA_RIVA = "nvidia_riva" SONIOX = "soniox" CEREBRAS = "cerebras" + NADIR = "nadir" AI21_CHAT = "ai21_chat" VOLCENGINE = "volcengine" CODESTRAL = "codestral" diff --git a/litellm/utils.py b/litellm/utils.py index 9ca19f61862..4ea0769ea11 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4828,6 +4828,13 @@ def get_optional_params( model=model, drop_params=bool(drop_params), ) + elif custom_llm_provider == "nadir": + optional_params = litellm.NadirConfig().map_openai_params( # rebind-ok: same optional_params rebinding every sibling provider branch does + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=bool(drop_params), + ) elif custom_llm_provider == "xai": optional_params = litellm.XAIChatConfig().map_openai_params( model=model, @@ -5690,6 +5697,8 @@ def _check_provider_match(model_info: dict, custom_llm_provider: str | None) -> elif custom_llm_provider == "github": # Allow github/ aliases to reuse existing provider metadata. return True + elif custom_llm_provider == "nadir": + return True else: return False @@ -6815,6 +6824,11 @@ def validate_environment( keys_in_environment = True else: missing_keys.append("CEREBRAS_API_KEY") + elif custom_llm_provider == "nadir": + if "NADIR_API_KEY" in os.environ: + keys_in_environment = True # rebind-ok: same flag rebinding every sibling provider branch does + else: + missing_keys.append("NADIR_API_KEY") elif custom_llm_provider == "baseten": if "BASETEN_API_KEY" in os.environ: keys_in_environment = True @@ -8473,6 +8487,7 @@ class ProviderConfigManager: LlmProviders.HUGGINGFACE: (lambda: litellm.HuggingFaceChatConfig(), False), LlmProviders.TOGETHER_AI: (lambda: litellm.TogetherAIChatConfig(), False), LlmProviders.OPENROUTER: (lambda: litellm.OpenrouterConfig(), False), + LlmProviders.NADIR: (lambda: litellm.NadirConfig(), False), LlmProviders.VERCEL_AI_GATEWAY: ( lambda: litellm.VercelAIGatewayConfig(), False, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index aa79cadae2a..5a207dc4c02 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3301,12 +3301,14 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "azure/ada": { + "deprecation_date": "2028-02-09", "input_cost_per_token": 1e-07, "litellm_provider": "azure", "max_input_tokens": 8191, "max_tokens": 8191, "mode": "embedding", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule" }, "azure/codex-mini": { "cache_read_input_token_cost": 3.75e-07, @@ -3340,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3347,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5151,6 +5155,7 @@ "supports_tool_choice": true }, "azure/gpt-35-turbo-16k": { + "deprecation_date": "2025-04-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 16385, @@ -5158,9 +5163,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-35-turbo-16k-0613": { + "deprecation_date": "2025-04-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 16385, @@ -5168,6 +5175,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 4e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5188,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5195,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5211,6 +5221,7 @@ "supports_tool_choice": true }, "azure/gpt-4-0613": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5218,6 +5229,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -5234,6 +5246,7 @@ "supports_tool_choice": true }, "azure/gpt-4-32k": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 6e-05, "litellm_provider": "azure", "max_input_tokens": 32768, @@ -5241,9 +5254,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.00012, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-4-32k-0613": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 6e-05, "litellm_provider": "azure", "max_input_tokens": 32768, @@ -5251,6 +5266,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 0.00012, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_tool_choice": true }, "azure/gpt-4-turbo": { @@ -5511,6 +5527,7 @@ }, "azure/gpt-4.5-preview": { "cache_read_input_token_cost": 3.75e-05, + "deprecation_date": "2025-07-14", "input_cost_per_token": 7.5e-05, "input_cost_per_token_batches": 3.75e-05, "litellm_provider": "azure", @@ -5520,6 +5537,7 @@ "mode": "chat", "output_cost_per_token": 0.00015, "output_cost_per_token_batches": 7.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/legacy-models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -15797,6 +15815,7 @@ "supports_tool_choice": true }, "computer-use-preview": { + "deprecation_date": "2026-07-23", "input_cost_per_token": 3e-06, "litellm_provider": "openai", "max_input_tokens": 8192, @@ -22838,6 +22857,37 @@ "/v1/images/generations" ] }, + "fal_ai/fal-ai/nano-banana-2": { + "litellm_provider": "fal_ai", + "metadata": { + "comment": "priced by the request's resolution field (0.5K, 1K default, 2K, 4K); the web search and high thinking surcharges are not modeled" + }, + "mode": "image_generation", + "output_cost_per_image": 0.08, + "output_cost_per_image_0.5K": 0.06, + "output_cost_per_image_1K": 0.08, + "output_cost_per_image_2K": 0.12, + "output_cost_per_image_4K": 0.16, + "source": "https://fal.ai/models/fal-ai/nano-banana-2", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, + "fal_ai/fal-ai/nano-banana-pro": { + "litellm_provider": "fal_ai", + "metadata": { + "comment": "priced by the request's resolution field (1K default, 2K, 4K); the web search surcharge is not modeled" + }, + "mode": "image_generation", + "output_cost_per_image": 0.15, + "output_cost_per_image_1K": 0.15, + "output_cost_per_image_2K": 0.15, + "output_cost_per_image_4K": 0.3, + "source": "https://fal.ai/models/fal-ai/nano-banana-pro", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "fal_ai/openai/gpt-image-2": { "litellm_provider": "fal_ai", "metadata": { @@ -28748,6 +28798,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 1e-05, "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -29823,6 +29874,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "rpm": 10000, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_modalities": [ @@ -32113,6 +32165,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-1.5": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -32133,6 +32186,7 @@ "supports_pdf_input": true }, "gpt-image-1.5-2025-12-16": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -32153,6 +32207,7 @@ "supports_pdf_input": true }, "gpt-image-2": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.25e-07, "input_cost_per_token": 5e-06, @@ -32171,12 +32226,14 @@ "supports_pdf_input": true }, "gpt-image-2-2026-04-21": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", "mode": "image_generation", "input_cost_per_image_token": 8e-06, "output_cost_per_image_token": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32185,6 +32242,7 @@ "supports_pdf_input": true }, "gpt-image-2.5-flare": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32200,6 +32258,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-flare-2026-09-08": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32215,6 +32274,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32230,6 +32290,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-2.5-sunburst-2026-09-08": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 5e-06, "litellm_provider": "openai", @@ -32657,6 +32718,7 @@ }, "gpt-5-codex": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 272000, @@ -32740,6 +32802,7 @@ }, "gpt-5.1-codex-mini": { "cache_read_input_token_cost": 2.5e-08, + "deprecation_date": "2026-07-23", "input_cost_per_token": 2.5e-07, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32770,6 +32833,7 @@ }, "gpt-5.1-codex-max": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32801,6 +32865,7 @@ }, "gpt-5.1-codex": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 400000, @@ -32832,6 +32897,7 @@ }, "gpt-5.1-chat-latest": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -32968,6 +33034,7 @@ }, "gpt-5.2-codex": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-07-23", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 272000, @@ -32999,6 +33066,7 @@ }, "gpt-5.2-chat-latest": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-08-10", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -34822,6 +34890,7 @@ }, "gpt-5.3-chat-latest": { "cache_read_input_token_cost": 1.75e-07, + "deprecation_date": "2026-08-10", "input_cost_per_token": 1.75e-06, "litellm_provider": "openai", "max_input_tokens": 128000, @@ -35054,6 +35123,7 @@ "supports_minimal_reasoning_effort": true }, "gpt-image-1": { + "cache_read_input_image_token_cost": 2.5e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-10-23", @@ -35071,6 +35141,7 @@ ] }, "gpt-image-1-mini": { + "cache_read_input_image_token_cost": 2.5e-07, "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_batches": 1e-07, "deprecation_date": "2026-12-01", @@ -35090,6 +35161,7 @@ "gpt-realtime": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 3.2e-05, @@ -35125,6 +35197,7 @@ "gpt-realtime-1.5": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35159,6 +35232,7 @@ "gpt-realtime-2": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35193,6 +35267,7 @@ "gpt-realtime-2.1": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, @@ -35229,6 +35304,7 @@ "gpt-realtime-2.1-mini": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, @@ -35265,6 +35341,7 @@ "gpt-realtime-mini": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1e-05, @@ -35300,6 +35377,7 @@ "gpt-realtime-2025-08-28": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, + "cache_read_input_image_token_cost": 5e-07, "cache_read_input_token_cost": 4e-07, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 3.2e-05, @@ -55976,6 +56054,7 @@ "gpt-realtime-mini-2025-12-15": { "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, + "cache_read_input_image_token_cost": 8e-08, "cache_read_input_token_cost": 6e-08, "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, @@ -56068,6 +56147,7 @@ ] }, "chatgpt-image-latest": { + "cache_read_input_image_token_cost": 2e-06, "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_batches": 6.3e-07, "deprecation_date": "2026-12-01", @@ -56451,6 +56531,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -56519,6 +56600,7 @@ "mode": "audio_speech", "output_cost_per_audio_token": 1e-05, "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" @@ -68835,6 +68917,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/NousResearch/Nous-Hermes-2-Mixtral-8x7B-DPO": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68844,6 +68927,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/QwQ-32B": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68853,6 +68937,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2-72B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 9e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68862,6 +68947,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2-VL-72B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68871,6 +68957,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-72B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68880,6 +68967,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-Coder-32B-Instruct": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 16384, @@ -68889,6 +68977,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-VL-72B-Instruct": { + "deprecation_date": "2026-01-05", "input_cost_per_token": 1.95e-06, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -68898,6 +68987,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "deprecation_date": "2026-06-04", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68907,6 +68997,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-Next-FP8": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68916,6 +69007,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68925,6 +69017,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68934,6 +69027,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-32B-Instruct": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68943,6 +69037,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-8B-Instruct": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68953,6 +69048,7 @@ }, "together_ai/Qwen/Qwen3.5-397B-A17B": { "cache_read_input_token_cost": 3.5e-07, + "deprecation_date": "2026-06-29", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -68962,6 +69058,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Llama-70B": { + "deprecation_date": "2025-12-23", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68971,6 +69068,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68980,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-14B": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 1.6e-06, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68989,6 +69088,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-V3.1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -68998,6 +69098,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/deepseek-coder-33b-instruct": { + "deprecation_date": "2024-08-22", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 16384, @@ -69007,6 +69108,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-2-27b-it": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 8e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69016,6 +69118,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69025,6 +69128,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Llama-3-8b-chat-hf": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69034,6 +69138,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Llama-4-Scout-17B-16E-Instruct": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -69043,6 +69148,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3-70B-Instruct-Turbo": { + "deprecation_date": "2025-12-23", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69052,6 +69158,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3-8B-Instruct": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 8192, @@ -69061,6 +69168,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69070,6 +69178,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69079,6 +69188,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-7B-Instruct-v0.1": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69088,6 +69198,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69097,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69107,6 +69219,7 @@ }, "together_ai/moonshotai/Kimi-K2.6": { "cache_read_input_token_cost": 2e-07, + "deprecation_date": "2026-08-19", "input_cost_per_token": 1.2e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69117,6 +69230,7 @@ }, "together_ai/moonshotai/Kimi-K2.7-Code": { "cache_read_input_token_cost": 1.9e-07, + "deprecation_date": "2026-08-27", "input_cost_per_token": 9.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69126,6 +69240,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/nvidia/Llama-3.1-Nemotron-70B-Instruct-HF": { + "deprecation_date": "2025-08-28", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "max_input_tokens": 32768, @@ -69145,6 +69260,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69154,6 +69270,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-4.5-Air-FP8": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 131072, @@ -69163,6 +69280,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-4.7": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 202752, @@ -69172,6 +69290,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5": { + "deprecation_date": "2026-06-22", "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", "max_input_tokens": 202752, @@ -69182,6 +69301,7 @@ }, "together_ai/zai-org/GLM-5.1": { "cache_read_input_token_cost": 2.6e-07, + "deprecation_date": "2026-07-10", "input_cost_per_token": 1.4e-06, "litellm_provider": "together_ai", "max_input_tokens": 202752, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 395b2db1137..35624045fdf 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -624,6 +624,10 @@ "type": "number", "minimum": 0 }, + "output_cost_per_image_0.5K": { + "type": "number", + "minimum": 0 + }, "output_cost_per_image_1024": { "type": "number", "minimum": 0 @@ -632,6 +636,18 @@ "type": "number", "minimum": 0 }, + "output_cost_per_image_1K": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_image_2K": { + "type": "number", + "minimum": 0 + }, + "output_cost_per_image_4K": { + "type": "number", + "minimum": 0 + }, "output_cost_per_image_512": { "type": "number", "minimum": 0 diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index b8d1621cde3..e6cb0592a15 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -457,6 +457,24 @@ "interactions": true } }, + "nadir": { + "display_name": "Nadir (`nadir`)", + "url": "https://docs.litellm.ai/docs/providers/nadir", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false, + "interactions": false + } + }, "cerebras": { "display_name": "Cerebras (`cerebras`)", "url": "https://docs.litellm.ai/docs/providers/cerebras", diff --git a/pyproject.toml b/pyproject.toml index ba72378989a..f2364b5e77b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,8 +75,8 @@ proxy = [ "mcp>=2.2.0,<3", "httpx2>=2.5.0,<3", "pydantic>=2.12.0,<3", - "litellm-proxy-extras==0.4.101", - "litellm-enterprise==0.1.70", + "litellm-proxy-extras==0.4.102", + "litellm-enterprise==0.1.71", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", @@ -171,11 +171,11 @@ proxy-runtime = [ "anthropic[vertex]>=0.84.0,<1.0", "grpcio==1.78.0", "prometheus-client>=0.20.0,<1.0", - "langfuse>=2.59.7,<3.0", - "opentelemetry-api==1.28.0", - "opentelemetry-sdk==1.28.0", - "opentelemetry-exporter-otlp==1.28.0", - "opentelemetry-instrumentation-fastapi==0.49b0", + "langfuse>=4.7,<5.0", + "opentelemetry-api==1.33.1", + "opentelemetry-sdk==1.33.1", + "opentelemetry-exporter-otlp==1.33.1", + "opentelemetry-instrumentation-fastapi==0.54b1", "ddtrace>=4.8.2,<5.0", "sentry-sdk>=2.21.0,<3.0", "mangum>=0.17.0,<1.0", @@ -222,11 +222,11 @@ dev = [ "types-PyYAML==6.0.12.20250915", "botocore-stubs==1.43.14", "types-boto3[bedrock,bedrock-agent,bedrock-runtime,kms,s3,sagemaker-runtime,sts]==1.43.30", - "opentelemetry-api==1.28.0", - "opentelemetry-sdk==1.28.0", - "opentelemetry-exporter-otlp==1.28.0", - "opentelemetry-instrumentation-fastapi==0.49b0", - "langfuse==2.59.7", + "opentelemetry-api==1.33.1", + "opentelemetry-sdk==1.33.1", + "opentelemetry-exporter-otlp==1.33.1", + "opentelemetry-instrumentation-fastapi==0.54b1", + "langfuse>=4.7,<5.0", "fastapi-offline==1.7.6", "fakeredis==2.34.1", "pytest-rerunfailures==15.1", @@ -249,10 +249,10 @@ proxy-dev = [ "prisma==0.11.0", "hypercorn==0.17.3", "prometheus-client==0.20.0", - "opentelemetry-api==1.28.0", - "opentelemetry-sdk==1.28.0", - "opentelemetry-exporter-otlp==1.28.0", - "opentelemetry-instrumentation-fastapi==0.49b0", + "opentelemetry-api==1.33.1", + "opentelemetry-sdk==1.33.1", + "opentelemetry-exporter-otlp==1.33.1", + "opentelemetry-instrumentation-fastapi==0.54b1", "azure-identity==1.25.2", "a2a-sdk==1.1.0", ] @@ -272,7 +272,7 @@ ci = [ "lunary==1.4.36; python_version == '3.10'", "lunary==1.4.37; python_version >= '3.11'", "logfire==4.6.0", - "traceloop-sdk==0.33.12", + "traceloop-sdk==0.34.0", "detect-secrets==1.5.0", "PyGithub==2.8.1", "aiodynamo==24.7", diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py index 2bb65072ad4..378b8e0876a 100644 --- a/scripts/check_type_discipline.py +++ b/scripts/check_type_discipline.py @@ -40,7 +40,8 @@ LIT003 noqa suppression without rule codes or without a reason. LIT004 pyright/mypy ignore without bracketed codes or without a reason. Required shape: `# pyright: ignore[reportArgumentType] # ` LIT005 A `# mutable-ok` / `# cast-ok` / `# guard-ok` / `# kwargs-ok` / - `# rebind-ok` / `# writable-ok` suppression without a reason. + `# rebind-ok` / `# writable-ok` / `# comprehension-ok` suppression + without a reason. LIT006 `cast(...)` call. typing.cast is an unchecked assertion (the moral equivalent of TypeScript's `as`); it lies to the type checker with zero runtime guarantee. Validate into a concrete frozen type at the boundary instead. @@ -103,6 +104,15 @@ LIT013 A `# -ok: ` suppression on a line where none of the rules that token suppresses fires. Like ruff's RUF100: a marker that suppresses nothing rots in place and hides real violations that land on the line later. Delete it. +LIT014 Comprehension with more than one `for` clause or more than one `if` clause, + in any of the four forms (list, set, dict, generator expression). Stacked + `for`s and `if`s read as nested loops and guards squashed onto one line; + split the comprehension into a helper generator, a named intermediate, or + a plain loop instead. A comprehension nested inside another's element or + iterable is its own node and is judged separately. Suppress with + `# comprehension-ok: ` on any line the comprehension spans. The + marker belongs to the innermost violating comprehension spanning that + line, and also to any single-line violating comprehension on that line. LIT000 Setup failure: a target file could not be read, or contains a syntax error. Reported as a violation rather than crashing the run. @@ -206,6 +216,7 @@ GUARD_OK_RE = re.compile(r"#\s*guard-ok(?::\s*(?P.*))?") KWARGS_OK_RE = re.compile(r"#\s*kwargs-ok(?::\s*(?P.*))?") REBIND_OK_RE = re.compile(r"#\s*rebind-ok(?::\s*(?P.*))?") WRITABLE_OK_RE = re.compile(r"#\s*writable-ok(?::\s*(?P.*))?") +COMPREHENSION_OK_RE = re.compile(r"#\s*comprehension-ok(?::\s*(?P.*))?") @dataclass(frozen=True, slots=True) class _OkToken: @@ -224,6 +235,7 @@ OK_SUPPRESSIONS: Final[tuple[_OkToken, ...]] = ( _OkToken("kwargs-ok", KWARGS_OK_RE, frozenset(("LIT008",))), _OkToken("rebind-ok", REBIND_OK_RE, frozenset(("LIT010", "LIT011"))), _OkToken("writable-ok", WRITABLE_OK_RE, frozenset(("LIT012",))), + _OkToken("comprehension-ok", COMPREHENSION_OK_RE, frozenset(("LIT014",))), ) @@ -1035,6 +1047,83 @@ def iter_typeddict_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: ) +# --------------------------------------------------------------------------- # +# Stacked comprehension clauses (LIT014) +# --------------------------------------------------------------------------- # + +COMPREHENSION_NODES = (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp) + + +def _span(node: ast.expr) -> range: + return range(node.lineno, (node.end_lineno or node.lineno) + 1) + + +def _clause_counts(node: ast.expr) -> tuple[int, int]: + return ( + len(node.generators), + sum(len(g.ifs) for g in node.generators), + ) + + +def _violates(node: ast.expr) -> bool: + for_count, if_count = _clause_counts(node) + return for_count > 1 or if_count > 1 + + +def _comprehension_owners(tree: ast.AST, ok_lines: frozenset[int]) -> Mapping[int, int]: + """id(node) -> marker line for each `# comprehension-ok` line's owner. + + Only violating comprehensions own markers. Each marker belongs to the + innermost violating comprehension whose span contains it (line span first, + column width breaks ties) plus every violating comprehension whose whole + span is that single line, so a comment inside a nested comprehension never + silences a multi-line enclosing one and a violation sharing its only line + can still be suppressed. + """ + violating: Final = tuple( + n for n in ast.walk(tree) if isinstance(n, COMPREHENSION_NODES) and _violates(n) + ) + + def nesting_key(node: ast.expr) -> tuple[int, int]: + return (len(_span(node)), (node.end_col_offset or node.col_offset) - node.col_offset) + + def owners(line: int) -> tuple[ast.expr, ...]: + containing: Final = tuple(n for n in violating if line in _span(n)) + innermost: Final = min(containing, key=nesting_key, default=None) + single_line: Final = tuple(n for n in violating if len(_span(n)) == 1 and n.lineno == line) + return (*single_line, *(() if innermost is None else (innermost,))) + + return MappingProxyType({id(o): line for line in ok_lines for o in owners(line)}) + + +def iter_comprehension_violations( + path: Path, tree: ast.AST, ok_lines: frozenset[int] +) -> Iterator[tuple[Violation, bool]]: + """(violation, owned) pairs for every violating comprehension. + + An owned comprehension reports at its marker's line so apply_suppressions + drops it and counts the marker as used; an unowned one reports at its own + line and is kept verbatim, since a marker suppresses only its owner even + when another violation shares that line. + """ + owners: Final = _comprehension_owners(tree, ok_lines) + for node in ast.walk(tree): + if not isinstance(node, COMPREHENSION_NODES) or not _violates(node): + continue + for_count, if_count = _clause_counts(node) + yield ( + Violation( + path, + owners.get(id(node), node.lineno), + "LIT014", + f"comprehension with {for_count} `for` clauses and {if_count} `if` clauses: " + f"at most one of each is allowed. Split it into a helper generator, a named " + f"intermediate, or a plain loop (suppress: `# comprehension-ok: `)", + ), + id(node) in owners, + ) + + # --------------------------------------------------------------------------- # # Suppression application and unused suppressions (LIT013) # --------------------------------------------------------------------------- # @@ -1087,8 +1176,13 @@ def check_file(path: Path) -> tuple[Violation, ...]: except SyntaxError as exc: return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}")) + comprehension_violations: Final = tuple( + iter_comprehension_violations(path, tree, suppressions["comprehension-ok"]) + ) + return ( *violations, + *(v for v, owned in comprehension_violations if not owned), *apply_suppressions( path, ( @@ -1099,6 +1193,7 @@ def check_file(path: Path) -> tuple[Violation, ...]: *iter_final_violations(path, tree), *iter_param_violations(path, tree), *iter_typeddict_violations(path, tree), + *(v for v, owned in comprehension_violations if owned), ), suppressions, ), diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py index 4ba1a2ea393..5acaf3994f7 100644 --- a/scripts/type_discipline_gate.py +++ b/scripts/type_discipline_gate.py @@ -15,8 +15,12 @@ without codes or reason), LIT006 (cast), LIT008 (`**kwargs`), LIT009 (inert (assignment without a Final declaration; suppress deliberate rebinding with `# rebind-ok: `), LIT011 (parameter rebinding or in-place mutation), and LIT012 (TypedDict field without a `ReadOnly[...]` qualifier; suppress with -`# writable-ok: `) carry limits at or above their current count to -ratchet down; LIT005 (`*-ok` suppression without a reason) is frozen at limit 0 +`# writable-ok: `), and LIT014 (comprehension with more than one `for` +or `if` clause; suppress with `# comprehension-ok: ` on a spanned +line, which belongs to the innermost violating comprehension spanning it and +to any single-line violating comprehension on that line) carry limits at +or above their current count to ratchet down; LIT005 (`*-ok` suppression +without a reason) is frozen at limit 0 so any net-new reasonless suppression trips the gate; LIT013 (`*-ok` suppression that suppresses nothing) is frozen at 0 for the same reason; and LIT007 (TypeGuard/TypeIs) is a hard zero. @@ -198,7 +202,8 @@ def cmd_check(base: str) -> None: "Remove the new violations, give each a reason (`# noqa: XXX # `, " "`# pyright: ignore[rule] # `, `# mutable-ok: `, " "`# cast-ok: `, `# guard-ok: `, `# kwargs-ok: `, " - "`# rebind-ok: `, `# writable-ok: `), or remove an equal " + "`# rebind-ok: `, `# writable-ok: `, " + "`# comprehension-ok: `), or remove an equal " "number elsewhere; the ceiling " "is the limit in type-discipline-budget.json." ) 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 d2ec3f1286e..7c247ae3303 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/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 7d3df57bb01..7bdf5b9c0e0 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -24,6 +24,8 @@ - {id: llm.chat_completions.anthropic.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude prompt caching"} - {id: llm.chat_completions.anthropic.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: anthropic, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude extended thinking"} - {id: llm.chat_completions.anthropic.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: anthropic, capability: structured_output, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude response_schema"} +- {id: llm.chat_completions.anthropic.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/anthropic/chat/transformation.py", rationale: "OpenAI-format chat to first-party Anthropic: flagged Claude 4.8+/5 must keep a mid-conversation role system reminder in messages; hoisting it into the top-level system field mutates the cached prefix and re-bills the conversation at cache-write pricing (#36559)", fail_before_fix: proven} +- {id: llm.chat_completions.anthropic.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/anthropic/chat/transformation.py", rationale: "OpenAI-format chat to first-party Anthropic: Claude <= 4.7 and Haiku 4.5 reject role system inside messages, so unflagged models must convert a mid-conversation reminder to a user turn in place (hoisting collapses the prompt cache) and still answer (#36559)", fail_before_fix: proven} - {id: llm.chat_completions.bedrock_converse.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Bedrock Converse unified"} - {id: llm.chat_completions.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming over Converse"} - {id: llm.chat_completions.bedrock_converse.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Converse function_calling; AWS adoption"} @@ -35,6 +37,8 @@ - {id: llm.chat_completions.bedrock_converse.batch_deployment.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: batch_deployment, streaming: nonstream, assertions: [works], source: "types/utils.py bedrock_batch_litellm_params", rationale: "A deployment carrying the documented batch-only S3 keys (s3_access_key_id, s3_secret_access_key, s3_encryption_key_id) must still serve ordinary chat; unregistered keys fall into optional_params and are forwarded as additionalModelRequestFields, which Bedrock 400s and which puts the S3 secret in the request body and debug log (LIT-8290)", fail_before_fix: proven} - {id: llm.chat_completions.bedrock_invoke.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "Regional inference-profile ids (us.anthropic.*) over the invoke route, the deployment shape behind a customer timeout report on v1.90.0"} - {id: llm.chat_completions.bedrock_invoke.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming with regional inference-profile ids over the invoke route"} +- {id: llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/anthropic/chat/transformation.py", rationale: "OpenAI-format chat to Bedrock Invoke builds the Anthropic request through AnthropicConfig.transform_request, so flagged Claude 4.8+/5 must keep a mid-conversation role system reminder in messages; hoisting mutates the cached prefix and collapses the prompt cache (#36559)", fail_before_fix: proven} +- {id: llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/anthropic/chat/transformation.py", rationale: "OpenAI-format chat to Bedrock Invoke: Claude <= 4.7 and Haiku 4.5 reject role system inside messages, so unflagged models must convert a mid-conversation reminder to a user turn in place (hoisting collapses the prompt cache) and still answer (#36559)", fail_before_fix: proven} - {id: llm.chat_completions.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Vertex AI"} - {id: llm.chat_completions.gemini.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini OpenAI-compatible chat translation"} - {id: llm.chat_completions.gemini.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini chat cost lands in SpendLogs"} @@ -81,6 +85,7 @@ - {id: llm.responses.vertex.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Vertex"} - {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"} - {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"} +- {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven} - {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"} - {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"} - {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"} diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index b7af8b6e9ec..0a2d4b6a10d 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -49,6 +49,7 @@ - {id: quota_management.spend_tracking.surface_consistency.matches_every_surface, module: quota_management, tier: P1, behavior: spend_tracking, variant: surface_consistency, assertions: [matches_every_surface], exercised_on: [chat_completions], source: "proxy/db/db_spend_update_writer.py", rationale: "One priced request lands the same response_cost on the spend log row, /key/info, /team/info, the usage export's /user/daily/activity/aggregated row, and the litellm_spend_metric Prometheus sample; each is a separate writer, so a rounding, dropped, or double-counted write on one drifts it from the rest (LIT-3620, LIT-5045)"} - {id: quota_management.spend_tracking.tags.attributes_spend, module: quota_management, tier: P1, behavior: spend_tracking, variant: tags, assertions: [attributes_spend], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "Request tags round-trip to spend rows and tag rollups match tagged logs"} - {id: quota_management.spend_tracking.end_user.attributes_spend, module: quota_management, tier: P1, behavior: spend_tracking, variant: end_user, assertions: [attributes_spend], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "user= attribution lands the end-user id on the spend row"} +- {id: quota_management.spend_tracking.end_user.attributes_responses_header, module: quota_management, tier: P1, behavior: spend_tracking, variant: end_user, assertions: [attributes_responses_header], exercised_on: [responses], source: "proxy/auth/auth_utils.py", rationale: "A /v1/responses call carrying x-litellm-customer-id or x-litellm-end-user-id plus x-litellm-tags, the headers Codex CLI attaches through its config.toml http_headers because it has no body field for the end user, lands the end user and the tags on a costed aresponses spend row whose spend the customer's /customer/info total matches (LIT-8575)"} - {id: quota_management.spend_tracking.per_model.writes_own_rows, module: quota_management, tier: P2, behavior: spend_tracking, variant: per_model, assertions: [writes_own_rows], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "Each model on a shared key gets its own spend row"} - {id: quota_management.spend_tracking.failure.writes_failure_row, module: quota_management, tier: P1, behavior: spend_tracking, variant: failure, assertions: [writes_failure_row], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_log_error_logger.py", rationale: "A failed call writes a failure-status spend row"} - {id: quota_management.spend_tracking.failure.writes_normalized_error, module: quota_management, tier: P1, behavior: spend_tracking, variant: failure, assertions: [writes_normalized_error], exercised_on: [chat_completions], source: "litellm_core_utils/error_normalization.py", rationale: "Failure rows carry a stable metadata.error_information.normalized_error key next to the unchanged error_message, so two upstream auth failures with different provider wording share one cluster key a dashboard can group by"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index e5626144fad..5fd19212ab7 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -66,6 +66,7 @@ LlmCapability = Literal[ "basic", "batch_deployment", "blank_s3_env", + "code_interpreter", "count_tokens", "govcloud_partition", "split_s3_credentials", diff --git a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md index 44d6e79122e..a18c81fa01d 100644 --- a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md +++ b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md @@ -48,7 +48,8 @@ most likely to silently break and the one a mock can't prove works. |----------|---------------|-----------|------------|-------------|--------| | Chat | live (spend suite) | live (spend suite) | gap | live | partial | | Embeddings | live (spend suite) | n/a | n/a | live | covered | -| Responses / image / audio / rerank / realtime | - | - | - | - | gap | +| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial | +| Image / audio / rerank / realtime | - | - | - | - | gap | ## This suite's files @@ -61,6 +62,7 @@ most likely to silently break and the one a mock can't prove works. | `test_anthropic_passthrough_streaming_logs_cost` | anthropic native, stream, cost | | `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost | | `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost | +| `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key | Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is added at runtime instead of declared in the gateway config: the test POSTs `/model/new` diff --git a/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py b/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py new file mode 100644 index 00000000000..480225b502e --- /dev/null +++ b/tests/e2e/llm_translation/test_chat_mid_conversation_system_e2e.py @@ -0,0 +1,322 @@ +"""Live e2e: mid-conversation ``role: "system"`` handling on the OpenAI-format +/v1/chat/completions path is model-aware for first-party Anthropic and Bedrock +Invoke, both of which build the Anthropic request through +``AnthropicConfig.transform_request`` (#36559). + +Only the leading run of system messages becomes the top-level ``system`` +parameter. A ``role: "system"`` entry that appears later in ``messages`` used to +be hoisted into that same field, which rewrote the cached prefix and re-billed +the whole conversation at cache-write pricing on every reminder. Models flagged +``supports_mid_conversation_system`` in the cost map (Claude 4.8+ and the 5 +family) must keep the reminder in ``messages`` as ``role: "system"``; models +without the flag (Claude 4.7 and older, Haiku 4.5) reject that role inside +``messages``, so the proxy must convert the reminder to a user turn in place, +prefixed with an operator note. Either way the prompt cache written on turn one +must be read back in full on turn two. + +The conversation shape mirrors what an OpenAI-SDK client sends mid-session: a +cached system prompt, a user turn carrying its own ``cache_control`` breakpoint, +an assistant turn, a ``role: "system"`` reminder, and a fresh user turn. The +message-turn breakpoint is what makes the cache assertion able to fail: a cache +entry whose prefix spans ``system`` plus message turns is invalidated when the +reminder is hoisted (the ``system`` field mutates and a turn disappears from +``messages``), while an entry ending at the system block itself would survive +the hoist and mask the regression. + +The provider-native ``cache_control`` request shape is not expressible with the +shared ``ChatBody`` (whose content parts carry no cache_control), so the body is +built from the typed content blocks shared in ``models.py``. +""" + +from __future__ import annotations + +import time + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import Result, unwrap +from lifecycle import ResourceManager +from models import CacheControl, ChatResponse, LiteLLMParamsBody, RichMessage, TextBlock, Usage +from passthrough_client import PassthroughClient + +pytestmark = [pytest.mark.e2e, pytest.mark.provider_live] + +CACHE_PRIMING_DEADLINE_SECONDS = 60.0 +CACHE_PRIMING_INTERVAL_SECONDS = 3.0 +CACHE_WARM_CONSECUTIVE_READS = 3 + + +class CacheChatRequest(BaseModel): + """OpenAI-format chat body whose content blocks carry ``cache_control``.""" + + model: str + messages: list[RichMessage] + max_tokens: int = 64 + cache: dict[str, bool] = {"no-cache": True} + + +def _anthropic_params(model: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody(model=model, api_key="os.environ/ANTHROPIC_API_KEY") + + +def _invoke_params(model: str, region: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody(model=model, aws_region_name=region) + + +def _cacheable_system_turn(marker: str) -> RichMessage: + """A system prompt comfortably above the 4096-token minimum cacheable size + of Haiku 4.5 (the smallest model here), unique per run so no other run's + cache entry can satisfy the read.""" + text = " ".join(f"Reference paragraph {index} for run {marker}." for index in range(300)) + return RichMessage(role="system", content=[TextBlock(text=text, cache_control=CacheControl())]) + + +def _user_turn(text: str, *, cached: bool = False) -> RichMessage: + block = TextBlock(text=text, cache_control=CacheControl() if cached else None) + return RichMessage(role="user", content=[block]) + + +def _assistant_turn(text: str) -> RichMessage: + return RichMessage(role="assistant", content=[TextBlock(text=text)]) + + +def _system_reminder_turn() -> RichMessage: + return RichMessage( + role="system", + content=[TextBlock(text="Answer with exactly one word.")], + ) + + +def _post_chat(client: PassthroughClient, key: str, body: CacheChatRequest) -> Result[ChatResponse]: + return client.proxy.transport.post( + "/v1/chat/completions", + headers=client.proxy.transport.bearer(key), + json=body, + response_type=ChatResponse, + ) + + +def _register_deployment(client: PassthroughClient, resources: ResourceManager, params: LiteLLMParamsBody) -> str: + model = f"e2e-chat-midsys-{unique_marker()}" + model_id = client.proxy.create_model(model, params) + resources.defer(lambda: client.proxy.delete_model(model_id)) + return model + + +def _first_turn_user_text(marker: str) -> str: + """A first user turn heavy enough (hundreds of tokens) that losing its cache + entry is unambiguous in the usage numbers, unique per attempt so priming + retries never depend on the proxy's response cache behavior.""" + notes = " ".join(f"Session note {index} for attempt {marker}." for index in range(100)) + return f"Reply with one word.\n{notes}" + + +def _cache_read_tokens(usage: Usage | None) -> int: + """Cache-read tokens however the chat usage reports them: the Anthropic-style + ``cache_read_input_tokens`` litellm forwards, or the OpenAI-style + ``prompt_tokens_details.cached_tokens`` it mirrors them into.""" + if usage is None: + return 0 + if usage.cache_read_input_tokens: + return usage.cache_read_input_tokens + if usage.prompt_tokens_details and usage.prompt_tokens_details.cached_tokens: + return usage.prompt_tokens_details.cached_tokens + return 0 + + +def _cache_creation_tokens(usage: Usage | None) -> int: + if usage is None: + return 0 + return usage.cache_creation_input_tokens or 0 + + +def _response_text(response: ChatResponse) -> str: + return "".join(choice.message.content or "" for choice in response.choices if choice.message) + + +def _response_role(response: ChatResponse) -> str | None: + first = response.choices[0].message if response.choices else None + return first.role if first else None + + +class PrimedCache(BaseModel): + first_user_text: str + prefix_read_tokens: int + first_turn_creation_tokens: int + + @property + def full_prefix_tokens(self) -> int: + return self.prefix_read_tokens + self.first_turn_creation_tokens + + +def _prime_prompt_cache(client: PassthroughClient, key: str, model: str, system_turn: RichMessage) -> PrimedCache: + """Send first-turn calls (fresh cache-marked user turn each attempt, + identical system prefix) until one both reads the system prefix back from + cache and writes its own user-turn chunk, then re-send that exact turn until + its own chunk reads back on three sends in a row, proving the cache is live + in both directions before the reminder turn goes out (a freshly written entry + can take a few seconds to become readable). Only the pre-reminder turn is + ever retried here, so retries can never warm a mutated-prefix cache entry and + mask the regression the second turn asserts on.""" + deadline = time.monotonic() + CACHE_PRIMING_DEADLINE_SECONDS + while True: + user_text = _first_turn_user_text(unique_marker()) + body = CacheChatRequest(model=model, messages=[system_turn, _user_turn(user_text, cached=True)]) + usage = unwrap(_post_chat(client, key, body)).usage + read_tokens = _cache_read_tokens(usage) + creation_tokens = _cache_creation_tokens(usage) + if read_tokens > 0 and creation_tokens > 0: + primed = PrimedCache( + first_user_text=user_text, + prefix_read_tokens=read_tokens, + first_turn_creation_tokens=creation_tokens, + ) + if _first_turn_reads_back(client, key, body, primed.full_prefix_tokens, deadline): + return primed + if time.monotonic() >= deadline: + pytest.fail( + f"{model}: prompt cache never became readable in full within " + f"{CACHE_PRIMING_DEADLINE_SECONDS}s (last usage: {usage})" + ) + time.sleep(CACHE_PRIMING_INTERVAL_SECONDS) + + +def _reads_full_prefix(client: PassthroughClient, key: str, body: CacheChatRequest, full_prefix_tokens: int) -> bool: + return _cache_read_tokens(unwrap(_post_chat(client, key, body)).usage) >= full_prefix_tokens + + +def _first_turn_reads_back( + client: PassthroughClient, + key: str, + body: CacheChatRequest, + full_prefix_tokens: int, + deadline: float, +) -> bool: + """True once the full prefix reads back on CACHE_WARM_CONSECUTIVE_READS sends in + a row. Some providers' global endpoints serve the prompt cache per region, so a + fresh entry can be missing from the region the next request lands on; each miss + re-creates the entry there, so the streak converges as the regions warm up.""" + while time.monotonic() < deadline: + if all(_reads_full_prefix(client, key, body, full_prefix_tokens) for _ in range(CACHE_WARM_CONSECUTIVE_READS)): + return True + time.sleep(CACHE_PRIMING_INTERVAL_SECONDS) + return False + + +def _reminder_turn_body(model: str, system_turn: RichMessage, primed: PrimedCache) -> CacheChatRequest: + """Turn two in OpenAI shape: the primed prefix, an assistant reply, the + mid-conversation system reminder, and a fresh cache-marked user turn.""" + return CacheChatRequest( + model=model, + messages=[ + system_turn, + _user_turn(primed.first_user_text, cached=True), + _assistant_turn("OK."), + _system_reminder_turn(), + _user_turn("Reply with one word again.", cached=True), + ], + ) + + +def _assert_flagged_model_keeps_cache( + client: PassthroughClient, resources: ResourceManager, params: LiteLLMParamsBody +) -> None: + model = _register_deployment(client, resources, params) + key = resources.key(models=[model]) + system_turn = _cacheable_system_turn(unique_marker()) + + primed = _prime_prompt_cache(client, key, model, system_turn) + + second = unwrap(_post_chat(client, key, _reminder_turn_body(model, system_turn, primed))) + read_tokens = _cache_read_tokens(second.usage) + + assert _response_role(second) == "assistant", f"{model}: unexpected role {_response_role(second)!r}" + assert _response_text(second).strip(), f"{model}: reminder turn returned no completion text" + assert read_tokens >= primed.full_prefix_tokens, ( + f"{model}: turn with a mid-conversation system reminder read {read_tokens} " + f"cached tokens, expected at least the {primed.full_prefix_tokens} cached on " + f"turn one ({primed.prefix_read_tokens} system prefix + " + f"{primed.first_turn_creation_tokens} first user turn); the reminder was " + f"hoisted into the top-level system field, which mutates the cached prefix " + f"and re-bills the conversation at cache-write pricing" + ) + + +def _assert_unflagged_model_converts_and_succeeds( + client: PassthroughClient, resources: ResourceManager, params: LiteLLMParamsBody +) -> None: + model = _register_deployment(client, resources, params) + key = resources.key(models=[model]) + system_turn = _cacheable_system_turn(unique_marker()) + + primed = _prime_prompt_cache(client, key, model, system_turn) + + second = unwrap(_post_chat(client, key, _reminder_turn_body(model, system_turn, primed))) + read_tokens = _cache_read_tokens(second.usage) + + assert _response_role(second) == "assistant", f"{model}: unexpected role {_response_role(second)!r}" + assert _response_text(second).strip(), ( + f"{model}: conversation with a mid-conversation system reminder returned " + f"no text; the reminder was forwarded in place to a model that rejects " + f"role 'system' inside messages instead of being converted to a user turn" + ) + assert read_tokens >= primed.full_prefix_tokens, ( + f"{model}: reminder turn read {read_tokens} cached tokens, expected at least " + f"the {primed.full_prefix_tokens} cached on turn one " + f"({primed.prefix_read_tokens} system prefix + " + f"{primed.first_turn_creation_tokens} first user turn); the reminder was " + f"hoisted into the top-level system field instead of being converted to a " + f"user turn in place, mutating the cached prefix and re-billing the " + f"conversation at cache-write pricing" + ) + + +class TestAnthropicChatMidConversationSystem: + FLAGGED_MODEL = "anthropic/claude-opus-4-8" + UNFLAGGED_MODEL = "anthropic/claude-haiku-4-5-20251001" + + @pytest.mark.covers( + "llm.chat_completions.anthropic.mid_conversation_system.nonstream.cache_hit", + exercised_on=[], + ) + def test_flagged_model_keeps_prompt_cache_across_system_reminder( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + _assert_flagged_model_keeps_cache(client, resources, _anthropic_params(self.FLAGGED_MODEL)) + + @pytest.mark.covers( + "llm.chat_completions.anthropic.mid_conversation_system.nonstream.works", + exercised_on=[], + ) + def test_unflagged_model_converts_system_reminder_and_succeeds( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + _assert_unflagged_model_converts_and_succeeds(client, resources, _anthropic_params(self.UNFLAGGED_MODEL)) + + +class TestBedrockInvokeChatMidConversationSystem: + FLAGGED_MODEL = "bedrock/invoke/us.anthropic.claude-sonnet-5" + UNFLAGGED_MODEL = "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0" + AWS_REGION = "us-east-1" + + @pytest.mark.covers( + "llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.cache_hit", + exercised_on=[], + ) + def test_flagged_model_keeps_prompt_cache_across_system_reminder( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + _assert_flagged_model_keeps_cache(client, resources, _invoke_params(self.FLAGGED_MODEL, self.AWS_REGION)) + + @pytest.mark.covers( + "llm.chat_completions.bedrock_invoke.mid_conversation_system.nonstream.works", + exercised_on=[], + ) + def test_unflagged_model_converts_system_reminder_and_succeeds( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + _assert_unflagged_model_converts_and_succeeds( + client, resources, _invoke_params(self.UNFLAGGED_MODEL, self.AWS_REGION) + ) diff --git a/tests/e2e/llm_translation/test_containers_e2e.py b/tests/e2e/llm_translation/test_containers_e2e.py new file mode 100644 index 00000000000..887aecb8df1 --- /dev/null +++ b/tests/e2e/llm_translation/test_containers_e2e.py @@ -0,0 +1,167 @@ +"""Live e2e: an Azure code_interpreter container's file, read back by its native +id with a team service-account key. + +The Azure container endpoints were first verified on one shape: a single Azure +deployment whose credentials came from ``AZURE_API_BASE`` in the proxy env, +containers created explicitly with LiteLLM-managed ids, and the master key as +the caller. The customer differs on all three axes at once, and this cell pins +that shape: + +- the Azure deployments carry their own ``api_base`` and ``api_key`` (the + pytest process reads both from its env and registers them literally), so a + proxy booted with no ``AZURE_API_BASE`` serves them; +- two Azure deployments are registered, the first with an invalid key, so a + container call that guesses a deployment instead of routing by container id + lands on the decoy and fails; +- the container is created implicitly by ``/v1/responses`` with the + ``code_interpreter`` tool, and every container call afterwards names it by + Azure's own ``cntr_`` id with only ``custom_llm_provider=azure`` beside + it, the way a client that stores provider ids does (the routing envelope + LiteLLM wraps around the id in the responses output is peeled off first); +- every LLM-side call is made with a service-account key of a team whose + member is a plain ``internal_user``; a service-account key belongs to the + team, not to a user, and the master key only does the setup. + +Fail-before-fix, proven against a local proxy booted with no Azure env: with +#28990 reverted the upload 403s (the ownership row written at creation no +longer matches a key without a user_id), and with #27921 reverted it fails +with "api_base is required for Azure AI Studio ... Passed `api_base=None`" +because the native id carries no model_id and nothing else names a deployment. +A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the +second regression, since the global-credential fallback then reaches the +container anyway. + +The streaming variant is not here: a streamed ``/v1/responses`` writes the +container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK +closes the connection at ``[DONE]``, so the write is cancelled and every +follow-up container call 403s (LIT-8612). That cell comes with its fix. +""" + +from __future__ import annotations + +import base64 +import binascii +import os +from types import MappingProxyType +from typing import Final + +import pytest +from e2e_config import REQUEST_TIMEOUT, unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from management.management_client import ManagementClient, build_client +from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody +from openai import OpenAI +from openai.types.responses import Response, ResponseCodeInterpreterToolCall +from openai.types.responses.tool_param import CodeInterpreter +from proxy_client import ProxyClient +from sdk_clients import NO_PROXY_CACHE, SdkClients + +pytestmark = [pytest.mark.e2e, pytest.mark.provider_live] + +AZURE_BACKEND: Final = "azure/gpt-5.4-nano" +AZURE_API_VERSION: Final = "v1" +AZURE_PROVIDER_QUERY: Final = MappingProxyType({"custom_llm_provider": "azure"}) +CODE_INTERPRETER: Final[CodeInterpreter] = {"type": "code_interpreter", "container": {"type": "auto"}} +PROMPT: Final = "Use python to compute 6*7 and reply with just the number." +CODE_INTERPRETER_TIMEOUT: Final = 3 * REQUEST_TIMEOUT + + +def _azure_credentials() -> tuple[str, str]: + api_base: Final = os.environ.get("AZURE_API_BASE", "") + api_key: Final = os.environ.get("AZURE_API_KEY", "") + if not api_base or not api_key: + pytest.fail("set AZURE_API_BASE and AZURE_API_KEY in the pytest env; the deployments are registered with them") + return api_base, api_key + + +def _azure_params(api_base: str, api_key: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody(model=AZURE_BACKEND, api_base=api_base, api_key=api_key, api_version=AZURE_API_VERSION) + + +def _register_two_azure_deployments(proxy: ProxyClient, resources: ResourceManager, marker: str) -> str: + api_base, api_key = _azure_credentials() + decoy_id: Final = proxy.create_model( + f"e2e-containers-decoy-{marker}", _azure_params(api_base, f"decoy-{marker}"), provider_live=True + ) + resources.defer(lambda: proxy.delete_model(decoy_id)) + model: Final = f"e2e-containers-{marker}" + model_id: Final = proxy.create_model(model, _azure_params(api_base, api_key), provider_live=True) + resources.defer(lambda: proxy.delete_model(model_id)) + return model + + +def _service_account_key( + proxy: ProxyClient, resources: ResourceManager, management: ManagementClient, marker: str, model: str +) -> str: + team_id: Final = management.create_team(TeamNewBody(team_alias=f"e2e-containers-{marker}", models=[model])) + resources.defer(lambda: management.delete_team(team_id)) + user_id: Final = management.create_user( + UserNewBody(user_email=f"e2e-containers-{marker}@example.com", user_role="internal_user") + ) + resources.defer(lambda: management.delete_user_strict(user_id)) + management.add_team_member(team_id, user_id) + resources.defer(lambda: management.delete_team_member(team_id, user_id)) + generated: Final = unwrap( + proxy.transport.post( + "/key/service-account/generate", + headers=proxy.management_headers(), + json=KeyGenerateBody(team_id=team_id, key_alias=f"e2e-containers-sa-{marker}", models=[model]), + response_type=KeyGenerateResponse, + ) + ) + resources.defer(lambda: management.delete_key_strict(generated.key)) + return generated.key + + +def _response_with_code_interpreter(client: OpenAI, model: str) -> Response: + return client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create( + model=model, input=PROMPT, tools=[CODE_INTERPRETER], tool_choice="required", extra_body=NO_PROXY_CACHE + ) + + +def _container_id(response: Response) -> str: + calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall)) + assert calls, f"no code_interpreter_call in the responses output: {response.output!r}" + return calls[0].container_id + + +def _routing_envelope(container_id: str) -> str | None: + try: + envelope: Final = base64.b64decode(container_id.removeprefix("cntr_"), validate=True).decode() + except (binascii.Error, UnicodeDecodeError): + return None + return envelope if envelope.startswith("litellm:") else None + + +def _native_container_id(container_id: str) -> str: + envelope: Final = _routing_envelope(container_id) + return container_id if envelope is None else envelope.rpartition("container_id:")[2] + + +def _assert_file_round_trip(client: OpenAI, native_id: str, marker: str) -> None: + payload: Final = f"hello from {marker}\n".encode() + uploaded: Final = client.containers.files.create( + native_id, file=(f"{marker}.txt", payload), extra_query=AZURE_PROVIDER_QUERY + ) + fetched: Final = client.containers.files.content.retrieve( + uploaded.id, container_id=native_id, extra_query=AZURE_PROVIDER_QUERY + ) + assert fetched.content == payload, f"container file content differs from the upload: {fetched.content!r}" + + +class TestAzureContainerFiles: + @pytest.mark.covers("llm.responses.azure_openai.code_interpreter.nonstream.works") + def test_service_account_key_reads_container_file_by_native_id( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + marker: Final = unique_marker() + model: Final = _register_two_azure_deployments(proxy, resources, marker) + key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model) + client: Final = sdk.openai(key) + native_id: Final = _native_container_id(_container_id(_response_with_code_interpreter(client, model))) + resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY)) + assert native_id.startswith("cntr_") and _routing_envelope(native_id) is None, ( + f"container id is not the provider's own id: {native_id}" + ) + _assert_file_round_trip(client, native_id, marker) diff --git a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md index 6baebc4c28c..32dc0c47dda 100644 --- a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md +++ b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md @@ -23,7 +23,7 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`. | per-model / per-provider attribution | `test_spend_tracking_utils.py` | unit | covered | yes (`test_each_model_on_a_shared_key_gets_its_own_row`) | | field population (model/tokens/api_key/team/org) | `test_spend_tracking_utils.py` | unit | partial | yes (asserts real values) | | `request_tags` propagation | `test_db_spend_update_writer.py` | unit | partial | yes (`test_request_tags_round_trip`) | -| `end_user` attribution | unit | unit | partial | yes (`test_end_user_spend_attributed_on_row`) | +| `end_user` attribution | unit | unit | partial | yes (`test_end_user_spend_attributed_on_row`, `test_end_user_header_attributes_responses_row`) | ## Cost calculation by modality @@ -67,6 +67,7 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`. | `test_request_tags_round_trip` | tags persist onto the row | | `test_tag_spend_matches_sum_of_tagged_logs` | `/spend/tags` SUM/COUNT == tagged rows | | `test_end_user_spend_attributed_on_row` | `end_user` attributed + costed | +| `test_end_user_header_attributes_responses_row` | `x-litellm-customer-id` / `x-litellm-end-user-id` + `x-litellm-tags` headers on `/v1/responses` (the Codex CLI `http_headers` shape) attributed + tagged + costed, and `/customer/info` spend equals the row | | `test_each_model_on_a_shared_key_gets_its_own_row` | per-model/provider rows, correct model + cost, distinct request_ids matching response id | | `test_failure_call_writes_failure_status_row` | failed call -> `status=failure`, `spend=0` | | `test_spend_calculate_returns_nonzero_cost` | cost-map smoke (no batch wait) | diff --git a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py index 8b63b063e14..e607c12b731 100644 --- a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py +++ b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py @@ -20,6 +20,7 @@ from typing import Final from e2e_config import unique_marker from e2e_http import ( + AuthHeaders, FileUploadForm, Headers, NoBody, @@ -36,6 +37,7 @@ from models import ( ChatMessage, ChatMetadata, ChatResponse, + CustomerInfoParams, DateRangeParams, EmbedBody, EmbedResponse, @@ -63,9 +65,10 @@ METRICS_PATH: Final = "/metrics/" __all__ = [ "BatchCreateBody", + "BatchObject", "CallbackLogMetadata", "CallbackLogPayload", - "BatchObject", + "ClientAttributionHeaders", "DailyActivityKeyBreakdown", "FileObject", "ProbeResult", @@ -80,6 +83,16 @@ __all__ = [ ] +class ClientAttributionHeaders(AuthHeaders): + """The attribution headers a coding agent attaches to every call from its own + config (Codex CLI's config.toml ``http_headers``, Claude Code's + ``ANTHROPIC_CUSTOM_HEADERS``) because it has no body field for the end user.""" + + x_litellm_customer_id: str | None = Field(default=None, alias="x-litellm-customer-id") + x_litellm_end_user_id: str | None = Field(default=None, alias="x-litellm-end-user-id") + x_litellm_tags: str | None = Field(default=None, alias="x-litellm-tags") + + class GeminiApiKeyHeaders(Headers): x_goog_api_key: str = Field(serialization_alias="x-goog-api-key") content_type: str = Field(default="application/json", serialization_alias="Content-Type") @@ -222,6 +235,10 @@ class TeamInfoSpendResponse(BaseModel): team_info: TeamInfoSpend +class CustomerSpendResponse(BaseModel): + spend: float | None = None + + def _chat_body( model: str, content: str, @@ -371,6 +388,31 @@ class SpendClient: ) return outcome.result if isinstance(outcome, Converged) else outcome.last_result + def customer_spend(self, customer_id: str) -> float: + """0.0 until the spend writer has upserted the end-user row, which /customer/info 404s before.""" + looked_up: Final = self.proxy.transport.get( + "/customer/info", + headers=self.proxy.transport.master, + params=CustomerInfoParams(end_user_id=customer_id), + response_type=CustomerSpendResponse, + ) + match looked_up: + case Success(data=data): + return data.spend or 0.0 + case _: + return 0.0 + + def poll_customer_spend(self, customer_id: str, *, minimum: float = 0.0) -> float: + outcome: Final = await_converged( + lambda: self.customer_spend(customer_id), + converged=lambda spend: spend > minimum, + timeout=self.proxy.poll_timeout, + interval=self.proxy.poll_interval, + now=time.monotonic, + sleep=time.sleep, + ) + return outcome.result if isinstance(outcome, Converged) else outcome.last_result + def scrape_metrics(self) -> Mapping[str, ProbeResult]: """GET /metrics/ on every replica in PROXY_REPLICA_URLS, keyed by replica. The counter is per pod, so the union of the replicas is the fleet's exposition; the @@ -479,9 +521,12 @@ class SpendClient: ) def send_responses(self, key: str, model: str, content: str) -> StreamingResponse: + return self.send_responses_with_headers(self.proxy.transport.bearer(key), model, content) + + def send_responses_with_headers(self, headers: AuthHeaders, model: str, content: str) -> StreamingResponse: return self.proxy.transport.send( "/v1/responses", - headers=self.proxy.transport.bearer(key), + headers=headers, json=ResponsesBody(model=model, input=content), ) diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py index 286421e2e3f..6633396b538 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py @@ -24,7 +24,14 @@ import pytest from e2e_http import RateLimitedError, Success from lifecycle import ResourceManager from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogs, SpendLogsParams -from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap +from spend_e2e_client import ( + ClientAttributionHeaders, + SpendClient, + SpendLogRow, + is_ok, + unique_marker, + unwrap, +) pytestmark = pytest.mark.e2e @@ -425,6 +432,41 @@ def test_end_user_spend_attributed_on_row( assert (row.spend or 0) > 0, f"end-user row should cost > 0: {_summarize(rows)}" +@pytest.mark.covers("quota_management.spend_tracking.end_user.attributes_responses_header") +@pytest.mark.parametrize("header", ["x-litellm-customer-id", "x-litellm-end-user-id"]) +def test_end_user_header_attributes_responses_row( + client: SpendClient, scoped_key: str, resources: ResourceManager, header: str +) -> None: + """Codex CLI has no body field for the end user, so its config.toml http_headers + attach the customer header (and x-litellm-tags) to every /v1/responses call. + A regression that stops reading either header on the Responses route, drops the + tags, costs the row at zero, or leaves the customer's own spend total behind the + row fails here.""" + customer = resources.customer(f"e2e-codex-{unique_marker()}") + tag = f"codex-{unique_marker()}" + headers = ClientAttributionHeaders.model_validate( + {"authorization": f"Bearer {scoped_key}", header: customer, "x-litellm-tags": tag} + ) + sent = client.send_responses_with_headers( + headers, "openai-responses-codex", f"one word {unique_marker()}" + ) + assert sent.ok, f"/v1/responses failed with {sent.status_code}: {sent.body[:300]}" + + rows = client.poll_logs_for_key( + scoped_key, predicate=lambda rs: any(r.end_user == customer for r in rs) + ) + row = _require_row( + rows, lambda r: r.end_user == customer, f"attributed to end_user {customer!r} via {header}" + ) + assert row.call_type == "aresponses", f"row is not a Responses row: {_summarize(rows)}" + assert tag in (row.request_tags or []), f"tag {tag!r} missing from {row.request_tags}" + assert (row.spend or 0) > 0, f"end-user row should cost > 0: {_summarize(rows)}" + customer_total = client.poll_customer_spend(customer) + assert _approx_equal(customer_total, row.spend or 0), ( + f"/customer/info spend {customer_total} != the row's {row.spend}: {_summarize(rows)}" + ) + + @pytest.mark.covers("quota_management.spend_tracking.per_model.writes_own_rows") def test_each_model_on_a_shared_key_gets_its_own_row( client: SpendClient, scoped_key: str diff --git a/tests/e2e/ui/fixtures/pages.ts b/tests/e2e/ui/fixtures/pages.ts index 56b2bed380d..ba5887f3113 100644 --- a/tests/e2e/ui/fixtures/pages.ts +++ b/tests/e2e/ui/fixtures/pages.ts @@ -20,6 +20,7 @@ export enum Page { RouterSettings = "router-settings", UiTheme = "ui-theme", CostTracking = "cost-tracking", + CostOptimization = "cost-optimization", ModelHubTable = "model-hub-table", Caching = "caching", Logs = "logs", diff --git a/tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts b/tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts new file mode 100644 index 00000000000..08dcc4c2c21 --- /dev/null +++ b/tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts @@ -0,0 +1,51 @@ +import { test, expect } from "@playwright/test"; +import { randomUUID } from "node:crypto"; +import { execFileSync } from "node:child_process"; +import * as path from "node:path"; +import { Page } from "../../fixtures/pages"; +import { navigateToPage } from "../../helpers/navigation"; + +test("cache leakage by model merges a deployment's resolved and requested model names into its model group", async ({ + page, +}) => { + const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master"; + const marker = `integration-browser-${randomUUID()}`; + const group = `${marker}-public`; + const deployment = `${marker}-backend`; + const apiKey = marker; + const support = (...args: string[]) => + execFileSync( + process.env.INTEGRATION_PYTHON ?? "python", + [ + path.resolve( + __dirname, + "../../../../integration/_support/daily_spend_rows.py", + ), + ...args, + ], + { encoding: "utf8", timeout: 10_000, killSignal: "SIGKILL" }, + ); + try { + support("seed", apiKey, group, deployment); + await page.goto("/ui/login"); + await page.getByPlaceholder("Enter your username").fill("admin"); + await page.getByPlaceholder("Enter your password").fill(master); + await page.getByRole("button", { name: "Login", exact: true }).click(); + await expect(page).toHaveURL( + (url) => + url.pathname.startsWith("/ui") && !url.pathname.includes("login"), + ); + await navigateToPage(page, Page.CostOptimization); + await page.getByRole("tab", { name: "Prompt Caching" }).click(); + await page.getByRole("tab", { name: "By model" }).click(); + const rows = page.getByRole("row").filter({ hasText: marker }); + await expect(rows).toHaveCount(1); + await expect(rows.first()).toContainText(group); + await expect(rows.first()).toContainText("275,000"); + await expect( + page.getByRole("row").filter({ hasText: deployment }), + ).toHaveCount(0); + } finally { + support("clear", apiKey); + } +}); diff --git a/tests/e2e/ui/tests/integrationCritical/expected.json b/tests/e2e/ui/tests/integrationCritical/expected.json index 189cdef9a93..b73c04acbe0 100644 --- a/tests/e2e/ui/tests/integrationCritical/expected.json +++ b/tests/e2e/ui/tests/integrationCritical/expected.json @@ -6,5 +6,6 @@ "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::pressing Enter on Update opens the credentials modal instead of the server editor", "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server with two per-user variables reports the remaining gap until both are saved", "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server without per-user variables shows no credential row", - "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page" + "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page", + "tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group" ] diff --git a/tests/integration/_support/daily_spend_rows.py b/tests/integration/_support/daily_spend_rows.py new file mode 100644 index 00000000000..9d530fb436a --- /dev/null +++ b/tests/integration/_support/daily_spend_rows.py @@ -0,0 +1,52 @@ +import json +import sys +from datetime import datetime, timezone +from typing import Final, LiteralString + +from integration._support.database import read_rows, write_rows + +SEED_QUERY: Final[LiteralString] = """ +INSERT INTO "LiteLLM_DailyUserSpend" ( + id, user_id, date, api_key, model, model_group, custom_llm_provider, + endpoint, mcp_namespaced_tool_name, + prompt_tokens, completion_tokens, spend, + api_requests, successful_requests, failed_requests, updated_at +) VALUES ( + gen_random_uuid()::text, %s, %s, %s, %s, %s, 'bedrock', + NULL, NULL, + %s, %s, %s, + %s, %s, %s, now() +) +""" + +CLEAR_COUNT_QUERY: Final[LiteralString] = 'SELECT count(*) AS count FROM "LiteLLM_DailyUserSpend" WHERE api_key=%s' +CLEAR_QUERY: Final[LiteralString] = 'DELETE FROM "LiteLLM_DailyUserSpend" WHERE api_key=%s' + + +def seed(api_key: str, group: str, deployment: str) -> int: + today: Final = datetime.now(timezone.utc).date().isoformat() + write_rows( + SEED_QUERY, + (api_key, today, api_key, deployment, group, "270000", "1000", "0.81", "270", "270", "0"), + ) + write_rows( + SEED_QUERY, + (api_key, today, api_key, group, "", "5000", "0", "0.0", "50", "0", "50"), + ) + return 2 + + +def clear(api_key: str) -> int: + count: Final = int(str(read_rows(CLEAR_COUNT_QUERY, (api_key,))[0]["count"])) + write_rows(CLEAR_QUERY, (api_key,)) + return count + + +if __name__ == "__main__": + command: Final = sys.argv[1] + if command == "seed": + sys.stdout.write(json.dumps({"affected": seed(sys.argv[2], sys.argv[3], sys.argv[4])}) + "\n") + elif command == "clear": + sys.stdout.write(json.dumps({"affected": clear(sys.argv[2])}) + "\n") + else: + raise SystemExit(f"unknown command: {command}") diff --git a/tests/integration/_support/database.py b/tests/integration/_support/database.py index e7f0ebdf603..461cdbda1ee 100644 --- a/tests/integration/_support/database.py +++ b/tests/integration/_support/database.py @@ -1,5 +1,5 @@ import os -from typing import Final +from typing import Final, LiteralString import psycopg from psycopg.rows import dict_row @@ -14,3 +14,8 @@ def read_rows( with psycopg.connect(database_url or os.environ["DATABASE_URL"], row_factory=dict_row) as connection: connection.execute("SET TRANSACTION READ ONLY") return ROWS.validate_python(connection.execute(query, parameters).fetchall()) + + +def write_rows(query: LiteralString, parameters: tuple[str, ...], *, database_url: str | None = None) -> None: + with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection: + connection.execute(query, parameters) diff --git a/tests/integration/authorization/test_key_alias_model_access.py b/tests/integration/authorization/test_key_alias_model_access.py new file mode 100644 index 00000000000..50fc53bd4a9 --- /dev/null +++ b/tests/integration/authorization/test_key_alias_model_access.py @@ -0,0 +1,72 @@ +import uuid +from typing import Final + +import httpx + +from tests.integration._support.client import Gateway, eventually, object_value, string_value + + +def _listed_model_ids(response: httpx.Response) -> frozenset[str]: + entries: Final = response.json()["data"] + assert isinstance(entries, list), response.text + return frozenset(string_value(object_value(entry)["id"]) for entry in entries) + + +def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None: + """Every id /v1/models lists for this key must be callable by the same key.""" + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + listed: Final = _listed_model_ids(response) + assert listed == frozenset({model, alias}), response.text + for model_id in sorted(listed): + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model_id, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}" + + +def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[model], aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team_id: Final = scenario.team(models=[model]) + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(team_id=team_id, aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + allowed: Final = scenario.model() + hidden: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[allowed], aliases={alias: hidden}) + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + assert _listed_model_ids(response) == frozenset({allowed}), response.text + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": alias, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 403, called.text + assert "key_model_access_denied" in called.text, called.text diff --git a/tests/integration/authorization/test_warmed_policy.py b/tests/integration/authorization/test_warmed_policy.py index a9bee196ddd..b03610c894b 100644 --- a/tests/integration/authorization/test_warmed_policy.py +++ b/tests/integration/authorization/test_warmed_policy.py @@ -1,14 +1,14 @@ +import os from collections.abc import Iterator from contextlib import ExitStack, contextmanager from hashlib import sha256 from typing import Final -import os import psycopg import pytest -from pydantic import JsonValue from hypothesis import strategies as st from hypothesis.stateful import RuleBasedStateMachine, invariant, rule, run_state_machine_as_test +from pydantic import JsonValue from tests.integration._support.client import Gateway, eventually, object_value from tests.integration._support.database import read_rows @@ -195,6 +195,71 @@ def test_warmed_team_role_demotion_prevents_later_management_writes(gateway: Gat assert_serving(gateway, model, caller, 200) +def _key_row(key: str) -> dict[str, JsonValue]: + rows: Final = read_rows( + 'SELECT max_budget, key_alias FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + assert len(rows) == 1 + return rows[0] + + +@pytest.mark.covers("mgmt.key.update.team_admin_member_key_budget_requires_opt_in") +def test_team_admin_changes_member_key_budget_only_when_opted_in(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + admin: Final = scenario.user(user_role="internal_user") + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team( + models=[model], + members_with_roles=[{"user_id": admin, "role": "admin"}, {"user_id": member, "role": "user"}], + ) + other_team: Final = scenario.team(models=[model], members_with_roles=[{"user_id": member, "role": "user"}]) + member_key: Final = scenario.key( + user_id=member, team_id=team, models=[model], max_budget=10, key_alias="member" + ) + personal_key: Final = scenario.key(user_id=member, models=[model], max_budget=10) + foreign_key: Final = scenario.key(user_id=member, team_id=other_team, models=[model], max_budget=10) + admin_key: Final = scenario.key( + user_id=admin, team_id=team, models=[model], allowed_routes=["/key/update", "/v1/chat/completions"] + ) + member_caller: Final = scenario.key( + user_id=member, team_id=team, models=[model], allowed_routes=["/key/update", "/v1/chat/completions"] + ) + assert_serving(gateway, model, member_key, 200) + with _team_admins_may_edit(gateway, []): + denied: Final = gateway.request("POST", "/key/update", {"key": member_key, "max_budget": 0}, key=admin_key) + assert denied.status_code == 403, denied.text + assert _key_row(member_key) == {"max_budget": 10.0, "key_alias": "member"} + with _team_admins_may_edit(gateway, ["member_key_budgets"]): + for target in (personal_key, foreign_key): + out_of_scope: Final = gateway.request( + "POST", "/key/update", {"key": target, "max_budget": 0}, key=admin_key + ) + assert out_of_scope.status_code == 403, out_of_scope.text + assert _key_row(target)["max_budget"] == 10.0 + by_member: Final = gateway.request( + "POST", "/key/update", {"key": admin_key, "max_budget": 0}, key=member_caller + ) + assert by_member.status_code == 403, by_member.text + not_budget: Final = gateway.request( + "POST", "/key/update", {"key": member_key, "key_alias": "renamed"}, key=admin_key + ) + assert not_budget.status_code == 403, not_budget.text + assert _key_row(member_key) == {"max_budget": 10.0, "key_alias": "member"} + changed: Final = gateway.request( + "POST", "/key/update", {"key": member_key, "max_budget": 0, "budget_duration": "30d"}, key=admin_key + ) + assert changed.status_code == 200, changed.text + assert _key_row(member_key) == {"max_budget": 0.0, "key_alias": "member"} + assert_serving(gateway, model, member_key, 422, "budget_exceeded") + restored: Final = gateway.request( + "POST", "/key/update", {"key": member_key, "max_budget": 10}, key=admin_key + ) + assert restored.status_code == 200, restored.text + assert_serving(gateway, model, member_key, 200) + + @pytest.mark.covers("mgmt.key.update.expiry_changes_reach_warmed_workers") def test_expiry_and_explicit_clear_reach_both_warmed_workers(gateway: Gateway, peer: Gateway) -> None: with gateway.scenario() as scenario: diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py new file mode 100644 index 00000000000..5a3ffdb9965 --- /dev/null +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -0,0 +1,270 @@ +import base64 +import json +import time +import uuid +from collections.abc import Sequence +from pathlib import Path +from typing import Final + +import yaml +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.common.v1.common_pb2 import KeyValue +from opentelemetry.proto.trace.v1.trace_pb2 import Span +from pydantic import BaseModel, TypeAdapter + +PUBLIC_KEY: Final = "pk-lf-integration" +SECRET_KEY: Final = "sk-lf-integration" +PROJECTS_PATH: Final = "/api/public/projects" +TRACES_PATH: Final = "/api/public/otel/v1/traces" +PROMPTS_PATH: Final = "/api/public/v2/prompts/" +_PROXY_CONFIG: Final = TypeAdapter(dict[str, object]) +_SETTINGS: Final = TypeAdapter(dict[str, object]) + + +class _ProviderBody(BaseModel): + messages: list[object] + + +def _completion(text: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": "chatcmpl-" + text, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ).encode() + ) + + +def _projects() -> Reply: + return Reply(body=json.dumps({"data": [{"id": "integration-project", "name": "integration"}]}).encode()) + + +def _text_prompt(name: str) -> Reply: + return Reply( + body=json.dumps( + { + "type": "text", + "name": name, + "version": 1, + "prompt": "Say {{word}}", + "config": {}, + "labels": ["production"], + "tags": [], + } + ).encode() + ) + + +def _langfuse_config(tmp_path: Path) -> Path: + config: Final = _PROXY_CONFIG.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + settings: Final = { + **_SETTINGS.validate_python(config["litellm_settings"]), + "success_callback": ["langfuse"], + "failure_callback": ["langfuse"], + } + path: Final = tmp_path / "langfuse.yaml" + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + return path + + +def _langfuse_environment(langfuse: Wire) -> dict[str, str]: + return { + "LANGFUSE_HOST": langfuse.url, + "LANGFUSE_PUBLIC_KEY": PUBLIC_KEY, + "LANGFUSE_SECRET_KEY": SECRET_KEY, + "LANGFUSE_FLUSH_INTERVAL": "1", + } + + +def _attribute(entries: Sequence[KeyValue], key: str) -> str | list[str] | None: + for entry in entries: + if entry.key != key: + continue + if entry.value.HasField("array_value"): + return [item.string_value for item in entry.value.array_value.values] + return entry.value.string_value + return None + + +def _spans(batches: Sequence[Request]) -> tuple[Span, ...]: + return tuple( + span + for batch in batches + if batch.target == TRACES_PATH and batch.headers.get("content-type") == "application/x-protobuf" + for resource_spans in ExportTraceServiceRequest.FromString(batch.body).resource_spans + for scope_spans in resource_spans.scope_spans + for span in scope_spans.spans + ) + + +def test_langfuse_callback_delivers_the_generation_over_otlp_v4_with_the_caller_trace_fields( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "langfuse" + uuid.uuid4().hex + trace_id: Final = uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + + def upstream(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {provider_secret}" + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker + "-question"}], + "metadata": { + "trace_id": trace_id, + "trace_name": marker + "-trace", + "generation_name": marker, + "trace_user_id": marker + "-user", + "session_id": marker + "-session", + "tags": [marker], + }, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def exported() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple(span for span in _spans(received) if span.name == marker) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=20) + span: Final = spans[0] + posts: Final = tuple(request for request in received if request.method == "POST") + assert {request.target for request in posts} == {TRACES_PATH}, [request.target for request in received] + basic: Final = "Basic " + base64.b64encode(f"{PUBLIC_KEY}:{SECRET_KEY}".encode()).decode() + for request in posts: + assert request.headers["authorization"] == basic + assert request.headers["content-type"] == "application/x-protobuf" + assert request.headers["x-langfuse-ingestion-version"] == "4" + assert provider_secret.encode() not in request.body + assert candidate.key.encode() not in request.body + + assert span.trace_id.hex() == trace_id + assert span.parent_span_id == b"" + attributes: Final = span.attributes + assert _attribute(attributes, "langfuse.observation.type") == "generation" + assert _attribute(attributes, "langfuse.trace.name") == marker + "-trace" + assert _attribute(attributes, "user.id") == marker + "-user" + assert _attribute(attributes, "session.id") == marker + "-session" + assert marker in (_attribute(attributes, "langfuse.trace.tags") or ()) + assert _attribute(attributes, "langfuse.observation.model.name") == "openai/gpt-4o-mini" + assert json.loads(str(_attribute(attributes, "langfuse.observation.usage_details"))) == { + "input": 11, + "output": 4, + "total": 15, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + } + assert marker + "-question" in str(_attribute(attributes, "langfuse.observation.input")) + assert marker + "-answer" in str(_attribute(attributes, "langfuse.observation.output")) + assert ( + _attribute(attributes, "langfuse.observation.metadata.litellm_call_id") + == response.headers["x-litellm-call-id"] + ) + + +def test_prompt_fetch_encodes_the_name_retries_a_5xx_once_and_keeps_langfuse_headers_off_the_client( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "prompt" + uuid.uuid4().hex + leak: Final = "leak-" + marker + flaky_prompt: Final = f"{marker}/what?" + encoded_flaky_prompt: Final = f"{marker}%2Fwhat%3F" + missing_prompt: Final = marker + "-missing" + + seen_prompt_gets: Final[list[str]] = [] # mutable-ok: the double counts attempts across requests + + def upstream(request: Request) -> Reply: + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + if request.method == "POST": + return Reply(body=b"", content_type="application/x-protobuf") + assert request.target.startswith(PROMPTS_PATH), request.target + assert request.headers["authorization"].startswith("Basic ") + if request.target.startswith(PROMPTS_PATH + encoded_flaky_prompt): + prior: Final = sum(1 for seen in seen_prompt_gets if seen.startswith(PROMPTS_PATH + encoded_flaky_prompt)) + seen_prompt_gets.append(request.target) + if prior == 0: + return Reply(status=503, body=b'{"message":"try later"}', headers={"retry-after": "30"}) + return _text_prompt(flaky_prompt) + seen_prompt_gets.append(request.target) + return Reply( + status=404, + body=b'{"message":"Prompt not found","error":"LangfuseNotFoundError"}', + headers={"set-cookie": f"session={leak}; Path=/", "x-upstream-internal": leak, "server": leak}, + ) + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + flaky: Final = scenario.model( + model="langfuse/gpt-4o-mini", prompt_id=flaky_prompt, api_base=provider.url + "/v1", api_key="synthetic" + ) + missing: Final = scenario.model( + model="langfuse/gpt-4o-mini", prompt_id=missing_prompt, api_base=provider.url + "/v1", api_key="synthetic" + ) + started: Final = time.monotonic() + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": flaky, "messages": [{"role": "user", "content": marker}], "prompt_variables": {"word": marker}}, + ) + elapsed: Final = time.monotonic() - started + assert response.status_code == 200, response.text + assert elapsed < 5, f"a retried cold prompt miss took {elapsed:.1f}s" + attempts: Final = tuple( + target for target in seen_prompt_gets if target.startswith(PROMPTS_PATH + encoded_flaky_prompt) + ) + assert len(attempts) == 2, seen_prompt_gets + assert all(target.split("?", 1)[0] == PROMPTS_PATH + encoded_flaky_prompt for target in attempts), attempts + sent: Final = _ProviderBody.model_validate_json(provider.drain()[-1].body).messages + assert any("Say " + marker in json.dumps(message) for message in sent), sent + + failure: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": missing, "messages": [{"role": "user", "content": marker}], "prompt_variables": {"word": marker}}, + ) + assert failure.status_code == 404, failure.text + assert "Prompt not found" in failure.text + assert leak not in failure.text + assert leak not in json.dumps(dict(failure.headers)) + assert "set-cookie" not in failure.headers and "x-upstream-internal" not in failure.headers + assert sum(1 for target in seen_prompt_gets if target.startswith(PROMPTS_PATH + missing_prompt)) == 1 diff --git a/tests/integration/observability/test_presidio_streaming_output.py b/tests/integration/observability/test_presidio_streaming_output.py index 5bc0a46427b..680617f1652 100644 --- a/tests/integration/observability/test_presidio_streaming_output.py +++ b/tests/integration/observability/test_presidio_streaming_output.py @@ -181,7 +181,7 @@ class Rig: "model": self.anthropic, "max_tokens": 64, "stream": True, - "messages": [{"role": "user", "content": "who designed it"}], + "messages": [{"role": "user", "content": f"who designed it {uuid.uuid4().hex}"}], **({"guardrails": list(guardrails)} if guardrails is not None else {}), } @@ -193,6 +193,13 @@ def anthropic_text(received: Received) -> str: return "".join(event["delta"]["text"] for event in events if event.get("type") == "content_block_delta") +def anthropic_message_id(received: Received) -> str: + events: Final = tuple( + json.loads(line.removeprefix("data: ")) for line in received.text.split("\n") if line.startswith("data: ") + ) + return "".join(event["message"]["id"] for event in events if event.get("type") == "message_start") + + @contextmanager def presidio_rig( gateway: Gateway, @@ -339,10 +346,10 @@ def test_native_gemini_unauthenticated_request_is_rejected_before_upstream(gatew assert rig.upstream.drain() == () -def anthropic_provider(chunks: tuple[bytes, ...]) -> Callable[[Request], Reply]: +def anthropic_provider(chunks: tuple[bytes, ...], *, pause_between_chunks: float = 0) -> Callable[[Request], Reply]: def provider(request: Request) -> Reply: assert request.target == "/v1/messages", request.target - return Reply(content_type="text/event-stream", chunks=chunks) + return Reply(content_type="text/event-stream", chunks=chunks, pause_between_chunks=pause_between_chunks) return provider @@ -376,6 +383,50 @@ def test_anthropic_messages_first_frame_split_across_transport_chunks_is_still_m assert received.text.count("event: message_start") == 1 +def test_anthropic_messages_first_frame_split_inside_a_utf8_character_is_still_masked( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "msg_" + uuid.uuid4().hex + whole: Final = anthropic_stream(identity, f"{PERSON} designed the caf\u00e9.") + delta: Final = whole[2].replace("\\u00e9".encode(), "\u00e9".encode()) + split_at: Final = delta.index("\u00e9".encode()) + 1 + assert delta[split_at - 1 : split_at] == b"\xc3", delta + chunks: Final = (whole[0] + whole[1] + delta[:split_at], delta[split_at:], *whole[3:]) + with presidio_rig(gateway, tmp_path, anthropic_provider(chunks, pause_between_chunks=0.5)) as rig: + received: Final = rig.stream("/v1/messages", rig.messages_body()) + assert received.status == 200, received.text + assert anthropic_text(received) == f"{MASK} designed the caf\u00e9." + assert PERSON not in received.text, received.text + assert anthropic_message_id(received) == identity, received.text + + +def test_anthropic_messages_stream_led_by_sse_comment_keepalive_is_still_masked( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "msg_" + uuid.uuid4().hex + chunks: Final = (b": keepalive\n\n", *anthropic_stream(identity, f"{PERSON} designed it.")) + with presidio_rig(gateway, tmp_path, anthropic_provider(chunks, pause_between_chunks=0.5)) as rig: + received: Final = rig.stream("/v1/messages", rig.messages_body()) + assert received.status == 200, received.text + assert received.text.startswith(": keepalive"), received.text[:200] + assert anthropic_text(received) == f"{MASK} designed it." + assert PERSON not in received.text, received.text + assert identity in received.text + + +def test_anthropic_messages_stream_led_by_data_less_ping_event_is_still_masked( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "msg_" + uuid.uuid4().hex + chunks: Final = (b"event: ping\n\n", *anthropic_stream(identity, f"{PERSON} designed it.")) + with presidio_rig(gateway, tmp_path, anthropic_provider(chunks, pause_between_chunks=0.5)) as rig: + received: Final = rig.stream("/v1/messages", rig.messages_body()) + assert received.status == 200, received.text + assert anthropic_text(received) == f"{MASK} designed it." + assert PERSON not in received.text, received.text + assert identity in received.text + + def test_anthropic_messages_stream_fails_closed_when_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None: identity: Final = "msg_" + uuid.uuid4().hex provider: Final = anthropic_provider(anthropic_stream(identity, f"{PERSON} designed it.")) @@ -468,7 +519,7 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t return Reply( content_type="text/event-stream", chunks=(gemini_frame(f"{PERSON} "), gemini_frame("designed it.")), - pause_between_chunks=0.05, + pause_between_chunks=0.5, ) with presidio_rig(gateway, tmp_path, provider, anonymize=flaky_anonymizer) as rig: @@ -514,7 +565,7 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t def test_native_gemini_keeps_streaming_after_one_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None: frames: Final = (gemini_frame("alive "), gemini_frame("still.")) - provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, pause_between_chunks=0.05)) + provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, pause_between_chunks=0.5)) with presidio_rig(gateway, tmp_path, provider) as rig: workers: Final = eventually( lambda: tuple( diff --git a/tests/integration/pricing/test_configured_prices.py b/tests/integration/pricing/test_configured_prices.py index 0e4efea3a15..93290a404fc 100644 --- a/tests/integration/pricing/test_configured_prices.py +++ b/tests/integration/pricing/test_configured_prices.py @@ -1,6 +1,9 @@ +import asyncio import json +import os import uuid from collections.abc import Iterator, Mapping +from hashlib import sha256 from pathlib import Path from typing import Final @@ -12,6 +15,96 @@ from litellm import get_model_info from tests.integration._support.client import Gateway, eventually, object_value, string_value from tests.integration._support.database import read_rows from tests.integration._support.process import owned_proxy +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse +from tests.integration.pricing.test_realtime_cached_audio_pricing import one_realtime_turn + +REALTIME_MODEL: Final = "gpt-realtime-2" +REALTIME_INPUT_TEXT_TOKENS: Final = 10 +REALTIME_INPUT_AUDIO_TOKENS: Final = 20 +REALTIME_OUTPUT_TEXT_TOKENS: Final = 5 +REALTIME_OUTPUT_AUDIO_TOKENS: Final = 7 + + +def _realtime_response_done() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": REALTIME_INPUT_TEXT_TOKENS + + REALTIME_INPUT_AUDIO_TOKENS + + REALTIME_OUTPUT_TEXT_TOKENS + + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_tokens": REALTIME_INPUT_TEXT_TOKENS + REALTIME_INPUT_AUDIO_TOKENS, + "output_tokens": REALTIME_OUTPUT_TEXT_TOKENS + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_token_details": { + "text_tokens": REALTIME_INPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_INPUT_AUDIO_TOKENS, + "cached_tokens": 0, + }, + "output_token_details": { + "text_tokens": REALTIME_OUTPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_OUTPUT_AUDIO_TOKENS, + }, + }, + }, + }, + ), + ) + + +@pytest.mark.parametrize( + ("input_text_rate", "input_audio_rate", "output_text_rate", "output_audio_rate"), + ((0.001, 0.002, 0.003, 0.004), (0.0, 0.0, 0.0, 0.0)), + ids=("custom_rates", "zero_rated"), +) +def test_realtime_session_is_charged_at_the_deployment_configured_rates( + gateway: Gateway, + input_text_rate: float, + input_audio_rate: float, + output_text_rate: float, + output_audio_rate: float, +) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"realtime-configured-price-{uuid.uuid4().hex[:12]}" + handle: Final = register_scenario(scenario_id, _realtime_response_done()) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/{REALTIME_MODEL}", + api_key=scenario_id, + api_base=gateway.upstream_url.rstrip("/"), + input_cost_per_token=input_text_rate, + input_cost_per_audio_token=input_audio_rate, + output_cost_per_token=output_text_rate, + output_cost_per_audio_token=output_audio_rate, + ) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + assert session.get("type") == "session.created", session + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s', + (sha256(key.encode()).hexdigest(),), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["call_type"] == "_arealtime", rows + assert float(str(rows[0]["spend"])) == pytest.approx( + REALTIME_INPUT_TEXT_TOKENS * input_text_rate + + REALTIME_INPUT_AUDIO_TOKENS * input_audio_rate + + REALTIME_OUTPUT_TEXT_TOKENS * output_text_rate + + REALTIME_OUTPUT_AUDIO_TOKENS * output_audio_rate, + abs=1e-9, + ), rows @pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates") diff --git a/tests/integration/pricing/test_realtime_cached_audio_pricing.py b/tests/integration/pricing/test_realtime_cached_audio_pricing.py index 4a7598d0cbf..42e90cf2c3e 100644 --- a/tests/integration/pricing/test_realtime_cached_audio_pricing.py +++ b/tests/integration/pricing/test_realtime_cached_audio_pricing.py @@ -92,7 +92,7 @@ def cached_audio_response_done() -> RealtimeResponse: ) -async def _one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: +async def one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: async with websockets.connect( f"{proxy_url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model}", additional_headers={"Authorization": f"Bearer {key}"}, @@ -115,7 +115,7 @@ def test_realtime_cached_audio_tokens_bill_at_audio_cache_read_rate_not_full_aud model: Final = scenario.model( model=f"openai/{MODEL}", api_key=scenario_id, api_base=gateway.upstream_url.rstrip("/") ) - session: Final = asyncio.run(_one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) assert session.get("type") == "session.created", session rows: Final = eventually( lambda: read_rows( diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index fb20cdf7e0e..92947fbf6fe 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -842,6 +842,7 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): """ - Unit test for `_get_trace_id` function in Logging obj """ + from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id from litellm.litellm_core_utils.litellm_logging import Logging litellm.success_callback = ["langfuse"] @@ -874,24 +875,18 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): time.sleep(3) assert litellm_logging_obj._get_trace_id(service_name="langfuse") is not None - ## if existing_trace_id exists + # langfuse addresses a trace by a 32-hex id, so the id litellm reports back is the + # resolved form of whichever source won; that is what the alerting deep link needs if langfuse_existing_trace_id is not None: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == langfuse_existing_trace_id - ) - ## if trace_id exists + expected_source = langfuse_existing_trace_id elif langfuse_trace_id is not None: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == langfuse_trace_id - ) - ## if no trace_id or existing_trace_id is provided, use litellm_trace_id + expected_source = langfuse_trace_id else: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == litellm_logging_obj.litellm_trace_id - ) + expected_source = litellm_logging_obj.litellm_trace_id + + assert litellm_logging_obj._get_trace_id(service_name="langfuse") == resolve_trace_id( + expected_source + ) def test_convert_model_response_object(): 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_alangfuse.py b/tests/local_testing/test_alangfuse.py index 7b1f7f203e3..a9d111843fd 100644 --- a/tests/local_testing/test_alangfuse.py +++ b/tests/local_testing/test_alangfuse.py @@ -11,6 +11,7 @@ logging.basicConfig(level=logging.DEBUG) import litellm from litellm import completion from litellm.caching import InMemoryCache +from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id litellm.num_retries = 3 litellm.success_callback = ["langfuse"] @@ -36,7 +37,7 @@ def langfuse_client(): langfuse_client = langfuse.Langfuse( public_key=os.environ["LANGFUSE_PUBLIC_KEY"], secret_key=os.environ["LANGFUSE_SECRET_KEY"], - host="https://us.cloud.langfuse.com", + host=os.environ.get("LANGFUSE_HOST", "https://us.cloud.langfuse.com"), ) litellm.in_memory_llm_clients_cache.set_cache( key=_langfuse_cache_key, @@ -227,29 +228,27 @@ async def test_langfuse_logging_without_request_response(stream, langfuse_client print(chunk) langfuse_client.flush() - await asyncio.sleep(5) - # get trace with _unique_trace_name - trace = langfuse_client.get_generations(trace_id=_unique_trace_name) - - print("trace_from_langfuse", trace) - - _trace_data = trace.data - - if ( - len(_trace_data) == 0 - ): # prevent infrequent list index out of range error from langfuse api - return + for _ in range(30): + _trace_data = langfuse_client.api.observations.get_many( + trace_id=resolve_trace_id(_unique_trace_name), + type="GENERATION", + fields="core,io", + ).data + if _trace_data: + break + await asyncio.sleep(3) print(f"_trace_data: {_trace_data}") - assert _trace_data[0].input == { + assert json.loads(_trace_data[0].input) == { "messages": [{"content": "redacted-by-litellm", "role": "user"}] } - assert _trace_data[0].output == { + assert json.loads(_trace_data[0].output) == { "role": "assistant", "content": "redacted-by-litellm", "function_call": None, "tool_calls": None, + "provider_specific_fields": None, } except Exception as e: diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion.json index e252e8a128f..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion.json @@ -1,99 +1,38 @@ { - "batch": [ - { - "id": "7e00e081-468b-4fe9-a409-eb12ac7d3d2d", - "type": "trace-create", - "body": { - "id": "litellm-test-793c217f-9417-4e77-84a7-8dcc16e5b72b", - "timestamp": "2025-01-16T19:28:55.124873Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-16T19:28:55.125002Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "b9ec2c0f-18df-46c7-9e90-624c60bf78ee", - "type": "generation-create", - "body": { - "name": "litellm-acompletion", - "startTime": "2025-01-16T11:28:54.796360-08:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 3.5e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 3.5e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-11-28-54-796360_chatcmpl-521e530f-5e29-4d0a-8d1a-58fca0a847c2", - "endTime": "2025-01-16T11:28:55.124353-08:00", - "completionStartTime": "2025-01-16T11:28:55.124353-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 - }, - "traceId": "litellm-test-6a51ae70-a4e7-499e-afcd-dce2a3b31850" - }, - "timestamp": "2025-01-16T19:28:55.125258Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-03734ab3-8790-4c09-b5fb-8c3b663413b6" + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" + } + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_bedrock_call.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_bedrock_call.json index dd49d9751f1..6f359380245 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_bedrock_call.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_bedrock_call.json @@ -1,85 +1,31 @@ { - "batch": [ - { - "id": "3c9b544f-ef3f-449e-8ec1-763acbb56bec", - "type": "trace-create", - "body": { - "id": "litellm-test-c4c1c850-e8c9-4b16-b5a4-bff2bf9fa4f6", - "timestamp": "2025-05-26T21:13:16.796768Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "tags": [] - }, - "timestamp": "2025-05-26T21:13:16.796875Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 6e-05 }, - { - "id": "90e6bc70-05d9-4444-8b87-4523a9a54c17", - "type": "generation-create", - "body": { - "traceId": "litellm-test-c4c1c850-e8c9-4b16-b5a4-bff2bf9fa4f6", - "name": "litellm-acompletion", - "startTime": "2025-05-26T14:13:16.469836-07:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": null, - "response_cost": 6e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "usage_object": null - }, - "litellm_response_cost": 6e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "level": "DEFAULT", - "id": "time-14-13-16-469836_chatcmpl-3803a9e9-aa68-4493-94d9-247f354830d6", - "endTime": "2025-05-26T14:13:16.795438-07:00", - "completionStartTime": "2025-05-26T14:13:16.795438-07:00", - "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "modelParameters": { - "aws_region": "us-east-1" - }, - "usage": { - "input": 10, - "output": 10, - "unit": "TOKENS", - "totalCost": 6e-05 - }, - "usageDetails": { - "input": 10, - "output": 10, - "total": 20, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-05-26T21:13:16.797156Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-3bfc4db9-217f-48e9-92e0-142566e3c204" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "langfuse.observation.model.parameters": { + "aws_region": "us-east-1" + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 10, + "total": 20, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_complex_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_complex_metadata.json index 15794de7a07..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_complex_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_complex_metadata.json @@ -1,138 +1,38 @@ { - "batch": [ - { - "id": "9ee9100b-c4aa-4e40-a10d-bc189f8b4242", - "type": "trace-create", - "body": { - "id": "litellm-test-c414db10-dd68-406e-9d9e-03839bc2f346", - "timestamp": "2025-01-22T17:27:51.702596Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:27:51.702716Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "f8d20489-ed58-429f-b609-87380e223746", - "type": "generation-create", - "body": { - "traceId": "litellm-test-c414db10-dd68-406e-9d9e-03839bc2f346", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:27:51.150898-08:00", - "metadata": { - "string_value": "hello", - "int_value": 42, - "float_value": 3.14, - "bool_value": true, - "nested_dict": { - "key1": "value1", - "key2": { - "inner_key": "inner_value" - } - }, - "list_value": [ - 1, - 2, - 3 - ], - "set_value": [ - 1, - 2, - 3 - ], - "complex_list": [ - { - "dict_in_list": "value" - }, - "simple_string", - [ - 1, - 2, - 3 - ] - ], - "user": { - "name": "John", - "age": 30, - "tags": [ - "customer", - "active" - ] - }, - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-27-51-150898_chatcmpl-b783291c-dc76-4660-bfef-b79be9d54e57", - "endTime": "2025-01-22T09:27:51.702048-08:00", - "completionStartTime": "2025-01-22T09:27:51.702048-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:27:51.703046Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_langfuse_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_langfuse_metadata.json index 8d5d08894ef..5ed49cde972 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_langfuse_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_langfuse_metadata.json @@ -1,116 +1,62 @@ { - "batch": [ - { - "id": "872a0a1c-4328-431b-80b6-fd55a8a44477", - "type": "trace-create", - "body": { - "id": "litellm-test-533ffb2d-a0a3-45b5-911c-7940466cdc8e", - "timestamp": "2025-01-22T17:19:11.234960Z", - "name": "test_trace_name", - "userId": "test_user_id", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "sessionId": "test_session_id", - "version": "test_trace_version", - "metadata": { - "test_key": "test_value" - }, - "tags": [ - "test_tag", - "test_tag_2" - ] - }, - "timestamp": "2025-01-22T17:19:11.235169Z" + "name": "test_generation_name", + "parent_span_id": "0d9cfbb24ef808cd", + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "18d6f044-e522-4376-96e0-7eec765677ed", - "type": "generation-create", - "body": { - "traceId": "litellm-test-533ffb2d-a0a3-45b5-911c-7940466cdc8e", - "name": "test_generation_name", - "startTime": "2025-01-22T09:19:10.957072-08:00", - "metadata": { - "tags": [ - "test_tag", - "test_tag_2" - ], - "parent_observation_id": "test_parent_observation_id", - "version": "test_version", - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 3.5e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 3.5e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "parentObservationId": "test_parent_observation_id", - "version": "test_version", - "id": "time-09-19-10-957072_chatcmpl-4da65aba-32e4-400d-aaa2-6bfe096d8141", - "endTime": "2025-01-22T09:19:11.234200-08:00", - "completionStartTime": "2025-01-22T09:19:11.234200-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:19:11.235541Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.release": "test_trace_release", + "langfuse.trace.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" + } + ] + }, + "langfuse.trace.metadata.test_key": "test_value", + "langfuse.trace.name": "test_trace_name", + "langfuse.trace.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.trace.tags": [ + "test_tag", + "test_tag_2" + ], + "langfuse.version": "test_trace_version", + "session.id": "test_session_id", + "user.id": "test_user_id" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_no_choices.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_no_choices.json index ff8419ee392..b5a0737cf39 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_no_choices.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_no_choices.json @@ -1,85 +1,31 @@ { - "batch": [ - { - "id": "1f1d7517-4602-4c59-a322-7fc0306f1b7a", - "type": "trace-create", - "body": { - "id": "litellm-test-dbadfdfc-f4e7-4f05-8992-984c37359166", - "timestamp": "2025-02-07T00:23:27.669634Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "tags": [] - }, - "timestamp": "2025-02-07T00:23:27.669809Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 1.9999999999999998e-05 }, - { - "id": "fbe610b6-f500-4c7d-8e34-d40a0e8c487b", - "type": "generation-create", - "body": { - "traceId": "litellm-test-dbadfdfc-f4e7-4f05-8992-984c37359166", - "name": "litellm-acompletion", - "startTime": "2025-02-06T16:23:27.220129-08:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 3.5e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 3.5e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "level": "DEFAULT", - "id": "time-16-23-27-220129_chatcmpl-565360d7-965f-4533-9c09-db789af77a7d", - "endTime": "2025-02-06T16:23:27.644253-08:00", - "completionStartTime": "2025-02-06T16:23:27.644253-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 10, - "unit": "TOKENS", - "totalCost": 1.9999999999999998e-05 - }, - "usageDetails": { - "input": 10, - "output": 10, - "total": 20, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-02-07T00:23:27.670175Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 10, + "total": 20, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_router.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_router.json index df99b11d26b..749796e0d04 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_router.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_router.json @@ -1,95 +1,33 @@ { - "batch": [ - { - "id": "45eb9b25-605c-4c4a-b2b3-8241e079cd31", - "type": "trace-create", - "body": { - "id": "litellm-test-32702f3d-8a1c-4912-a3d6-286e59a9c568", - "timestamp": "2025-05-24T17:01:19.408179Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "tags": [] - }, - "timestamp": "2025-05-24T17:01:19.408284Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 1.9999999999999998e-05 }, - { - "id": "9f5e9b7d-0cea-4776-b4b9-5c2e8f4bad3c", - "type": "generation-create", - "body": { - "traceId": "litellm-test-32702f3d-8a1c-4912-a3d6-286e59a9c568", - "name": "litellm-acompletion", - "startTime": "2025-05-24T10:01:19.142356-07:00", - "metadata": { - "model_group": "gpt-3.5-turbo", - "model_group_size": 1, - "deployment": "gpt-3.5-turbo", - "model_info": { - "id": "0f1cd8f9e6a22e499303d479486395563ea04decade83fe7334dc2f079a857c2", - "db_model": false - }, - "api_base": null, - "hidden_params": { - "model_id": "0f1cd8f9e6a22e499303d479486395563ea04decade83fe7334dc2f079a857c2", - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 3.5e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 3.5e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "level": "DEFAULT", - "id": "time-10-01-19-142356_chatcmpl-16b215b7-e51e-47b0-8fe5-9dd6f226fda1", - "endTime": "2025-05-24T10:01:19.406531-07:00", - "completionStartTime": "2025-05-24T10:01:19.406531-07:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "stream": false, - "max_retries": 0, - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 10, - "unit": "TOKENS", - "totalCost": 1.9999999999999998e-05 - }, - "usageDetails": { - "input": 10, - "output": 10, - "total": 20, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-05-24T17:01:19.408586Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-3bfc4db9-217f-48e9-92e0-142566e3c204" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "stream": false, + "max_retries": 0, + "extra_body": "{}" + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 10, + "total": 20, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags.json index fd3d3194a5b..39e26f5957d 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags.json @@ -1,106 +1,42 @@ { - "batch": [ - { - "id": "42be960a-5dde-47df-9cbc-1fdd0fdcaa7d", - "type": "trace-create", - "body": { - "id": "litellm-test-f3ab679b-1e1d-43fd-9a9a-f11287aeb339", - "timestamp": "2025-01-22T15:31:28.963419Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [ - "test_tag", - "test_tag_2" - ] - }, - "timestamp": "2025-01-22T15:31:28.963706Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "5486df5a-3776-4adf-abd0-bd22e51f7fb4", - "type": "generation-create", - "body": { - "traceId": "litellm-test-f3ab679b-1e1d-43fd-9a9a-f11287aeb339", - "name": "litellm-acompletion", - "startTime": "2025-01-22T07:31:28.960749-08:00", - "metadata": { - "tags": [ - "test_tag", - "test_tag_2" - ], - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-07-31-28-960749_chatcmpl-f06338f0-8c49-45d8-be35-2854a89723c1", - "endTime": "2025-01-22T07:31:28.962389-08:00", - "completionStartTime": "2025-01-22T07:31:28.962389-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T15:31:28.964179Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion", + "langfuse.trace.tags": [ + "test_tag", + "test_tag_2" + ] } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags_stream.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags_stream.json index af15f351189..5223538f919 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags_stream.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_tags_stream.json @@ -1,106 +1,42 @@ { - "batch": [ - { - "id": "06b8fa9f-151b-4e74-9fbf-8af5222a7f40", - "type": "trace-create", - "body": { - "id": "litellm-test-54368a51-a382-493c-b0a8-3f1af23e18c4", - "timestamp": "2025-01-22T16:38:26.016582Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [ - "test_tag_stream", - "test_tag_2_stream" - ] - }, - "timestamp": "2025-01-22T16:38:26.016828Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "4ca1fd78-53e3-41b5-95d9-417b09e3f0eb", - "type": "generation-create", - "body": { - "traceId": "litellm-test-54368a51-a382-493c-b0a8-3f1af23e18c4", - "name": "litellm-acompletion", - "startTime": "2025-01-22T08:38:25.665692-08:00", - "metadata": { - "tags": [ - "test_tag_stream", - "test_tag_2_stream" - ], - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-08-38-25-665692_chatcmpl-8b67ffb8-4326-4e1b-bf4a-f70930c11c00", - "endTime": "2025-01-22T08:38:26.015666-08:00", - "completionStartTime": "2025-01-22T08:38:26.015666-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T16:38:26.017252Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion", + "langfuse.trace.tags": [ + "test_tag_stream", + "test_tag_2_stream" + ] } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json index 5998c52659c..d7d292390a0 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json @@ -1,83 +1,29 @@ { - "batch": [ - { - "id": "7d33d536-2730-4815-8957-80866c09c053", - "type": "trace-create", - "body": { - "id": "litellm-test-72861437-ff5b-4c48-89c0-a143534d9e7a", - "timestamp": "2025-05-26T21:15:40.610459Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "tags": [] - }, - "timestamp": "2025-05-26T21:15:40.610603Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "ebb5079c-7726-4adb-9616-e1862735e1d8", - "type": "generation-create", - "body": { - "traceId": "litellm-test-72861437-ff5b-4c48-89c0-a143534d9e7a", - "name": "litellm-acompletion", - "startTime": "2025-05-26T14:15:40.349639-07:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": null, - "response_cost": 3.5e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "vertex_ai/gemini-3-flash-preview", - "usage_object": null - }, - "litellm_response_cost": 3.5e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "level": "DEFAULT", - "id": "time-14-15-40-349639_chatcmpl-59a988d0-7ef1-4dc4-bc18-d2e78961817f", - "endTime": "2025-05-26T14:15:40.607266-07:00", - "completionStartTime": "2025-05-26T14:15:40.607266-07:00", - "model": "gemini-3-flash-preview", - "modelParameters": {}, - "usage": { - "input": 10, - "output": 10, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 10, - "total": 20, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-05-26T21:15:40.610953Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-3bfc4db9-217f-48e9-92e0-142566e3c204" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gemini-3-flash-preview", + "langfuse.observation.model.parameters": {}, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 10, + "total": 20, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata.json index 82a115a0899..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata.json @@ -1,113 +1,38 @@ { - "batch": [ - { - "id": "ddf567e5-a1b5-4e38-8a7c-f48bc847f721", - "type": "trace-create", - "body": { - "id": "litellm-test-46551fc7-c916-4a83-aeef-4274b5582ce1", - "timestamp": "2025-01-22T17:59:39.367430Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:59:39.367707Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "d3eb2c9e-e123-419d-b27b-c8283a505ae8", - "type": "generation-create", - "body": { - "traceId": "litellm-test-46551fc7-c916-4a83-aeef-4274b5582ce1", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:59:39.362554-08:00", - "metadata": { - "int": 42, - "str": "hello", - "list": [ - 1, - 2, - 3 - ], - "set": [ - 4, - 5 - ], - "dict": { - "nested": "value" - }, - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-59-39-362554_chatcmpl-d20ba1d9-cda6-4773-822e-921ebcd426a0", - "endTime": "2025-01-22T09:59:39.365756-08:00", - "completionStartTime": "2025-01-22T09:59:39.365756-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:59:39.368310Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata_2.json b/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata_2.json index 33e6b01bee3..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata_2.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/complex_metadata_2.json @@ -1,105 +1,38 @@ { - "batch": [ - { - "id": "ea3d694a-ce6b-417e-86e3-23ac17c6f6c6", - "type": "trace-create", - "body": { - "id": "litellm-test-38dcf290-8742-4fc5-ad03-c5d47e91dec0", - "timestamp": "2025-01-22T18:06:50.959206Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T18:06:50.959409Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "5fe03133-5798-4f87-8eec-ae0264f1eccc", - "type": "generation-create", - "body": { - "traceId": "litellm-test-38dcf290-8742-4fc5-ad03-c5d47e91dec0", - "name": "litellm-acompletion", - "startTime": "2025-01-22T10:06:50.957097-08:00", - "metadata": { - "list": [ - "list", - "not", - "a", - "dict" - ], - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-10-06-50-957097_chatcmpl-62d4ad7c-291b-4fc7-a8a4-3ed0fc3912a5", - "endTime": "2025-01-22T10:06:50.958374-08:00", - "completionStartTime": "2025-01-22T10:06:50.958374-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T18:06:50.959850Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/empty_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/empty_metadata.json index f4040f1f8fc..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/empty_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/empty_metadata.json @@ -1,99 +1,38 @@ { - "batch": [ - { - "id": "28d0c943-284b-4151-bf0d-8acf0f449865", - "type": "trace-create", - "body": { - "id": "litellm-test-d9506624-457c-40bc-9a37-578b896fa22a", - "timestamp": "2025-01-22T17:59:32.888622Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:59:32.888940Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "384e9fb4-3516-47b2-a4ae-1666337ec4a7", - "type": "generation-create", - "body": { - "traceId": "litellm-test-d9506624-457c-40bc-9a37-578b896fa22a", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:59:32.878577-08:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-59-32-878577_chatcmpl-1195f870-fd4d-4e38-8dc8-99dd3da5ab0b", - "endTime": "2025-01-22T09:59:32.880691-08:00", - "completionStartTime": "2025-01-22T09:59:32.880691-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:59:32.889548Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_function.json b/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_function.json index 77ca252c86d..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_function.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_function.json @@ -1,99 +1,38 @@ { - "batch": [ - { - "id": "88b1898a-cc5d-4e8e-93bc-3e71300c5e8d", - "type": "trace-create", - "body": { - "id": "litellm-test-a46356d9-ecff-44c8-a3da-fed3588b5128", - "timestamp": "2025-01-22T17:59:36.162545Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:59:36.162702Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "96bb77a6-a350-431b-bfd8-425491259728", - "type": "generation-create", - "body": { - "traceId": "litellm-test-a46356d9-ecff-44c8-a3da-fed3588b5128", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:59:36.161090-08:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-59-36-161090_chatcmpl-1ee988c9-9133-4655-bbe4-b97ffb6e3dc9", - "endTime": "2025-01-22T09:59:36.161959-08:00", - "completionStartTime": "2025-01-22T09:59:36.161959-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:59:36.162997Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_lock.json b/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_lock.json index f4040f1f8fc..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_lock.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/metadata_with_lock.json @@ -1,99 +1,38 @@ { - "batch": [ - { - "id": "28d0c943-284b-4151-bf0d-8acf0f449865", - "type": "trace-create", - "body": { - "id": "litellm-test-d9506624-457c-40bc-9a37-578b896fa22a", - "timestamp": "2025-01-22T17:59:32.888622Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:59:32.888940Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "384e9fb4-3516-47b2-a4ae-1666337ec4a7", - "type": "generation-create", - "body": { - "traceId": "litellm-test-d9506624-457c-40bc-9a37-578b896fa22a", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:59:32.878577-08:00", - "metadata": { - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-59-32-878577_chatcmpl-1195f870-fd4d-4e38-8dc8-99dd3da5ab0b", - "endTime": "2025-01-22T09:59:32.880691-08:00", - "completionStartTime": "2025-01-22T09:59:32.880691-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:59:32.889548Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/nested_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/nested_metadata.json index f4a1bb9dcea..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/nested_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/nested_metadata.json @@ -1,105 +1,38 @@ { - "batch": [ - { - "id": "44f179be-e3b9-486f-986f-030fc50614f0", - "type": "trace-create", - "body": { - "id": "litellm-test-8a04085c-1859-48fa-9fd8-1ec487fe455e", - "timestamp": "2025-01-22T17:55:28.854927Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:55:28.855187Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "2175ee64-58a3-41ab-96df-405b76695f5f", - "type": "generation-create", - "body": { - "traceId": "litellm-test-8a04085c-1859-48fa-9fd8-1ec487fe455e", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:55:28.852503-08:00", - "metadata": { - "a": { - "nested_a": 1 - }, - "b": { - "nested_b": 2 - }, - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-55-28-852503_chatcmpl-131cf0da-a47b-4cd1-850b-50fa077362ac", - "endTime": "2025-01-22T09:55:28.853979-08:00", - "completionStartTime": "2025-01-22T09:55:28.853979-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:55:28.855732Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata.json b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata.json index d895378e2c6..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata.json @@ -1,105 +1,38 @@ { - "batch": [ - { - "id": "02c74119-76b7-4f79-91cb-c55f1495c100", - "type": "trace-create", - "body": { - "id": "litellm-test-e58116c7-ead0-417e-9f86-b35f1e5bc242", - "timestamp": "2025-01-22T17:53:53.754012Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:53:53.754178Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "097968e0-52e9-46b5-9e8e-e6e08dd00e72", - "type": "generation-create", - "body": { - "traceId": "litellm-test-e58116c7-ead0-417e-9f86-b35f1e5bc242", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:53:53.752422-08:00", - "metadata": { - "a": { - "nested_a": 1 - }, - "b": { - "nested_b": 2 - }, - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-53-53-752422_chatcmpl-e99bc1d3-a393-493f-8afe-4507c0acff15", - "endTime": "2025-01-22T09:53:53.753431-08:00", - "completionStartTime": "2025-01-22T09:53:53.753431-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:53:53.754511Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata2.json b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata2.json index 87eba33cfff..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata2.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata2.json @@ -1,109 +1,38 @@ { - "batch": [ - { - "id": "1a55383a-e6fa-41f9-81fe-e7aa58c55f40", - "type": "trace-create", - "body": { - "id": "litellm-test-08fd1578-4a67-49b4-ac23-2dff1c112c80", - "timestamp": "2025-01-22T17:56:35.477276Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:56:35.477571Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "13ba66e8-f72b-4f57-a6cc-57c0be2829b1", - "type": "generation-create", - "body": { - "traceId": "litellm-test-08fd1578-4a67-49b4-ac23-2dff1c112c80", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:56:35.474752-08:00", - "metadata": { - "a": [ - 1, - 2, - 3 - ], - "b": [ - 4, - 5, - 6 - ], - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-56-35-474752_chatcmpl-9b152610-3d1e-4731-a84e-d0341ea69a0f", - "endTime": "2025-01-22T09:56:35.476236-08:00", - "completionStartTime": "2025-01-22T09:56:35.476236-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:56:35.478171Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata3.json b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata3.json index dd3bb4a301f..906b1a42a8f 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata3.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/simple_metadata3.json @@ -1,113 +1,38 @@ { - "batch": [ - { - "id": "7fb1f295-a7af-47af-afbd-e2f2d08280aa", - "type": "trace-create", - "body": { - "id": "litellm-test-c3acc34b-3c06-4868-bcee-87a3c4c1367e", - "timestamp": "2025-01-22T17:56:38.786515Z", - "name": "litellm-acompletion", - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "tags": [] - }, - "timestamp": "2025-01-22T17:56:38.786742Z" + "name": "litellm-acompletion", + "parent_span_id": null, + "attributes": { + "langfuse.observation.cost_details": { + "total": 3.5e-05 }, - { - "id": "412870bc-fc50-4426-a0dc-9e8b016e14bb", - "type": "generation-create", - "body": { - "traceId": "litellm-test-c3acc34b-3c06-4868-bcee-87a3c4c1367e", - "name": "litellm-acompletion", - "startTime": "2025-01-22T09:56:38.784548-08:00", - "metadata": { - "a": [ - 1, - 2 - ], - "b": [ - 3, - 4 - ], - "c": { - "d": [ - 5, - 6 - ] - }, - "hidden_params": { - "model_id": null, - "cache_key": null, - "api_base": "https://api.openai.com", - "response_cost": 5.4999999999999995e-05, - "additional_headers": {}, - "litellm_overhead_time_ms": null, - "batch_models": null, - "litellm_model_name": "gpt-3.5-turbo", - "usage_object": null - }, - "litellm_response_cost": 5.4999999999999995e-05, - "cache_hit": false, - "requester_metadata": {} - }, - "input": { - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] - }, - "output": { - "content": "Hello! How can I assist you today?", - "role": "assistant", - "tool_calls": null, - "function_call": null, - "provider_specific_fields": null - }, - "level": "DEFAULT", - "id": "time-09-56-38-784548_chatcmpl-438c8727-86b3-44d9-9b46-42330922cf50", - "endTime": "2025-01-22T09:56:38.785762-08:00", - "completionStartTime": "2025-01-22T09:56:38.785762-08:00", - "model": "gpt-3.5-turbo", - "modelParameters": { - "extra_body": "{}" - }, - "usage": { - "input": 10, - "output": 20, - "unit": "TOKENS", - "totalCost": 3.5e-05 - }, - "usageDetails": { - "input": 10, - "output": 20, - "total": 30, - "cache_creation_input_tokens": 0, - "cache_read_input_tokens": 0 + "langfuse.observation.input": { + "messages": [ + { + "role": "user", + "content": "Hello!" } - }, - "timestamp": "2025-01-22T17:56:38.787196Z" - } - ], - "metadata": { - "batch_size": 2, - "sdk_integration": "litellm", - "sdk_name": "python", - "sdk_version": "2.44.1", - "public_key": "pk-lf-e02aaea3-8668-4c9f-8c69-771a4ea1f5c9" + ] + }, + "langfuse.observation.level": "DEFAULT", + "langfuse.observation.model.name": "gpt-3.5-turbo", + "langfuse.observation.model.parameters": { + "extra_body": "{}" + }, + "langfuse.observation.output": { + "content": "Hello! How can I assist you today?", + "role": "assistant", + "tool_calls": null, + "function_call": null, + "provider_specific_fields": null + }, + "langfuse.observation.type": "generation", + "langfuse.observation.usage_details": { + "input": 10, + "output": 20, + "total": 30, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0 + }, + "langfuse.trace.name": "litellm-acompletion" } -} \ No newline at end of file +} diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py index 2346a5ee047..92466a9470c 100644 --- a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py +++ b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py @@ -1,6 +1,3 @@ -import sys -from types import ModuleType, SimpleNamespace - import litellm from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler @@ -51,37 +48,29 @@ def test_resolve_langfuse_credentials_keeps_env_for_global_config(monkeypatch): assert host == "https://admin-configured.example" -def test_upstream_langfuse_debug_env_is_passed(monkeypatch): +def test_upstream_langfuse_env_only_warns_and_opens_no_second_channel(monkeypatch, caplog): + """UPSTREAM_LANGFUSE_* configured a second v2 ingestion client. v4 has one export channel per + credential set, so the values are ignored with a startup warning and never build anything.""" + from litellm.integrations.langfuse import langfuse_sdk from litellm.integrations.langfuse.langfuse import LangFuseLogger - class FakeLangfuse: - instances = [] - - def __init__(self, **kwargs): - self.kwargs = kwargs - FakeLangfuse.instances.append(self) - - fake_langfuse_module = ModuleType("langfuse") - fake_langfuse_module.Langfuse = FakeLangfuse - fake_langfuse_module.version = SimpleNamespace(__version__="2.6.0") - - monkeypatch.setitem(sys.modules, "langfuse", fake_langfuse_module) monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) + monkeypatch.setattr(langfuse_sdk, "_TRACING", {}) monkeypatch.setenv("LANGFUSE_MOCK", "true") monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret") monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public") monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") - monkeypatch.setenv("UPSTREAM_LANGFUSE_RELEASE", "release") - monkeypatch.setenv("UPSTREAM_LANGFUSE_DEBUG", "true") - logger = LangFuseLogger( - langfuse_public_key="public", - langfuse_secret="secret", - langfuse_host="https://langfuse.example", - ) + with caplog.at_level("WARNING", logger="LiteLLM"): + logger = LangFuseLogger( + langfuse_public_key="public", + langfuse_secret="secret", + langfuse_host="https://langfuse.example", + ) - assert logger.upstream_langfuse_debug == "true" - assert FakeLangfuse.instances[-1].kwargs["debug"] is True + assert any("UPSTREAM_LANGFUSE_* is no longer supported" in record.getMessage() for record in caplog.records) + assert [lease.tracing for lease in langfuse_sdk._TRACING.values()] == [logger.tracing] + assert all(key.public_key == "public" for key in langfuse_sdk._TRACING) def test_langfuse_handler_accepts_secret_key_alias(monkeypatch): diff --git a/tests/logging_callback_tests/test_langfuse_e2e_test.py b/tests/logging_callback_tests/test_langfuse_e2e_test.py index 76ebd2b9a28..61a93a175d9 100644 --- a/tests/logging_callback_tests/test_langfuse_e2e_test.py +++ b/tests/logging_callback_tests/test_langfuse_e2e_test.py @@ -1,166 +1,140 @@ import asyncio -import copy import json import logging import os import threading -from typing import Any, Optional +from collections.abc import Mapping +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.common.v1.common_pb2 import AnyValue logging.basicConfig(level=logging.DEBUG) import litellm -from litellm import completion -from litellm.caching import InMemoryCache +from litellm.integrations.langfuse.langfuse_sdk import resolve_observation_id, resolve_trace_id from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler litellm.num_retries = 3 litellm.success_callback = ["langfuse"] os.environ["LANGFUSE_DEBUG"] = "True" -import time import pytest import pytest_asyncio +LANGFUSE_EXPORT_POST: Final = "litellm.llms.custom_httpx.http_handler.HTTPHandler.post" +LANGFUSE_EXPORT_PATH: Final = "/api/public/otel/v1/traces" + +_PER_RUN_ATTRIBUTES: Final = frozenset( + { + "langfuse.observation.completion_start_time", + "langfuse.observation.metadata.applied_guardrails", + "langfuse.observation.metadata.cache_hit", + "langfuse.observation.metadata.hidden_params", + "langfuse.observation.metadata.litellm_call_id", + "langfuse.observation.metadata.litellm_response_cost", + "langfuse.observation.metadata.requester_metadata", + "langfuse.observation.metadata.response_id", + "langfuse.observation.metadata.usage_object", + } +) + + +def _decode_attribute(value: AnyValue) -> object: + match value.WhichOneof("value"): + case "string_value": + try: + return json.loads(value.string_value) + except json.JSONDecodeError: + return value.string_value + case "bool_value": + return value.bool_value + case "int_value": + return value.int_value + case "double_value": + return value.double_value + case "array_value": + return [_decode_attribute(item) for item in value.array_value.values] + case _: + return None + + +def _exported_spans(mock_post: MagicMock) -> list[dict[str, object]]: + spans: list[dict[str, object]] = [] + for call in mock_post.call_args_list: + url: str = call.args[0] if call.args else call.kwargs["url"] + assert url.endswith(LANGFUSE_EXPORT_PATH), url + request = ExportTraceServiceRequest.FromString(call.kwargs["data"]) + for resource_spans in request.resource_spans: + for scope_spans in resource_spans.scope_spans: + for span in scope_spans.spans: + spans.append( + { + "name": span.name, + "trace_id": span.trace_id.hex(), + "span_id": span.span_id.hex(), + "parent_span_id": span.parent_span_id.hex() or None, + "attributes": { + attribute.key: _decode_attribute(attribute.value) for attribute in span.attributes + }, + } + ) + return spans + + +def _comparable(span: Mapping[str, object]) -> dict[str, object]: + attributes = span["attributes"] + assert isinstance(attributes, dict) + return { + "name": span["name"], + "parent_span_id": span["parent_span_id"], + "attributes": {key: value for key, value in sorted(attributes.items()) if key not in _PER_RUN_ATTRIBUTES}, + } + def assert_langfuse_request_matches_expected( - actual_request_body: dict, + spans: list[dict[str, object]], expected_file_name: str, - trace_id: Optional[str] = None, + trace_id: str, ): - """ - Helper function to compare actual Langfuse request body with expected JSON file. - - Args: - actual_request_body (dict): The actual request body received from the API call - expected_file_name (str): Name of the JSON file containing expected request body (e.g., "transcription.json") - """ - # Get the current directory and read the expected request body + """Compare the generation langfuse exported for ``trace_id`` with the expected JSON file.""" pwd = os.path.dirname(os.path.realpath(__file__)) - expected_body_path = os.path.join( - pwd, "langfuse_expected_request_body", expected_file_name - ) - + expected_body_path = os.path.join(pwd, "langfuse_expected_request_body", expected_file_name) with open(expected_body_path, "r") as f: - expected_request_body = json.load(f) + expected_generation = json.load(f) - # Filter out events that don't match the trace_id - if trace_id: - actual_request_body["batch"] = [ - item - for item in actual_request_body["batch"] - if (item["type"] == "trace-create" and item["body"].get("id") == trace_id) - or ( - item["type"] == "generation-create" - and item["body"].get("traceId") == trace_id - ) - ] - - # When aggregating from multiple flush cycles, deduplicate by keeping - # only one trace-create and one generation-create per trace_id. - seen_types: dict = {} - deduped_batch: list = [] - for item in actual_request_body["batch"]: - item_type = item["type"] - if item_type not in seen_types: - seen_types[item_type] = True - deduped_batch.append(item) - actual_request_body["batch"] = deduped_batch - - # Ensure canonical order: trace-create first, generation-create second - actual_request_body["batch"].sort( - key=lambda x: 0 if x["type"] == "trace-create" else 1 + otel_trace_id: Final = resolve_trace_id(trace_id) + generations: Final = [ + span + for span in spans + if span["trace_id"] == otel_trace_id and span["attributes"]["langfuse.observation.type"] == "generation" # pyright: ignore[reportIndexIssue] # built as dict in _exported_spans + ] + assert len(generations) == 1, ( + f"Expected exactly one generation for trace_id={trace_id} ({otel_trace_id}), " + f"got {len(generations)}. Spans: {json.dumps(spans, indent=2)}" ) - print( - "actual_request_body after filtering", json.dumps(actual_request_body, indent=4) + actual_generation: Final = _comparable(generations[0]) + assert actual_generation == expected_generation, ( + f"Difference in exported generation: {json.dumps(actual_generation, indent=2)} " + f"!= {json.dumps(expected_generation, indent=2)}" ) - assert len(actual_request_body["batch"]) >= 2, ( - f"Expected at least 2 batch items (trace-create + generation-create) " - f"after filtering by trace_id={trace_id}, " - f"but got {len(actual_request_body['batch'])}. " - f"Items: {json.dumps(actual_request_body['batch'], indent=2)}" - ) - - # Replace dynamic values in actual request body - for item in actual_request_body["batch"]: - - # Replace IDs with expected IDs - if item["type"] == "trace-create": - item["id"] = expected_request_body["batch"][0]["id"] - item["body"]["id"] = expected_request_body["batch"][0]["body"]["id"] - item["timestamp"] = expected_request_body["batch"][0]["timestamp"] - item["body"]["timestamp"] = expected_request_body["batch"][0]["body"][ - "timestamp" - ] - elif item["type"] == "generation-create": - item["id"] = expected_request_body["batch"][1]["id"] - item["body"]["id"] = expected_request_body["batch"][1]["body"]["id"] - item["timestamp"] = expected_request_body["batch"][1]["timestamp"] - item["body"]["startTime"] = expected_request_body["batch"][1]["body"][ - "startTime" - ] - item["body"]["endTime"] = expected_request_body["batch"][1]["body"][ - "endTime" - ] - item["body"]["completionStartTime"] = expected_request_body["batch"][1][ - "body" - ]["completionStartTime"] - if trace_id is None: - print("popping traceId") - item["body"].pop("traceId") - else: - item["body"]["traceId"] = trace_id - expected_request_body["batch"][1]["body"]["traceId"] = trace_id - - # Replace SDK version with expected version - actual_request_body["batch"][0]["body"].pop("release", None) - actual_request_body["metadata"]["sdk_version"] = expected_request_body["metadata"][ - "sdk_version" - ] - # replace "public_key" with expected public key - actual_request_body["metadata"]["public_key"] = expected_request_body["metadata"][ - "public_key" - ] - actual_request_body["batch"][1]["body"]["metadata"] = expected_request_body[ - "batch" - ][1]["body"]["metadata"] - actual_request_body["metadata"]["sdk_integration"] = expected_request_body[ - "metadata" - ]["sdk_integration"] - actual_request_body["metadata"]["batch_size"] = expected_request_body["metadata"][ - "batch_size" - ] - # Assert the entire request body matches - assert ( - actual_request_body == expected_request_body - ), f"Difference in request bodies: {json.dumps(actual_request_body, indent=2)} != {json.dumps(expected_request_body, indent=2)}" - class TestLangfuseLogging: @pytest_asyncio.fixture async def mock_setup(self): """Common setup for Langfuse logging tests""" from litellm._uuid import uuid - from unittest.mock import AsyncMock, patch - import httpx - # Create a mock Response object - mock_response = AsyncMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = {"status": "success"} - - # Create mock for httpx.Client.post - mock_post = AsyncMock() - mock_post.return_value = mock_response + mock_post = MagicMock(return_value=MagicMock(ok=True, status_code=200)) litellm.set_verbose = True litellm.success_callback = ["langfuse"] - return {"trace_id": f"litellm-test-{str(uuid.uuid4())}", "mock_post": mock_post} + return {"trace_id": f"litellm-test-{uuid.uuid4()!s}", "mock_post": mock_post} async def _verify_langfuse_call( self, @@ -168,41 +142,16 @@ class TestLangfuseLogging: expected_file_name: str, trace_id: str, ): - """Helper method to verify Langfuse API calls""" - await asyncio.sleep(3) - - # Verify at least one call was made - assert mock_post.call_count >= 1 - - # Aggregate batch items from ALL calls — the Langfuse SDK may split - # trace-create and generation-create across separate HTTP flushes. - langfuse_url = "https://us.cloud.langfuse.com/api/public/ingestion" - all_batch_items: list = [] - metadata: Optional[dict] = None - for call in mock_post.call_args_list: - url = call[0][0] - if url != langfuse_url: - continue - request_body = call[1].get("content") - if request_body: - body = json.loads(request_body) - all_batch_items.extend(body.get("batch", [])) - if metadata is None: - metadata = body.get("metadata") - - assert len(all_batch_items) > 0, "No Langfuse ingestion calls found" - assert metadata is not None, "No metadata found in Langfuse calls" - - actual_request_body = { - "batch": all_batch_items, - "metadata": metadata, - } - - print("\nMocked Request Details (aggregated from all calls):") - print(f"Request Body: {json.dumps(actual_request_body, indent=4)}") + """Wait for the batch processor to export, then compare the generation it shipped.""" + otel_trace_id: Final = resolve_trace_id(trace_id) + for _ in range(100): + if any(span["trace_id"] == otel_trace_id for span in _exported_spans(mock_post)): + break + await asyncio.sleep(0.1) + assert mock_post.call_count >= 1, "langfuse exported nothing" assert_langfuse_request_matches_expected( - actual_request_body, + _exported_spans(mock_post), expected_file_name, trace_id, ) @@ -212,23 +161,21 @@ class TestLangfuseLogging: async def test_langfuse_logging_completion(self, mock_setup): """Test Langfuse logging for chat completion""" setup = mock_setup - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], mock_response="Hello! How can I assist you today?", metadata={"trace_id": setup["trace_id"]}, ) - await self._verify_langfuse_call( - setup["mock_post"], "completion.json", setup["trace_id"] - ) + await self._verify_langfuse_call(setup["mock_post"], "completion.json", setup["trace_id"]) @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) async def test_langfuse_logging_completion_with_tags(self, mock_setup): """Test Langfuse logging for chat completion with tags""" setup = mock_setup - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], @@ -238,16 +185,14 @@ class TestLangfuseLogging: "tags": ["test_tag", "test_tag_2"], }, ) - await self._verify_langfuse_call( - setup["mock_post"], "completion_with_tags.json", setup["trace_id"] - ) + await self._verify_langfuse_call(setup["mock_post"], "completion_with_tags.json", setup["trace_id"]) @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) async def test_langfuse_logging_completion_with_tags_stream(self, mock_setup): """Test Langfuse logging for chat completion with tags""" setup = mock_setup - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], @@ -263,12 +208,33 @@ class TestLangfuseLogging: setup["trace_id"], ) + @pytest.mark.asyncio + @pytest.mark.flaky(retries=3, delay=1) + async def test_langfuse_generation_id_metadata_names_the_exported_observation(self, mock_setup): + """v2 let callers pick the generation id; v4 only has span ids, so the requested id must become one.""" + setup = mock_setup + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): + await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hello!"}], + mock_response="Hello! How can I assist you today?", + metadata={"trace_id": setup["trace_id"], "generation_id": "my-generation"}, + ) + await self._verify_langfuse_call(setup["mock_post"], "completion.json", setup["trace_id"]) + + generation: Final = next( + span + for span in _exported_spans(setup["mock_post"]) + if span["trace_id"] == resolve_trace_id(setup["trace_id"]) + ) + assert generation["span_id"] == resolve_observation_id("my-generation") + @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) async def test_langfuse_logging_completion_with_langfuse_metadata(self, mock_setup): """Test Langfuse logging for chat completion with metadata for langfuse""" setup = mock_setup - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], @@ -297,12 +263,12 @@ class TestLangfuseLogging: @pytest.mark.flaky(retries=3, delay=1) async def test_langfuse_logging_with_non_serializable_metadata(self, mock_setup): """Test Langfuse logging with metadata that requires preparation (Pydantic models, sets, etc)""" - from pydantic import BaseModel - from typing import Set import datetime + from pydantic import BaseModel + class UserPreferences(BaseModel): - favorite_colors: Set[str] + favorite_colors: set[str] last_login: datetime.datetime settings: dict @@ -325,8 +291,8 @@ class TestLangfuseLogging: "trace_id": setup["trace_id"], } - with patch("httpx.Client.post", setup["mock_post"]): - response = await litellm.acompletion( + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): + await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], mock_response="Hello! How can I assist you today?", @@ -375,18 +341,14 @@ class TestLangfuseLogging: ], ) @pytest.mark.flaky(retries=6, delay=1) - async def test_langfuse_logging_with_various_metadata_types( - self, mock_setup, test_metadata, response_json_file - ): + async def test_langfuse_logging_with_various_metadata_types(self, mock_setup, test_metadata, response_json_file): """Test Langfuse logging with various metadata types including non-serializable objects""" - import threading - setup = mock_setup if test_metadata is not None: test_metadata["trace_id"] = setup["trace_id"] - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello!"}], @@ -402,13 +364,11 @@ class TestLangfuseLogging: @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) - async def test_langfuse_logging_completion_with_malformed_llm_response( - self, mock_setup - ): + async def test_langfuse_logging_completion_with_malformed_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup litellm._turn_on_debug() - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], usage=litellm.Usage( @@ -426,19 +386,15 @@ class TestLangfuseLogging: mock_response=mock_response, metadata={"trace_id": setup["trace_id"]}, ) - await self._verify_langfuse_call( - setup["mock_post"], "completion_with_no_choices.json", setup["trace_id"] - ) + await self._verify_langfuse_call(setup["mock_post"], "completion_with_no_choices.json", setup["trace_id"]) @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) - async def test_langfuse_logging_completion_with_bedrock_llm_response( - self, mock_setup - ): + async def test_langfuse_logging_completion_with_bedrock_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup litellm._turn_on_debug() - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], usage=litellm.Usage( @@ -467,13 +423,11 @@ class TestLangfuseLogging: @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) - async def test_langfuse_logging_completion_with_vertex_llm_response( - self, mock_setup - ): + async def test_langfuse_logging_completion_with_vertex_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup litellm._turn_on_debug() - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], usage=litellm.Usage( @@ -525,7 +479,7 @@ class TestLangfuseLogging: mock_async_client = AsyncHTTPHandler() mock_async_client.post = AsyncMock(return_value=mock_vllm_response) - with patch("httpx.Client.post", setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): await litellm.aembedding( model="hosted_vllm/BAAI/bge-small-en-v1.5", input=["Hello from litellm!"], @@ -539,9 +493,7 @@ class TestLangfuseLogging: actual_vllm_request = mock_async_client.post.call_args.kwargs["json"] pwd = os.path.dirname(os.path.realpath(__file__)) - expected_body_path = os.path.join( - pwd, "langfuse_expected_request_body", "embedding_with_vllm.json" - ) + expected_body_path = os.path.join(pwd, "langfuse_expected_request_body", "embedding_with_vllm.json") with open(expected_body_path, "r") as f: expected_vllm_request = json.load(f) @@ -568,7 +520,7 @@ class TestLangfuseLogging: } ] ) - with patch("httpx.Client.post", mock_setup["mock_post"]): + with patch(LANGFUSE_EXPORT_POST, mock_setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], usage=litellm.Usage( diff --git a/tests/logging_callback_tests/test_langfuse_unit_tests.py b/tests/logging_callback_tests/test_langfuse_unit_tests.py index 405b6e9e48e..61316204fc3 100644 --- a/tests/logging_callback_tests/test_langfuse_unit_tests.py +++ b/tests/logging_callback_tests/test_langfuse_unit_tests.py @@ -306,35 +306,63 @@ def test_get_langfuse_flush_interval(): def test_langfuse_e2e_sync(monkeypatch): - from litellm import completion - import litellm - import respx - import httpx + """A sync completion must reach langfuse over the wire, not just build a span. + + v4 exports OTLP over ``requests`` rather than the v2 ingestion endpoint over + httpx, so this stands up a real receiver and asserts langfuse posted to it. + """ + import threading import time + from http.server import BaseHTTPRequestHandler, HTTPServer - litellm.disable_aiohttp_transport = ( - True # since this uses respx, we need to set use_aiohttp_transport to False - ) + import litellm + from litellm import completion + from litellm.integrations.langfuse.langfuse import LangFuseLogger + from litellm.integrations.langfuse.langfuse_prompt_management import langfuse_client_init + from litellm.litellm_core_utils import litellm_logging - litellm._turn_on_debug() + received_paths = [] + + class _Receiver(BaseHTTPRequestHandler): + def do_POST(self): + received_paths.append(self.path) + self.rfile.read(int(self.headers.get("Content-Length") or 0)) + self.send_response(200) + self.send_header("Content-Length", "0") + self.end_headers() + + def log_message(self, *args): + pass + + server = HTTPServer(("127.0.0.1", 0), _Receiver) + threading.Thread(target=server.serve_forever, daemon=True).start() + monkeypatch.setenv("LANGFUSE_HOST", f"http://127.0.0.1:{server.server_port}") + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-e2e-sync") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-e2e-sync") monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setattr(litellm_logging, "langFuseLogger", None) + monkeypatch.setattr(litellm_logging, "in_memory_dynamic_logger_cache", DynamicLoggingCache()) + monkeypatch.setattr(litellm_logging, "_in_memory_loggers", []) + langfuse_client_init.cache_clear() - with respx.mock: - # Mock Langfuse - # Mock any Langfuse endpoint - langfuse_mock = respx.post( - "https://*.cloud.langfuse.com/api/public/ingestion" - ).mock(return_value=httpx.Response(200)) + try: completion( model="openai/my-fake-endpoint", messages=[{"role": "user", "content": "hello from litellm"}], stream=False, mock_response="Hello from litellm 2", ) + for logger in litellm.logging_callback_manager._get_all_callbacks(): + if isinstance(logger, LangFuseLogger): + logger.flush() + deadline = time.time() + 10 + while not received_paths and time.time() < deadline: + time.sleep(0.1) + finally: + server.shutdown() - time.sleep(3) - - assert langfuse_mock.called + assert received_paths, "langfuse exported nothing" + assert all(path.endswith("/api/public/otel/v1/traces") for path in received_paths) def test_get_chat_content_for_langfuse(): diff --git a/tests/test_keys.py b/tests/test_keys.py index 7a5b2502cfd..c1785b88822 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint): assert len(model_list["data"]) > 0 if model_access == "gpt-3.5-turbo": if model_endpoint == "/v1/models": - assert ( - len(model_list["data"]) == 1 - ), "model_access={}, model_access_level={}".format( + assert {entry["id"] for entry in model_list["data"]} == { + model_access, + "mistral-7b", + }, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format( model_access, model_access_level ) - assert model_list["data"][0]["id"] == model_access elif model_endpoint == "/model/info": assert isinstance(model_list["data"], list) assert len(model_list["data"]) == 1 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/caching/test_evicted_client_closer.py b/tests/test_litellm/caching/test_evicted_client_closer.py index 939be5f3d6b..7679e276621 100644 --- a/tests/test_litellm/caching/test_evicted_client_closer.py +++ b/tests/test_litellm/caching/test_evicted_client_closer.py @@ -10,9 +10,11 @@ never closed, because litellm does not own its lifecycle. import asyncio import gc import weakref +from unittest.mock import AsyncMock import httpx import pytest +from redis.asyncio import ConnectionPool, Redis from litellm.caching.evicted_client_closer import EvictedClientCloser from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler @@ -72,6 +74,33 @@ def make_closer(clock: FakeClock, grace_seconds: float = 60.0) -> EvictedClientC return EvictedClientCloser(grace_seconds=grace_seconds, clock=clock) +@pytest.mark.asyncio +async def test_redis_client_is_closed_only_after_its_subscription_releases_the_connection(): + closer = EvictedClientCloser(grace_seconds=0) + pool = ConnectionPool() + client = Redis.from_pool(pool) + closed = asyncio.Event() + connection = AsyncMock() + connection.disconnect.side_effect = closed.set + pool._available_connections.append(connection) + borrowed = pool.get_available_connection() + closer.mark_owned(client) + closer.schedule(client) + + closer.reap() + await asyncio.sleep(0.05) + + connection.disconnect.assert_not_awaited() + assert closer.pending_count == 1 + + await pool.release(borrowed) + closer.reap() + await asyncio.wait_for(closed.wait(), timeout=1) + + connection.disconnect.assert_awaited_once() + assert closer.pending_count == 0 + + async def _trickling_upstream(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: """Serves a chunked body slowly, so a request stays on the wire long enough to observe.""" await reader.read(4096) diff --git a/tests/test_litellm/caching/test_redis_cluster_cache.py b/tests/test_litellm/caching/test_redis_cluster_cache.py index 0763b5110d5..ba1deabd0e1 100644 --- a/tests/test_litellm/caching/test_redis_cluster_cache.py +++ b/tests/test_litellm/caching/test_redis_cluster_cache.py @@ -1,13 +1,20 @@ +import asyncio from importlib import import_module import json -from unittest.mock import MagicMock, patch +import ssl +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.testclient import TestClient +from redis.asyncio import Redis, RedisCluster +from redis.asyncio.cluster import ClusterNode +from redis.asyncio.connection import SSLConnection from litellm.caching.redis_cache import RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.caching.evicted_client_closer import EvictedClientCloser +from litellm.caching.llm_caching_handler import LLMClientCache @patch("litellm._redis.init_redis_cluster") @@ -175,3 +182,167 @@ def test_router_create_redis_cache_cluster_detection( with patch.object(RedisCache, "__init__", _mock_redis_cache_init): redis_cache = Router._create_redis_cache(cache_config) assert isinstance(redis_cache, expected_cache_type) + + +def _isolated_redis_cache(host: str) -> RedisCache: + """RedisCache whose sync client and pool are stubbed out.""" + with ( + patch("litellm._redis.get_redis_client", return_value=MagicMock()), + patch("litellm._redis.get_redis_connection_pool", return_value=MagicMock()), + ): + return RedisCache(host=host, port=6379) + + +def _cluster_for_pubsub(startup_node_host: str = "10.9.9.9") -> RedisCluster: + """Uninitialized RedisCluster carrying the connection kwargs a real one would.""" + return RedisCluster( + startup_nodes=[ClusterNode(host=startup_node_host, port=7000)], + password="cluster-secret", + socket_timeout=7.0, + ) + + +def test_init_pubsub_client_derives_a_node_client_for_cluster_backend() -> None: + """LIT-8543: a cluster-backed cache must return a pub/sub-capable client. + + The derived client pins a plain Redis connection pool to the cluster's + default node, inheriting the connection kwargs minus cluster-only keys. + """ + cache = _isolated_redis_cache("cluster-pubsub-default-node") + cluster = _cluster_for_pubsub() + node = ClusterNode(host="10.1.2.3", port=7001) + cluster.nodes_manager.default_node = node + cache.init_async_client = MagicMock(return_value=cluster) + + client = cache.init_pubsub_client() + + assert isinstance(client, Redis) and not isinstance(client, RedisCluster) + kwargs = client.connection_pool.connection_kwargs + assert kwargs["host"] == "10.1.2.3" + assert kwargs["port"] == 7001 + assert kwargs["password"] == "cluster-secret" + assert kwargs["socket_timeout"] == 7.0 + assert "response_callbacks" not in kwargs + + +def test_init_pubsub_client_falls_back_to_first_startup_node() -> None: + """Before cluster initialization there is no default node; the first + startup node is a valid pub/sub target.""" + cache = _isolated_redis_cache("cluster-pubsub-startup-fallback") + cluster = _cluster_for_pubsub(startup_node_host="10.8.8.8") + cache.init_async_client = MagicMock(return_value=cluster) + + client = cache.init_pubsub_client() + + assert isinstance(client, Redis) + assert client.connection_pool.connection_kwargs["host"] == "10.8.8.8" + + +def test_init_pubsub_client_returns_the_same_cached_client_on_repeat_calls() -> None: + cache = _isolated_redis_cache("cluster-pubsub-caching") + cluster = _cluster_for_pubsub() + cluster.nodes_manager.default_node = ClusterNode(host="10.1.2.3", port=7001) + cache.init_async_client = MagicMock(return_value=cluster) + + first = cache.init_pubsub_client() + second = cache.init_pubsub_client() + + assert first is second + + +def test_init_pubsub_client_returns_the_shared_async_client_for_standalone() -> None: + cache = _isolated_redis_cache("standalone-pubsub") + standalone = Redis() + cache.init_async_client = MagicMock(return_value=standalone) + + assert cache.init_pubsub_client() is standalone + + +def test_init_pubsub_client_preserves_tls_and_authentication() -> None: + cache = _isolated_redis_cache("cluster-pubsub-tls") + cluster = RedisCluster( + startup_nodes=[ClusterNode(host="redis.example.test", port=7000)], + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + username="pubsub-user", + password="test-password", + socket_connect_timeout=3.0, + socket_keepalive=True, + ) + cache.init_async_client = MagicMock(return_value=cluster) + + client = cache.init_pubsub_client() + connection = client.connection_pool.make_connection() + + assert isinstance(connection, SSLConnection) + assert connection.ssl_context.cert_reqs == ssl.CERT_REQUIRED + assert connection.ssl_context.check_hostname is True + assert connection.username == "pubsub-user" + assert connection.password == "test-password" + assert connection.socket_connect_timeout == 3.0 + assert connection.socket_keepalive is True + + +def test_init_pubsub_client_rejects_missing_nodes_and_can_retry() -> None: + cache = _isolated_redis_cache("cluster-pubsub-no-nodes") + cluster = _cluster_for_pubsub() + cluster.nodes_manager.startup_nodes = {} + cache.init_async_client = MagicMock(return_value=cluster) + + with pytest.raises(ValueError, match="no default node and no startup nodes"): + cache.init_pubsub_client() + + cluster.nodes_manager.default_node = ClusterNode(host="recovered.example.test", port=7001) + client = cache.init_pubsub_client() + + assert client.connection_pool.connection_kwargs["host"] == "recovered.example.test" + + +@pytest.mark.parametrize("close_fails", [False, True]) +@pytest.mark.parametrize("has_shared_pool", [False, True]) +def test_disconnect_closes_derived_pubsub_connections_even_when_pool_close_fails( + close_fails: bool, has_shared_pool: bool +) -> None: + cache = _isolated_redis_cache(f"cluster-pubsub-close-{close_fails}-{has_shared_pool}") + cache.async_redis_conn_pool = AsyncMock() if has_shared_pool else None + cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub()) + + async def exercise() -> None: + client = cache.init_pubsub_client() + connection = AsyncMock() + connection.disconnect.side_effect = ConnectionError("connection close failed") if close_fails else None + client.connection_pool._available_connections.append(connection) + + await cache.disconnect() + + connection.disconnect.assert_awaited_once() + if has_shared_pool: + cache.async_redis_conn_pool.disconnect.assert_awaited_once_with(inuse_connections=True) + cache.redis_client.close.assert_called_once() + + asyncio.run(exercise()) + + +def test_expired_pubsub_client_closes_connections_after_eviction() -> None: + cache = _isolated_redis_cache("cluster-pubsub-expired") + cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub()) + clients = LLMClientCache(evicted_client_closer=EvictedClientCloser(grace_seconds=0)) + + async def exercise() -> None: + client = cache.init_pubsub_client() + closed = asyncio.Event() + connection = AsyncMock() + connection.disconnect.side_effect = closed.set + client.connection_pool._available_connections.append(connection) + cache_key = clients.update_cache_key_with_event_loop(f"{cache._get_async_client_cache_key()}-pubsub") + clients.ttl_dict[cache_key] = 0 + + replacement = cache.init_pubsub_client() + await asyncio.wait_for(closed.wait(), timeout=1) + + assert replacement is not client + connection.disconnect.assert_awaited_once() + + with patch("litellm.in_memory_llm_clients_cache", clients): + asyncio.run(exercise()) 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/SlackAlerting/test_slack_alerting_utils.py b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py index 403cd51701d..b02dbea64b8 100644 --- a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py +++ b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py @@ -1,15 +1,9 @@ -import json -from typing import Optional -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest # Adds the grandparent directory to sys.path to allow importing project modules - import litellm -from litellm.integrations.langfuse.langfuse_prompt_management import ( - LangfusePromptManagement, -) from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager @@ -34,3 +28,95 @@ async def test_langfuse_not_initialized_returns_none_early(): # Verify the litellm_logging_obj was never accessed (early return) request_data["litellm_logging_obj"].assert_not_called() + + +@pytest.mark.asyncio +async def test_langfuse_trace_url_uses_the_request_host_without_building_a_logger(monkeypatch): + """Key-scoped callbacks point at their own Langfuse host; the alert link follows it. + + The lookup must not construct a LangFuseLogger per alert, or an alert storm + exhausts the initialized-client ceiling and takes the callback down with it. + """ + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) + logging_obj = MagicMock() + logging_obj._get_trace_id.return_value = "abc123" + logging_obj.standard_callback_dynamic_params = {"langfuse_host": "http://127.0.0.1:1"} + + result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + + assert result == "http://127.0.0.1:1/trace/abc123" + assert litellm.initialized_langfuse_clients == 0 + + +@pytest.mark.asyncio +async def test_langfuse_trace_url_falls_back_to_the_env_host(monkeypatch): + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setenv("LANGFUSE_HOST", "langfuse.internal:3000") + logging_obj = MagicMock() + logging_obj._get_trace_id.return_value = "abc123" + logging_obj.standard_callback_dynamic_params = {} + + assert await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) == ( + "http://langfuse.internal:3000/trace/abc123" + ) + + +@pytest.mark.asyncio +async def test_langfuse_trace_url_when_callback_registered_as_logger_instance(monkeypatch): + from litellm.integrations.langfuse.langfuse import LangFuseLogger + + logger = LangFuseLogger( + langfuse_public_key="pk-slack-instance", + langfuse_secret="sk-slack-instance", + langfuse_host="http://127.0.0.1:1", + ) + monkeypatch.setattr(litellm, "success_callback", [logger]) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setenv("LANGFUSE_HOST", "http://env-host.invalid") + logging_obj = MagicMock() + logging_obj._get_trace_id.return_value = "trace-from-instance" + logging_obj.standard_callback_dynamic_params = {} + + result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + + assert result == "http://127.0.0.1:1/trace/trace-from-instance" + + +@pytest.mark.asyncio +async def test_langfuse_trace_url_when_prompt_management_is_the_registered_callback(monkeypatch): + """Prompt management registers a LangFuseLogger subclass; the alert must read its host, not crash.""" + from litellm.integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement + + prompt_callback = LangfusePromptManagement( + langfuse_public_key="pk-slack-prompt", + langfuse_secret="sk-slack-prompt", + langfuse_host="http://127.0.0.1:2", + ) + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "callbacks", [prompt_callback]) + monkeypatch.setenv("LANGFUSE_HOST", "http://env-host.invalid") + logging_obj = MagicMock() + logging_obj._get_trace_id.return_value = "trace-from-prompt-callback" + logging_obj.standard_callback_dynamic_params = {} + + result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + + assert result == "http://127.0.0.1:2/trace/trace-from-prompt-callback" + + +@pytest.mark.asyncio +async def test_langfuse_trace_url_absent_when_trace_id_never_arrives(monkeypatch): + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setattr("litellm.integrations.SlackAlerting.utils.asyncio.sleep", AsyncMock()) + logging_obj = MagicMock() + logging_obj._get_trace_id.return_value = None + logging_obj.standard_callback_dynamic_params = {"langfuse_host": "http://127.0.0.1:1"} + + assert await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) is None diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py b/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py index 7dea4e67cdd..a7a553b2d9f 100644 --- a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py +++ b/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py @@ -1,9 +1,13 @@ -from types import MappingProxyType +import sys +from datetime import datetime, timezone from typing import Final from unittest.mock import MagicMock, patch import pytest +# langfuse_client_init imports this lazily; cache it before any test mocks +# sys.modules["langfuse"], or a single-file run dies on the real import +import litellm.integrations.langfuse.langfuse_sdk # noqa: F401 from litellm.integrations.langfuse.langfuse_prompt_management import ( LangfusePromptManagement, langfuse_client_init, @@ -17,9 +21,7 @@ class TestLangfusePromptManagement: # This also prevents test-ordering issues when earlier tests remove sys.modules["langfuse"]. self._mock_langfuse = MagicMock() self._mock_langfuse.version.__version__ = "3.0.0" - self._langfuse_patcher = patch.dict( - "sys.modules", {"langfuse": self._mock_langfuse} - ) + self._langfuse_patcher = patch.dict("sys.modules", {"langfuse": self._mock_langfuse}) self._langfuse_patcher.start() def teardown_method(self): @@ -31,9 +33,7 @@ class TestLangfusePromptManagement: patch.object( langfuse_prompt_management, "should_run_prompt_management" ) as mock_should_run_prompt_management, - patch.object( - langfuse_prompt_management, "_get_prompt_from_id" - ) as mock_get_prompt_from_id, + patch.object(langfuse_prompt_management, "_get_prompt_from_id") as mock_get_prompt_from_id, ): mock_should_run_prompt_management.return_value = True langfuse_prompt_management.get_chat_completion_prompt( @@ -51,9 +51,7 @@ class TestLangfusePromptManagement: def test_log_failure_event_runs_async_logger(self): langfuse_prompt_management = LangfusePromptManagement() - with patch( - "litellm.integrations.langfuse.langfuse_prompt_management.run_async_function" - ) as mock_run_async: + with patch("litellm.integrations.langfuse.langfuse_prompt_management.run_async_function") as mock_run_async: kwargs = {"standard_callback_dynamic_params": {}} start_time, end_time = 1, 2 @@ -65,10 +63,7 @@ class TestLangfusePromptManagement: ) mock_run_async.assert_called_once() - assert ( - mock_run_async.call_args[0][0] - == langfuse_prompt_management.async_log_failure_event - ) + assert mock_run_async.call_args[0][0] == langfuse_prompt_management.async_log_failure_event def test_langfuse_client_init_passes_dedicated_httpx_client(self): import httpx @@ -76,35 +71,28 @@ class TestLangfusePromptManagement: from litellm.llms.custom_httpx.http_handler import _get_httpx_client shared_client = _get_httpx_client().client - - mock_langfuse_class = MagicMock() + built = MagicMock() with ( patch( "litellm.integrations.langfuse.langfuse_prompt_management.resolve_langfuse_credentials", return_value=("pk-1234", "sk-1234", "https://localhost"), ), patch( - "litellm.integrations.langfuse.langfuse_prompt_management.LangFuseLogger._get_langfuse_flush_interval", - return_value=1, - ), - patch.dict("sys.modules", {"langfuse": self._mock_langfuse}), + "litellm.integrations.langfuse.langfuse_sdk.build_langfuse_client", built + ), # test-quality-ok: the REST client is built where langfuse_client_init resolves it; the transport it gets is the behavior under test patch( "litellm.llms.custom_httpx.http_handler.get_ssl_configuration", return_value=False, ) as mock_get_ssl, ): - self._mock_langfuse.Langfuse = mock_langfuse_class - langfuse_client_init( langfuse_public_key="pk-1234", langfuse_secret="sk-1234", langfuse_host="https://localhost", ) - mock_langfuse_class.assert_called_once() - call_kwargs = mock_langfuse_class.call_args[1] - assert "httpx_client" in call_kwargs - passed_client = call_kwargs["httpx_client"] + built.assert_called_once() + passed_client = built.call_args.kwargs["httpx_client"] assert isinstance(passed_client, httpx.Client) assert passed_client is not shared_client mock_get_ssl.assert_called_once() @@ -112,28 +100,181 @@ class TestLangfusePromptManagement: langfuse_client_init.cache_clear() -class _RecordingLangfuseForEnv: - last_environment: str | None = None - - def __init__(self, *, environment: str | None = None, **parameters: object) -> None: # kwargs-ok: records only environment out of whatever langfuse_client_init forwards - type(self).last_environment = environment - - @pytest.mark.parametrize( ("env_value", "expected"), (("Production", "default"), ("production ", "production"), ("prod", "prod")), ) -def test_langfuse_client_init_resolves_deployment_environment(monkeypatch, env_value, expected): - mock_langfuse_module: Final = MagicMock() - mock_langfuse_module.version.__version__ = "2.60.0" - mock_langfuse_module.Langfuse = _RecordingLangfuseForEnv +def test_prompt_management_logger_exports_the_resolved_deployment_environment(monkeypatch, env_value, expected): + from langfuse import LangfuseOtelSpanAttributes + + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-test") + monkeypatch.setenv("LANGFUSE_HOST", "http://127.0.0.1:1") + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_TRACING_ENVIRONMENT", env_value) + langfuse_client_init.cache_clear() + logger = LangfusePromptManagement() + langfuse_client_init.cache_clear() + assert logger.tracing.provider.resource.attributes[LangfuseOtelSpanAttributes.ENVIRONMENT] == expected + + +def test_langfuse_client_init_warns_that_upstream_langfuse_is_ignored(monkeypatch, caplog): + """The YAML `callbacks: ["langfuse"]` path builds its client here, not through LangFuseLogger.__init__, + so an operator who still sets UPSTREAM_LANGFUSE_* must get the same startup warning on this path.""" monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-test") monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-test") monkeypatch.setenv("LANGFUSE_HOST", "https://test.langfuse.com") - monkeypatch.setenv("LANGFUSE_TRACING_ENVIRONMENT", env_value) - monkeypatch.setattr(_RecordingLangfuseForEnv, "last_environment", None) - with patch.dict("sys.modules", MappingProxyType({"langfuse": mock_langfuse_module})): + monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "sk-upstream") + monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") + with caplog.at_level("WARNING", logger="LiteLLM"): langfuse_client_init.cache_clear() langfuse_client_init() langfuse_client_init.cache_clear() - assert _RecordingLangfuseForEnv.last_environment == expected + assert any("UPSTREAM_LANGFUSE_* is no longer supported" in record.getMessage() for record in caplog.records) + + +def test_langfuse_client_init_mock_mode_makes_no_network_calls(monkeypatch): + """LANGFUSE_MOCK promises full execution without egress. + + The registry maps the "langfuse" callback to LangfusePromptManagement, so + this logger is the one the standard proxy path emits observations through; + they travel over litellm's own OTLP exporter, which the httpx mock cannot see. + """ + import threading + from http.server import BaseHTTPRequestHandler, HTTPServer + + import litellm + + received = [] + + class _Receiver(BaseHTTPRequestHandler): + def do_POST(self): + received.append(self.path) + self.rfile.read(int(self.headers.get("Content-Length") or 0)) + self.send_response(200) + self.send_header("Content-Length", "0") + self.end_headers() + + def log_message(self, *args): + pass + + server = HTTPServer(("127.0.0.1", 0), _Receiver) + threading.Thread(target=server.serve_forever, daemon=True).start() + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_HOST", f"http://127.0.0.1:{server.server_port}") + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-pm-mock-egress") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-pm-mock-egress") + langfuse_client_init.cache_clear() + now: Final = datetime.now(timezone.utc) + + try: + logger = LangfusePromptManagement() + logged = logger.log_event_on_langfuse( + kwargs={ + "litellm_call_id": "call-pm-mock-egress", + "call_type": "completion", + "litellm_params": {"metadata": {"trace_id": "a" * 32}}, + "messages": [{"role": "user", "content": "hi"}], + "optional_params": {}, + }, + response_obj=litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "ok"}}]), + start_time=now, + end_time=now, + ) + logger.flush() + finally: + server.shutdown() + langfuse_client_init.cache_clear() + + assert logged["trace_id"] == "a" * 32 + assert received == [], f"LANGFUSE_MOCK still sent spans to the configured host: {received}" + + +def test_langfuse_debug_reaches_the_export_channel_through_the_registered_callback(monkeypatch): + """The registry maps ``langfuse`` to this class, whose constructor never runs ``LangFuseLogger.__init__``, + so wiring ``LANGFUSE_DEBUG`` only there left the flag a no-op on the YAML callback path.""" + import logging + + from litellm.integrations.langfuse.langfuse_sdk import release_langfuse_tracing + + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_HOST", "http://127.0.0.1:1") + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-pm-debug-wire") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-pm-debug-wire") + monkeypatch.setenv("LANGFUSE_DEBUG", "true") + langfuse_client_init.cache_clear() + langfuse_logger: Final = logging.getLogger("langfuse") + level_before: Final = langfuse_logger.level + langfuse_logger.setLevel(logging.WARNING) + try: + logger = LangfusePromptManagement() + assert langfuse_logger.level == logging.DEBUG + release_langfuse_tracing(logger.tracing, grace_seconds=0.0) + finally: + langfuse_logger.setLevel(level_before) + langfuse_client_init.cache_clear() + + +@pytest.mark.asyncio +async def test_async_log_failure_event_records_trace_id_for_alerting(monkeypatch): + from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id + from litellm.litellm_core_utils.specialty_caches.service_trace_id_cache import in_memory_trace_id_cache + + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_HOST", "http://127.0.0.1:1") + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-pm-trace-cache") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-pm-trace-cache") + langfuse_client_init.cache_clear() + call_id: Final = "call-trace-cache-1" + now: Final = datetime.now(timezone.utc) + kwargs: Final = { + "litellm_call_id": call_id, + "model": "gpt-5.4", + "messages": [{"role": "user", "content": "hi"}], + "litellm_params": {"metadata": {"trace_id": "alert-trace-1"}}, + "optional_params": {}, + "standard_callback_dynamic_params": {}, + "exception": RuntimeError("provider down"), + } + + try: + await LangfusePromptManagement().async_log_failure_event( + kwargs=kwargs, response_obj=None, start_time=now, end_time=now + ) + finally: + langfuse_client_init.cache_clear() + + assert in_memory_trace_id_cache.get_cache(litellm_call_id=call_id, service_name="langfuse") == resolve_trace_id( + "alert-trace-1" + ) + + +def test_old_sdk_fails_with_the_upgrade_message_before_the_otel_module_is_imported(monkeypatch): + """On a v2 install `langfuse_sdk` itself fails to import, so the version gate must run first.""" + import litellm.integrations.langfuse.langfuse_prompt_management as pm_module + + monkeypatch.setattr(pm_module, "installed_langfuse_version", lambda: "2.59.7") + monkeypatch.setitem(sys.modules, "litellm.integrations.langfuse.langfuse_sdk", None) + + with pytest.raises(ImportError) as raised: + LangfusePromptManagement( + langfuse_public_key="pk-old", langfuse_secret="sk-old", langfuse_host="http://127.0.0.1:1" + ) + + assert "2.59.7" in str(raised.value) + assert "langfuse_otel" in str(raised.value) + + +@pytest.mark.parametrize("raw", ["abc", "2.5"], ids=["text", "fraction"]) +def test_prompt_cache_ttl_typo_is_named_instead_of_reported_as_not_installed(monkeypatch, raw): + """The v4 SDK runs ``int()`` on this variable at import, and ``langfuse_client_init`` wraps any import + failure as "Langfuse not installed", so the gate has to run before that import.""" + monkeypatch.setenv("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS", raw) + monkeypatch.setitem(sys.modules, "litellm.integrations.langfuse.langfuse_sdk", None) + langfuse_client_init.cache_clear() + + with pytest.raises(ValueError, match="LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS") as raised: + langfuse_client_init(langfuse_public_key="pk-ttl", langfuse_secret="sk-ttl", langfuse_host="http://127.0.0.1:1") + + assert "not installed" not in str(raised.value) + assert repr(raw) in str(raised.value) diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py b/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py new file mode 100644 index 00000000000..c15a12c07cb --- /dev/null +++ b/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py @@ -0,0 +1,1759 @@ +"""Covers litellm's own Langfuse export channel: plain OTel spans carrying the v2 contracts. + +The timestamp assertions are the regression guard for the migration: the v4 SDK's public +API has no observation start time, so a callback running after the model call would +otherwise record its own duration instead of the call's. +""" + +import json +import logging +import threading +import uuid +from base64 import b64encode +from datetime import datetime, timedelta, timezone +from time import monotonic, sleep +from types import MappingProxyType +from typing import Final + +import httpx +import opentelemetry.trace as otel_trace +import pytest +from langfuse import LangfuseOtelSpanAttributes as A +from langfuse.api.core.api_error import ApiError +from langfuse.api.core.request_options import RequestOptions +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.sdk.trace import SpanProcessor, TracerProvider +from opentelemetry.sdk.trace.export import SpanExporter, SpanExportResult +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from litellm.integrations.langfuse.langfuse import ( + MINIMUM_LANGFUSE_VERSION, + installed_langfuse_version, + raise_if_unsupported_langfuse_version, +) +from litellm.integrations.langfuse.langfuse_sdk import ( + DiscardingSpanExporter, + LangfuseApiClient, + LangfusePromptError, + LangfuseSpanExporter, + LangfuseTracing, + _build_span_exporter, + _encode, + acquire_langfuse_tracing, + build_langfuse_client, + build_langfuse_tracing, + configured_flush_at, + configured_prompt_cache_ttl, + configured_sample_rate, + enable_langfuse_debug_logging, + flush_langfuse_tracing, + observation_attributes, + release_langfuse_tracing, + resolve_observation_id, + resolve_trace_id, + start_child_span, + start_generation, + to_unix_nanos, + trace_attributes, +) +from litellm.llms.custom_httpx.http_handler import HTTPHandler + +CALL_START = datetime(2024, 3, 1, 12, 0, 0, tzinfo=timezone.utc) +FIRST_TOKEN = CALL_START + timedelta(seconds=5) +CALL_END = CALL_START + timedelta(seconds=20) +TRACE_A = "a" * 32 +PARENT_C = "c" * 16 + + +@pytest.fixture(autouse=True) +def _own_channel_registry(monkeypatch: pytest.MonkeyPatch) -> None: + """Channels leaked by other test modules would otherwise take part in every process-wide flush here.""" + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk._TRACING", {}) + + +@pytest.fixture(name="channel") +def _channel() -> tuple[LangfuseTracing, InMemorySpanExporter]: + exporter = InMemorySpanExporter() + return ( + build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ), + exporter, + ) + + +def _only_span(exporter, name): + return next(s for s in exporter.get_finished_spans() if s.name == name) + + +def _generation( + tracing, + *, + name="gen", + trace_id=TRACE_A, + parent=None, + existing=False, + observation_id=None, + public=None, + attributes=None, +): + return start_generation( + tracing=tracing, + trace_id=trace_id, + parent_observation_id=parent, + existing_trace=existing, + observation_id=observation_id, + name=name, + start_time=CALL_START, + public=public, + attributes=attributes if attributes is not None else {}, + ) + + +def test_generation_records_the_model_call_window_not_the_callback(channel): + tracing, exporter = channel + attributes = observation_attributes(observation_type="generation", completion_start_time=FIRST_TOKEN) + _generation(tracing, attributes=attributes).end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "gen") + assert span.start_time == to_unix_nanos(CALL_START) + assert span.end_time == to_unix_nanos(CALL_END) + assert (span.end_time - span.start_time) == 20 * 1_000_000_000 + assert datetime.fromisoformat(json.loads(span.attributes[A.OBSERVATION_COMPLETION_START_TIME])) == FIRST_TOKEN + + +@pytest.mark.parametrize( + "supplied", + [1709294400.5, datetime(2024, 3, 1, 12, 0, 0, 500000, tzinfo=timezone.utc)], + ids=["unix-seconds-float", "datetime"], +) +def test_timestamps_accept_both_shapes_guardrails_and_callbacks_use(supplied): + """Guardrail entries carry unix seconds as floats, the callback carries datetimes.""" + assert to_unix_nanos(supplied) == 1709294400500000000 + + +def test_guardrail_span_with_float_timestamps_keeps_its_own_window_under_the_generation(channel): + tracing, exporter = channel + guardrail_start = 1709294400.0 + generation = _generation(tracing) + start_child_span( + tracing=tracing, parent=generation, name="guardrail", start_time=guardrail_start, attributes={} + ).end(guardrail_start + 2) + generation.end(CALL_END) + tracing.flush() + + guardrail = _only_span(exporter, "guardrail") + exported_generation = _only_span(exporter, "gen") + assert (guardrail.end_time - guardrail.start_time) == 2 * 1_000_000_000 + assert guardrail.context.trace_id == exported_generation.context.trace_id + assert guardrail.parent.span_id == exported_generation.context.span_id + + +def test_requested_trace_id_is_the_exported_trace_id_and_the_generation_is_its_root(channel): + """v2 ``trace(id=...)``: the caller's id is the trace and the generation has no parent.""" + tracing, exporter = channel + generation = _generation(tracing) + generation.end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "gen") + assert generation.trace_id == TRACE_A + assert format(span.context.trace_id, "032x") == TRACE_A + assert span.parent is None + + +def test_parent_observation_id_nests_the_generation_under_the_callers_observation(channel): + tracing, exporter = channel + _generation(tracing, name="child-gen", parent=PARENT_C).end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "child-gen") + assert format(span.context.trace_id, "032x") == TRACE_A + assert format(span.parent.span_id, "016x") == PARENT_C + assert span.parent.is_remote + + +def test_existing_trace_is_appended_to_rather_than_rewritten(channel): + """v2 ``existing_trace_id``: the generation joins the trace without becoming its root.""" + tracing, exporter = channel + _generation(tracing, name="continued", existing=True).end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "continued") + assert format(span.context.trace_id, "032x") == TRACE_A + assert span.parent is not None + assert span.parent.span_id != 0 + + +def test_generation_does_not_hang_under_the_callers_active_span(channel): + """The caller's own OTel span must stay untouched and must not become the generation's parent.""" + tracing, exporter = channel + app_tracer = TracerProvider().get_tracer("app") + with app_tracer.start_as_current_span("app-span") as app_span: + _generation(tracing, trace_id="b" * 32).end(CALL_END) + attributes_after = dict(app_span.attributes or {}) + tracing.flush() + + span = _only_span(exporter, "gen") + assert span.parent is None + assert format(span.context.trace_id, "032x") == "b" * 32 + assert attributes_after == {} + + +def test_requested_observation_id_becomes_the_exported_span_id(channel): + """v2 ``generation(id=...)``: the caller's id is what the export carries and what ``.id`` returns.""" + tracing, exporter = channel + requested = resolve_observation_id("chatcmpl-123") + + generation = _generation(tracing, observation_id=requested) + generation.end(CALL_END) + tracing.flush() + + assert generation.id == requested + assert format(_only_span(exporter, "gen").context.span_id, "016x") == requested + + +def test_requested_ids_do_not_leak_into_the_next_span(channel): + tracing, exporter = channel + requested = resolve_observation_id("chatcmpl-123") + + _generation(tracing, name="first", observation_id=requested).end(CALL_END) + second = _generation(tracing, name="second", trace_id=resolve_trace_id(None)) + second.end(CALL_END) + child = start_child_span(tracing=tracing, parent=second, name="child", start_time=CALL_END, attributes={}) + child.end(CALL_END) + tracing.flush() + + assert second.id != requested + assert child.id not in (requested, second.id) + assert len({span.context.span_id for span in exporter.get_finished_spans()}) == 3 + + +@pytest.mark.parametrize("public", [True, False], ids=["public", "private"]) +def test_trace_public_flag_lands_on_the_root_observation(channel, public): + tracing, exporter = channel + _generation(tracing, public=public, attributes=trace_attributes(public=public)).end(CALL_END) + tracing.flush() + assert _only_span(exporter, "gen").attributes[A.TRACE_PUBLIC] is public + + +def test_trace_public_flag_is_absent_when_not_requested(channel): + tracing, exporter = channel + _generation(tracing, attributes=trace_attributes(public=None)).end(CALL_END) + tracing.flush() + assert A.TRACE_PUBLIC not in _only_span(exporter, "gen").attributes + + +@pytest.mark.parametrize("public", [True, False, None], ids=["public", "private", "unset"]) +def test_child_span_repeats_the_generation_public_flag(channel, public): + """The server folds ``public`` across observations and reads a missing value as False. + + A guardrail span without the flag turned a ``trace_public: true`` request private on Langfuse Cloud. + """ + tracing, exporter = channel + generation = _generation(tracing, public=public) + start_child_span(tracing=tracing, parent=generation, name="guardrail", start_time=CALL_END, attributes={}).end() + generation.end(CALL_END) + tracing.flush() + + assert _only_span(exporter, "guardrail").attributes.get(A.TRACE_PUBLIC) is public + + +def test_trace_attributes_carry_the_v2_trace_fields(): + attributes = trace_attributes( + name="trace-name", + user_id="user-1", + session_id="session-1", + version="v2", + release="rel-1", + tags=("a", "b"), + metadata={"tenant": "t1", "nested": {"k": 1}}, + input={"messages": []}, + output="answer", + ) + assert attributes[A.TRACE_NAME] == "trace-name" + assert attributes[A.TRACE_USER_ID] == "user-1" + assert attributes[A.TRACE_SESSION_ID] == "session-1" + assert attributes[A.VERSION] == "v2" + assert attributes[A.RELEASE] == "rel-1" + assert attributes[A.TRACE_TAGS] == ("a", "b") + assert attributes[f"{A.TRACE_METADATA}.tenant"] == "t1" + assert json.loads(attributes[f"{A.TRACE_METADATA}.nested"]) == {"k": 1} + assert json.loads(attributes[A.TRACE_INPUT]) == {"messages": []} + assert attributes[A.TRACE_OUTPUT] == "answer" + + +def test_trace_attributes_skip_what_the_request_did_not_supply(): + assert dict(trace_attributes()) == {} + + +def test_non_mapping_metadata_is_carried_whole_instead_of_raising(): + """A truthy non-dict ``trace_metadata`` used to blow up the callback on ``**`` unpacking.""" + attributes = trace_attributes(metadata=("not", "a", "dict")) + assert json.loads(attributes[A.TRACE_METADATA]) == ["not", "a", "dict"] + + +def test_observation_attributes_serialize_the_generation_fields(): + attributes = observation_attributes( + observation_type="generation", + input=[{"role": "user", "content": "hi"}], + output={"role": "assistant", "content": "hello"}, + metadata=MappingProxyType({"litellm_call_id": "call-1", "cache_hit": False}), + level="ERROR", + status_message="boom", + model="gpt-4o", + model_parameters={"temperature": 0.1}, + usage_details={"input": 1, "output": 2}, + cost_details={"total": 0.01}, + prompt="not-a-prompt-client", + ) + assert attributes[A.OBSERVATION_TYPE] == "generation" + assert attributes[A.OBSERVATION_LEVEL] == "ERROR" + assert attributes[A.OBSERVATION_STATUS_MESSAGE] == "boom" + assert attributes[A.OBSERVATION_MODEL] == "gpt-4o" + assert json.loads(attributes[A.OBSERVATION_INPUT]) == [{"role": "user", "content": "hi"}] + assert json.loads(attributes[A.OBSERVATION_OUTPUT]) == {"role": "assistant", "content": "hello"} + assert json.loads(attributes[A.OBSERVATION_MODEL_PARAMETERS]) == {"temperature": 0.1} + assert json.loads(attributes[A.OBSERVATION_USAGE_DETAILS]) == {"input": 1, "output": 2} + assert json.loads(attributes[A.OBSERVATION_COST_DETAILS]) == {"total": 0.01} + assert attributes[f"{A.OBSERVATION_METADATA}.litellm_call_id"] == "call-1" + assert attributes[f"{A.OBSERVATION_METADATA}.cache_hit"] is False + assert A.OBSERVATION_PROMPT_NAME not in attributes + + +@pytest.mark.parametrize( + "supplied, expected", + [ + ("0123456789abcdef0123456789abcdef", "0123456789abcdef0123456789abcdef"), + ("0123456789ABCDEF0123456789ABCDEF", "0123456789abcdef0123456789abcdef"), + ("3fe0c940-b69a-de3b-a77c-06102505349a", "3fe0c940b69ade3ba77c06102505349a"), + ], + ids=["already-hex", "uppercase-hex", "uuid-with-dashes"], +) +def test_trace_id_passes_through_when_it_is_already_usable(supplied, expected): + assert resolve_trace_id(supplied) == expected + + +def test_arbitrary_trace_id_is_hashed_deterministically(): + first = resolve_trace_id("order-4471") + assert first == resolve_trace_id("order-4471") + assert len(first) == 32 and first == first.lower() + assert first != resolve_trace_id("order-4472") + + +@pytest.mark.parametrize("supplied", [12345, 12.5, True, False], ids=["int", "float", "true", "false"]) +def test_non_string_trace_id_is_normalized(supplied): + resolved = resolve_trace_id(supplied) + + assert len(resolved) == 32 + assert resolved == resolve_trace_id(supplied) + + +@pytest.mark.parametrize("supplied", [12345, 12.5, True, False], ids=["int", "float", "true", "false"]) +def test_non_string_observation_id_is_normalized(supplied): + resolved = resolve_observation_id(supplied) + + assert len(resolved) == 16 + assert resolved == resolve_observation_id(supplied) + + +def test_all_zero_ids_are_hashed_instead_of_passed_through(): + zero_trace = "0" * 32 + zero_span = "0" * 16 + + assert resolve_trace_id(zero_trace) != zero_trace + assert resolve_trace_id(zero_trace) == resolve_trace_id(zero_trace) + assert int(resolve_trace_id(zero_trace), 16) != 0 + assert resolve_observation_id(zero_span) != zero_span + assert int(resolve_observation_id(zero_span), 16) != 0 + + +def test_hyphen_only_trace_ids_are_deterministic(): + assert resolve_trace_id("---") == resolve_trace_id("---") + + +def test_trace_id_with_trailing_newline_is_hashed(): + supplied = "a" * 32 + "\n" + + resolved = resolve_trace_id(supplied) + + assert resolved != supplied + assert len(resolved) == 32 + + +def test_missing_trace_id_still_yields_a_valid_trace_id(): + generated = resolve_trace_id(None) + assert len(generated) == 32 + assert int(generated, 16) >= 0 + + +@pytest.mark.parametrize( + "supplied, expected", + [ + ("0123456789abcdef", "0123456789abcdef"), + (None, None), + ("", None), + ], + ids=["already-hex", "none", "empty"], +) +def test_observation_id_normalisation(supplied, expected): + assert resolve_observation_id(supplied) == expected + + +def test_arbitrary_observation_id_is_hashed_to_a_span_id(): + resolved = resolve_observation_id("my-parent-observation") + assert len(resolved) == 16 + assert resolved == resolve_observation_id("my-parent-observation") + + +@pytest.mark.parametrize("unsupported", ["2.59.7", "3.15.0", "5.0.0"], ids=["v2", "v3", "v5"]) +def test_unsupported_sdk_fails_loudly_rather_than_dropping_every_event(unsupported): + with pytest.raises(ImportError) as raised: + raise_if_unsupported_langfuse_version(unsupported) + assert unsupported in str(raised.value) + assert MINIMUM_LANGFUSE_VERSION in str(raised.value) + + +def test_supported_sdk_is_accepted(): + assert raise_if_unsupported_langfuse_version(installed_langfuse_version()) is None + + +def test_channel_carries_environment_and_release_on_the_resource(): + tracing = build_langfuse_tracing( + exporter=DiscardingSpanExporter(), + environment="staging", + release="v9", + sample_rate=1.0, + flush_interval_millis=10, + ) + attributes = tracing.provider.resource.attributes + assert attributes[A.ENVIRONMENT] == "staging" + assert attributes[A.RELEASE] == "v9" + + +def _generations_exported_at(sample_rate: float, trace_ids: tuple[str, ...]) -> frozenset[str]: + exporter: Final = InMemorySpanExporter() + tracing: Final = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=sample_rate, flush_interval_millis=10 + ) + for trace_id in trace_ids: + _generation(tracing, name="sampled", trace_id=trace_id).end(CALL_END) + tracing.flush() + return frozenset(format(span.context.trace_id, "032x") for span in exporter.get_finished_spans()) + + +def test_sample_rate_zero_drops_and_one_keeps_every_trace(): + trace_ids: Final = tuple(resolve_trace_id(uuid.uuid4()) for _ in range(20)) + assert _generations_exported_at(0, trace_ids) == frozenset() + assert _generations_exported_at(1, trace_ids) == frozenset(trace_ids) + + +def test_fractional_sample_rate_keeps_a_deterministic_share_of_uuid_trace_ids(): + trace_ids: Final = tuple(resolve_trace_id(uuid.uuid4()) for _ in range(400)) + kept: Final = _generations_exported_at(0.5, trace_ids) + assert 140 <= len(kept) <= 260 + assert _generations_exported_at(0.5, trace_ids) == kept + assert kept < _generations_exported_at(0.9, trace_ids) + + +@pytest.mark.parametrize("raw", ["1.5", "-0.5", "abc"]) +def test_unusable_sample_rate_warns_and_exports_everything( + monkeypatch: pytest.MonkeyPatch, raw: str, caplog: pytest.LogCaptureFixture +): + monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", raw) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert configured_sample_rate() == 1.0 + assert "LANGFUSE_SAMPLE_RATE" in caplog.text + + +def test_configured_sample_rate_reads_the_env_var(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LANGFUSE_SAMPLE_RATE", raising=False) + assert configured_sample_rate() == 1.0 + monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "0.25") + assert configured_sample_rate() == 0.25 + + +def _exported_generations(exporter: InMemorySpanExporter, tracing: LangfuseTracing, count: int) -> tuple: + for _ in range(count): + _generation(tracing, trace_id=resolve_trace_id(uuid.uuid4())).end(CALL_END) + tracing.flush() + return exporter.get_finished_spans() + + +def test_full_sample_rate_exports_every_trace_even_when_the_host_turned_otel_sampling_off( + monkeypatch: pytest.MonkeyPatch, +): + """A provider built without a sampler reads ``OTEL_TRACES_SAMPLER``, which belongs to the host's tracing.""" + monkeypatch.setenv("OTEL_TRACES_SAMPLER", "always_off") + exporter = InMemorySpanExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + assert len(_exported_generations(exporter, tracing, 5)) == 5 + + +@pytest.mark.parametrize( + ("variable", "value"), + [ + ("OTEL_SPAN_ATTRIBUTE_COUNT_LIMIT", "4"), + ("OTEL_ATTRIBUTE_COUNT_LIMIT", "4"), + ("OTEL_ATTRIBUTE_VALUE_LENGTH_LIMIT", "8"), + ("OTEL_SPAN_ATTRIBUTE_VALUE_LENGTH_LIMIT", "8"), + ], +) +def test_host_otel_span_limits_do_not_truncate_langfuse_observations( + monkeypatch: pytest.MonkeyPatch, variable: str, value: str +): + monkeypatch.setenv(variable, value) + exporter = InMemorySpanExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + attributes = {f"langfuse.observation.metadata.k{i}": "v" * 32 for i in range(40)} + _generation(tracing, attributes=attributes).end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "gen") + assert span.dropped_attributes == 0 + assert all(span.attributes[key] == "v" * 32 for key in attributes) + + +def test_host_otel_resource_env_does_not_reach_the_langfuse_resource(monkeypatch: pytest.MonkeyPatch): + """``OTEL_RESOURCE_ATTRIBUTES`` and ``OTEL_SERVICE_NAME`` belong to the host's tracing; Langfuse files a + trace under any ``deployment.environment`` it finds on the resource, and v2 shipped no resource at all.""" + monkeypatch.setenv("OTEL_RESOURCE_ATTRIBUTES", "team.secret.note=internal-only,deployment.environment=hijack") + monkeypatch.setenv("OTEL_SERVICE_NAME", "the-hosts-own-service") + tracing = build_langfuse_tracing( + exporter=InMemorySpanExporter(), environment="prod", release="r1", sample_rate=1.0, flush_interval_millis=10 + ) + + assert dict(tracing.provider.resource.attributes) == {A.ENVIRONMENT: "prod", A.RELEASE: "r1"} + + +@pytest.mark.parametrize( + ("value", "encoded"), + [ + (2**53 - 1, ("int_value", 2**53 - 1)), + (-(2**53) + 1, ("int_value", -(2**53) + 1)), + (2**53, ("string_value", str(2**53))), + (2**63 - 1, ("string_value", str(2**63 - 1))), + (2**63, ("string_value", str(2**63))), + (10**20, ("string_value", str(10**20))), + (-(2**63) - 1, ("string_value", str(-(2**63) - 1))), + (True, ("bool_value", True)), + ], + ids=[ + "json-safe-max", + "json-safe-min", + "json-safe-plus-one", + "int64-max", + "int64-max-plus-one", + "huge", + "int64-min-minus-one", + "bool", + ], +) +def test_metadata_ints_past_the_json_safe_range_reach_the_wire_as_strings(value, encoded): + """OTLP carries int64 only and its encoder silently drops any attribute it cannot fit, while the export + still succeeds, and Langfuse's reader rounds ints past 2**53 (int64 max read back as 9223372036854776000 + on 2026-09-21, where the v2 leg showed the exact digits as a string), so both ranges go as strings.""" + from opentelemetry.exporter.otlp.proto.common.trace_encoder import encode_spans + + exporter = InMemorySpanExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + attributes = observation_attributes(observation_type="generation", metadata={"order_id": value, "sibling": "kept"}) + _generation(tracing, attributes=attributes).end(CALL_END) + tracing.flush() + + (encoded_span,) = encode_spans(exporter.get_finished_spans()).resource_spans[0].scope_spans[0].spans + wire = {kv.key: kv.value for kv in encoded_span.attributes} + order_id = wire[f"{A.OBSERVATION_METADATA}.order_id"] + carried = { + "int_value": order_id.int_value, + "string_value": order_id.string_value, + "bool_value": order_id.bool_value, + } + assert (order_id.WhichOneof("value"), carried[order_id.WhichOneof("value")]) == encoded + assert wire[f"{A.OBSERVATION_METADATA}.sibling"].string_value == "kept" + + +@pytest.mark.parametrize( + ("raw", "expected"), + [(None, 60.0), ("5", 5.0), ("0", 0.0), (" -1 ", 60.0), ("2.5", 60.0), ("abc", 60.0)], + ids=["unset", "whole", "zero", "negative", "fraction", "text"], +) +def test_prompt_cache_ttl_env_falls_back_instead_of_raising(monkeypatch: pytest.MonkeyPatch, raw, expected, caplog): + """The SDK reads this knob as whole seconds; a negative one passes its import but must not cache forever, + and anything else falls back rather than raising out of logger construction.""" + if raw is None: + monkeypatch.delenv("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS", raising=False) + else: + monkeypatch.setenv("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS", raw) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert configured_prompt_cache_ttl() == expected + assert ("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS" in caplog.text) is (expected == 60.0 and raw is not None) + + +def test_many_metadata_keys_never_evict_the_generation_input_and_output(): + """OTel's default 128-attribute cap drops the earliest attributes, and v2 never capped metadata.""" + exporter = InMemorySpanExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + attributes = { + A.OBSERVATION_INPUT: "the-prompt", + A.OBSERVATION_OUTPUT: "the-completion", + **{f"langfuse.observation.metadata.k{i}": str(i) for i in range(300)}, + } + _generation(tracing, attributes=attributes).end(CALL_END) + tracing.flush() + + span = _only_span(exporter, "gen") + assert span.dropped_attributes == 0 + assert span.attributes[A.OBSERVATION_INPUT] == "the-prompt" + assert span.attributes[A.OBSERVATION_OUTPUT] == "the-completion" + assert span.attributes["langfuse.observation.metadata.k299"] == "299" + + +def test_otel_sdk_disabled_still_wins_but_is_called_out(monkeypatch: pytest.MonkeyPatch, caplog): + monkeypatch.setenv("OTEL_SDK_DISABLED", "true") + exporter = InMemorySpanExporter() + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + assert "OTEL_SDK_DISABLED" in caplog.text + assert _exported_generations(exporter, tracing, 3) == () + + +def test_spans_carry_the_langfuse_sdk_scope_name(channel): + """Langfuse keys on the SDK's instrumentation scope (langfuse 4.15.2, ``langfuse/_client/constants.py``, + read 2026-09-17); any other scope is foreign OTel traffic whose raw attributes get echoed into metadata.""" + tracing, exporter = channel + _generation(tracing).end(CALL_END) + tracing.flush() + assert _only_span(exporter, "gen").instrumentation_scope.name == "langfuse-sdk" + + +class _GatedExporter(SpanExporter): + """Hold the export thread until released, so spans pile up in the processor queue.""" + + def __init__(self) -> None: + self.gate = threading.Event() + self.batches: list[int] = [] + + def export(self, spans) -> SpanExportResult: + self.gate.wait(timeout=30) + self.batches.append(len(spans)) + return SpanExportResult.SUCCESS + + def shutdown(self) -> None: + return None + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return True + + +def test_export_queue_holds_a_v2_sized_burst_while_the_destination_stalls(): + """v2 queued 100k events; OTel's default 2048 dropped most of a burst during a destination stall.""" + exporter = _GatedExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + for _ in range(6000): + _generation(tracing, trace_id=resolve_trace_id(uuid.uuid4())).end(CALL_END) + exporter.gate.set() + assert tracing.flush(timeout_millis=30_000) is True + assert sum(exporter.batches) == 6000 + + +@pytest.mark.parametrize( + ("raw", "expected"), + [(None, 512), ("64", 64), ("0", 512), ("-5", 512), ("abc", 512), ("100001", 512), ("100000", 100_000)], + ids=["unset", "valid", "zero", "negative", "text", "over-queue", "at-queue"], +) +def test_langfuse_flush_at_is_parsed_like_the_sdk_did(monkeypatch: pytest.MonkeyPatch, raw, expected, caplog): + if raw is None: + monkeypatch.delenv("LANGFUSE_FLUSH_AT", raising=False) + else: + monkeypatch.setenv("LANGFUSE_FLUSH_AT", raw) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert configured_flush_at() == expected + assert ("LANGFUSE_FLUSH_AT" in caplog.text) is (raw is not None and str(expected) != raw) + + +def test_langfuse_flush_at_sizes_the_export_batches(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LANGFUSE_FLUSH_AT", "64") + exporter = _GatedExporter() + exporter.gate.set() + tracing = build_langfuse_tracing( + exporter=exporter, + environment=None, + release=None, + sample_rate=1.0, + flush_interval_millis=60_000, + flush_at=configured_flush_at(), + ) + for _ in range(200): + _generation(tracing, trace_id=resolve_trace_id(uuid.uuid4())).end(CALL_END) + tracing.flush() + assert sum(exporter.batches) == 200 + assert max(exporter.batches) == 64 + + +def test_acquired_channel_reads_langfuse_flush_at(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LANGFUSE_FLUSH_AT", "7") + small = _acquire(public_key="pk-flush-at-test") + monkeypatch.setenv("LANGFUSE_FLUSH_AT", "9") + assert _acquire(public_key="pk-flush-at-test") is not small + + +def test_channel_does_not_take_over_the_process_tracer_provider(): + provider_before = otel_trace.get_tracer_provider() + + tracing = acquire_langfuse_tracing( + public_key="pk-global-test", + secret_key="sk", + base_url="http://127.0.0.1:1", + environment=None, + release=None, + flush_interval=1.0, + mock_mode=True, + ) + + assert otel_trace.get_tracer_provider() is provider_before + assert tracing.provider is not provider_before + + +def _acquire(**overrides): + parameters = { + "public_key": "pk-cache-test", + "secret_key": "sk-cache", + "base_url": "http://127.0.0.1:1", + "environment": None, + "release": None, + "flush_interval": 1.0, + "mock_mode": True, + } + return acquire_langfuse_tracing(**{**parameters, **overrides}) + + +def test_same_credentials_share_one_channel(): + assert _acquire() is _acquire() + + +@pytest.mark.parametrize( + "override", + [ + {"secret_key": "sk-rotated"}, + {"base_url": "http://127.0.0.1:2"}, + {"environment": "staging"}, + {"mock_mode": False}, + ], + ids=["secret", "host", "environment", "mock-to-live"], +) +def test_changed_credentials_or_settings_get_their_own_channel(override): + assert _acquire() is not _acquire(**override) + + +class _RecordsShutdown(InMemorySpanExporter): + def __init__(self) -> None: + super().__init__() + self.shutdowns = 0 + + def shutdown(self) -> None: + self.shutdowns += 1 + super().shutdown() + + +def _acquire_recorded(monkeypatch: pytest.MonkeyPatch, public_key: str) -> tuple[LangfuseTracing, _RecordsShutdown]: + exporter = _RecordsShutdown() + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk._build_span_exporter", lambda **_: exporter) + return _acquire(public_key=public_key, mock_mode=False, flush_interval=600.0), exporter + + +def test_channel_is_retired_only_after_its_last_holder_releases_it(monkeypatch: pytest.MonkeyPatch): + """Two loggers on one credential set share the channel: the first release must leave it + exporting for the second, and the last release must shut the batch thread down and drop + the registry entry so the next logger gets a fresh channel instead of a dead one.""" + first, exporter = _acquire_recorded(monkeypatch, "pk-lease-test") + second = _acquire(public_key="pk-lease-test", mock_mode=False, flush_interval=600.0) + assert second is first + + release_langfuse_tracing(first, grace_seconds=0.0) + second.tracer.start_span("generation").end() + assert exporter.shutdowns == 0 + assert flush_langfuse_tracing() is True + assert len(exporter.get_finished_spans()) == 1 + + release_langfuse_tracing(second, grace_seconds=0.0) + assert exporter.shutdowns == 1 + assert _acquire(public_key="pk-lease-test", mock_mode=False, flush_interval=600.0) is not first + + +def test_release_flushes_the_queued_spans_before_the_channel_goes_away(monkeypatch: pytest.MonkeyPatch): + tracing, exporter = _acquire_recorded(monkeypatch, "pk-lease-flush-test") + tracing.tracer.start_span("generation").end() + + release_langfuse_tracing(tracing, grace_seconds=0.0) + + assert len(exporter.get_finished_spans()) == 1 + + +def test_channel_reacquired_within_the_grace_is_kept(monkeypatch: pytest.MonkeyPatch): + """A logger rebuilt for the same credentials right after the old one expired, and a callback + that fetched the old logger just before expiry, both keep exporting through the same channel.""" + tracing, exporter = _acquire_recorded(monkeypatch, "pk-lease-grace-test") + + release_langfuse_tracing(tracing, grace_seconds=0.2) + assert _acquire(public_key="pk-lease-grace-test", mock_mode=False, flush_interval=600.0) is tracing + + threading.Event().wait(0.5) + tracing.tracer.start_span("generation").end() + assert exporter.shutdowns == 0 + assert flush_langfuse_tracing() is True + assert len(exporter.get_finished_spans()) == 1 + + +def test_retire_timer_of_an_earlier_release_cannot_kill_a_reacquired_channel(monkeypatch: pytest.MonkeyPatch): + """release, re-acquire, release: the first timer used to fire into a channel that a later holder still + counted on for its own grace period, shutting the batch thread down while spans were still queued.""" + tracing, exporter = _acquire_recorded(monkeypatch, "pk-lease-race-test") + + release_langfuse_tracing(tracing, grace_seconds=0.2) + assert _acquire(public_key="pk-lease-race-test", mock_mode=False, flush_interval=600.0) is tracing + release_langfuse_tracing(tracing, grace_seconds=600.0) + + threading.Event().wait(0.5) + assert exporter.shutdowns == 0 + assert _acquire(public_key="pk-lease-race-test", mock_mode=False, flush_interval=600.0) is tracing + tracing.tracer.start_span("generation").end() + assert tracing.flush() is True + assert len(exporter.get_finished_spans()) == 1 + + +class _RejectsEverything(SpanExporter): + def export(self, spans) -> SpanExportResult: + return SpanExportResult.FAILURE + + def shutdown(self) -> None: + return None + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return True + + +def test_flush_is_false_when_the_destination_rejected_a_batch(monkeypatch: pytest.MonkeyPatch): + """The shutdown hook logs "channels flushed" off this value; a drained queue whose batches all + failed at the destination is a loss, not a flush.""" + monkeypatch.setattr( + "litellm.integrations.langfuse.langfuse_sdk._build_span_exporter", lambda **_: _RejectsEverything() + ) + tracing = _acquire(public_key="pk-flush-truth-test", mock_mode=False, flush_interval=600.0) + tracing.tracer.start_span("generation").end() + + assert tracing.flush() is False + assert flush_langfuse_tracing() is True, "an empty queue after the loss has nothing left to fail" + + +def test_release_of_a_channel_the_registry_never_handed_out_is_a_no_op(): + exporter = InMemorySpanExporter() + tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + + release_langfuse_tracing(tracing, grace_seconds=0.0) + tracing.tracer.start_span("generation").end() + + assert tracing.flush() is True + assert len(exporter.get_finished_spans()) == 1 + + +def test_flush_langfuse_tracing_exports_the_queued_spans_of_every_channel(monkeypatch: pytest.MonkeyPatch): + """The proxy shutdown hook flushes through this, so a span finished just before a + graceful restart must reach the exporter without waiting for the batch interval.""" + exporters: Final[ + list[InMemorySpanExporter] + ] = [] # mutable-ok: collects the exporters the patched builder hands out + + def build_in_memory(*, public_key: str, secret_key: str, base_url: str) -> InMemorySpanExporter: + exporters.append(InMemorySpanExporter()) + return exporters[-1] + + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk._build_span_exporter", build_in_memory) + for public_key in ("pk-flush-test-a", "pk-flush-test-b"): + _acquire(public_key=public_key, mock_mode=False, flush_interval=600.0).tracer.start_span("generation").end() + + assert [len(exporter.get_finished_spans()) for exporter in exporters] == [0, 0] + assert flush_langfuse_tracing() is True + assert [len(exporter.get_finished_spans()) for exporter in exporters] == [1, 1] + + +def test_flush_langfuse_tracing_flushes_channels_concurrently_under_one_deadline(monkeypatch: pytest.MonkeyPatch): + """A channel stuck on an unreachable host must not spend the whole deadline before the + next channel gets its turn; the first exporter here only returns once the second exported.""" + second_exported = threading.Event() + + class WaitsForTheOther(SpanExporter): + def export(self, spans): + return SpanExportResult.SUCCESS if second_exported.wait(timeout=5.0) else SpanExportResult.FAILURE + + def shutdown(self) -> None: + return None + + class Unblocks(SpanExporter): + def export(self, spans): + second_exported.set() + return SpanExportResult.SUCCESS + + def shutdown(self) -> None: + return None + + exporters = iter((WaitsForTheOther(), Unblocks())) + + def build_next(*, public_key: str, secret_key: str, base_url: str) -> SpanExporter: + return next(exporters) + + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk._build_span_exporter", build_next) + for public_key in ("pk-concurrent-flush-a", "pk-concurrent-flush-b"): + _acquire(public_key=public_key, mock_mode=False, flush_interval=600.0).tracer.start_span("generation").end() + + assert flush_langfuse_tracing(timeout_millis=2_000) is True + assert second_exported.is_set() + + +def test_flush_langfuse_tracing_leaves_an_overrunning_channel_on_a_daemon_thread(): + """A channel whose flush outlives the deadline is reported as failed and must not be able to + hold up interpreter exit, so the thread still flushing it has to be a daemon.""" + release = threading.Event() + + class BlocksUntilReleased(SpanProcessor): + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return release.wait(timeout=10.0) + + _acquire(public_key="pk-overrunning-flush", mock_mode=True, flush_interval=600.0).provider.add_span_processor( + BlocksUntilReleased() + ) + try: + assert flush_langfuse_tracing(timeout_millis=200) is False + stuck = [thread for thread in threading.enumerate() if thread.name.startswith("langfuse-flush")] + assert stuck and all(thread.daemon for thread in stuck) + finally: + release.set() + + +def test_a_changed_sample_rate_rebuilds_the_channel(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "0.25") + quarter = _acquire(public_key="pk-resample-test") + monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "1") + full = _acquire(public_key="pk-resample-test") + + assert full is not quarter + assert quarter.provider.sampler.get_description() == "TraceIdHashSampler{0.25}" + assert "TraceIdHashSampler" not in full.provider.sampler.get_description() + + +def _recording_transport(requests: list[httpx.Request], status: int = 401) -> httpx.Client: + def record(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(status, json=_PROJECTS_BODY if status == 200 else {"message": "unauthorized"}) + + return httpx.Client(transport=httpx.MockTransport(record)) + + +_PROJECTS_BODY: Final = { + "data": [{"id": "proj-under-test", "name": "p", "metadata": {}, "organization": {"id": "o", "name": "o"}}] +} + + +def test_rest_client_authenticates_with_the_credentials_it_was_built_with(): + """Two loggers for one public key but different secrets or hosts each talk to their own project.""" + requests: list[httpx.Request] = [] + build_langfuse_client( + public_key="pk-rest-test", + secret_key="sk-first", + base_url="http://127.0.0.1:1", + httpx_client=_recording_transport(requests), + ) + rotated = build_langfuse_client( + public_key="pk-rest-test", + secret_key="sk-second", + base_url="http://127.0.0.1:2", + httpx_client=_recording_transport(requests), + ) + + assert rotated.auth_check() is not None + assert requests[-1].url.host == "127.0.0.1" and requests[-1].url.port == 2 + assert requests[-1].headers["authorization"] == "Basic " + b64encode(b"pk-rest-test:sk-second").decode() + + +def test_rest_client_without_keys_fails_auth_check_instead_of_raising(monkeypatch): + """``/health/services?service=langfuse`` with no credentials must report a failed check, not crash.""" + for name in ("LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY"): + monkeypatch.delenv(name, raising=False) + client = build_langfuse_client(public_key=None, secret_key=None, base_url="http://127.0.0.1:1", httpx_client=None) + assert client.auth_check() is not None + + +def test_auth_check_names_the_servers_rejection(caplog): + """``/health/services`` used to print the 401 verbatim; a generic credentials message hides a 403 or a 500.""" + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=_recording_transport([], status=401), + ) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + failure = client.auth_check() + assert failure is not None + assert failure.reason == "status_code: 401, body: {'message': 'unauthorized'}" + assert failure.reason in caplog.text + + +def test_auth_check_names_an_unreachable_destination_rather_than_the_keys(): + def refuse(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("connection refused by lf.internal.example", request=request) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://lf.internal.example", + httpx_client=httpx.Client(transport=httpx.MockTransport(refuse)), + ) + failure = client.auth_check() + assert failure is not None + assert "connection refused by lf.internal.example" in failure.reason + + +def test_auth_check_fails_when_the_keys_reach_no_project(): + """A 200 with an empty project list is what the SDK's own ``auth_check`` raises on; it is not a pass.""" + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(lambda _: httpx.Response(200, json={"data": []}))), + ) + failure = client.auth_check() + assert failure is not None + assert "no project" in failure.reason + + +@pytest.mark.parametrize("status", [500, 503, 429], ids=["http-500", "http-503", "http-429"]) +def test_auth_check_and_project_id_make_one_round_trip_when_langfuse_is_down(status): + """Both run on the event loop; the generated client's default retries sleep for seconds, or for Retry-After.""" + requests: list[httpx.Request] = [] + + def fail(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(status, request=request, headers={"retry-after": "20"}, json={"message": "down"}) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(fail)), + ) + + started = monotonic() + failure = client.auth_check() + with pytest.raises(ApiError): + client.project_id() + assert failure is not None and f"status_code: {status}" in failure.reason + assert len(requests) == 2 + assert monotonic() - started < 0.5 + + +@pytest.mark.parametrize( + ("status", "round_trips"), + [(500, 2), (503, 2), (429, 1), (404, 1)], + ids=["http-500", "http-503", "http-429", "http-404"], +) +def test_cold_prompt_miss_never_sleeps_when_langfuse_is_down(status: int, round_trips: int): + """A cold ``get_prompt`` fetches inline on the event loop; with the generated client's default retries a + 429 carrying ``Retry-After: 30`` used to hold the loop for a minute. A 5xx gets the v2 client's one + quick retry, a 429 or 4xx none.""" + requests: list[httpx.Request] = [] + + def fail(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(status, request=request, headers={"retry-after": "30"}, json={"message": "down"}) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(fail)), + ) + + started = monotonic() + with pytest.raises(LangfusePromptError) as caught: + client.get_prompt("greeting") + assert len(requests) == round_trips + assert monotonic() - started < 0.5 + assert caught.value.status_code == status + + +@pytest.mark.parametrize("first_failure", [503, "connect-error"], ids=["http-503", "connect-error"]) +def test_one_transient_failure_on_a_cold_prompt_miss_does_not_fail_the_call(first_failure: int | str): + """The v2 client retried a cold fetch once; a single Langfuse blip must not fail the LLM call.""" + requests: list[httpx.Request] = [] + + def flaky(request: httpx.Request) -> httpx.Response: + requests.append(request) + if len(requests) > 1: + return httpx.Response(200, request=request, json=_TEXT_PROMPT_BODY) + if isinstance(first_failure, int): + return httpx.Response(first_failure, request=request, json={"message": "down"}) + raise httpx.ConnectError("refused", request=request) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(flaky)), + ) + + started = monotonic() + assert client.get_prompt("greeting").compile() == "hello" + assert len(requests) == 2 + assert monotonic() - started < 0.5 + assert client.get_prompt("greeting").compile() == "hello", "the retried prompt is cached like any other" + assert len(requests) == 2 + + +def test_prompt_fetch_error_carries_status_and_body_but_no_upstream_headers(): + """The proxy forwards an exception's ``headers`` to its client and prints ``str(e)``; the generated + ``ApiError`` carries Langfuse's response headers in both.""" + upstream_headers = {"server": "langfuse-edge", "set-cookie": "session=abc; HttpOnly", "x-upstream-internal": "1"} + + def not_found(request: httpx.Request) -> httpx.Response: + return httpx.Response(404, request=request, headers=upstream_headers, json={"message": "Prompt not found"}) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(not_found)), + ) + + with pytest.raises(Exception, match="Prompt not found") as caught: + client.get_prompt("missing") + + error = caught.value + assert getattr(error, "headers", None) is None + assert getattr(error, "status_code", None) == 404 + assert not any(header in str(error) for header in upstream_headers) + assert error.__cause__ is None and error.__suppress_context__, "the header-bearing ApiError must not ride along" + + +_TEXT_PROMPT_BODY: Final[dict[str, object]] = { + "type": "text", + "name": "n", + "version": 1, + "config": {}, + "labels": ["production"], + "tags": [], + "prompt": "hello", +} + + +@pytest.mark.parametrize( + ("name", "encoded"), + [ + ("what?", "what%3F"), + ("folder/greeting", "folder%2Fgreeting"), + ("my-prompt?label=staging", "my-prompt%3Flabel%3Dstaging"), + ("100% sure#1", "100%25%20sure%231"), + ], + ids=["question-mark", "folder-slash", "query-injection", "percent-space-hash"], +) +def test_prompt_name_is_url_encoded_into_the_request_path(name: str, encoded: str): + """The v2 client quoted the name before building the path and the v4 SDK's ``get_prompt`` does too; the + generated client alone puts the raw name into the URL, so ``what?`` fetched prompt ``what`` and + ``a/b`` left the prompts route.""" + requests: list[httpx.Request] = [] + + def record(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, request=request, json=_TEXT_PROMPT_BODY) + + client = build_langfuse_client( + public_key="pk", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=httpx.Client(transport=httpx.MockTransport(record)), + ) + + client.get_prompt(name, label="staging") + + (request,) = requests + assert request.url.raw_path == f"/api/public/v2/prompts/{encoded}?label=staging".encode() + + +def test_rest_client_reports_the_project_id_and_a_passing_auth_check(): + requests: list[httpx.Request] = [] + client = build_langfuse_client( + public_key="pk-project-test", + secret_key="sk", + base_url="http://127.0.0.1:1", + httpx_client=_recording_transport(requests, status=200), + ) + assert client.project_id() == "proj-under-test" + assert client.auth_check() is None + + +def test_rest_client_leaves_a_host_applications_langfuse_client_alone(): + """The SDK hands every ``Langfuse()`` built for one public key the same resource bundle, so a + litellm-built SDK client used to make a host application's client fetch through litellm's + host, secret and httpx client. litellm now speaks REST directly and registers nothing.""" + from langfuse import Langfuse + + requests: list[httpx.Request] = [] + litellm_client = build_langfuse_client( + public_key="pk-shared-with-host", + secret_key="sk-litellm", + base_url="http://litellm.example", + httpx_client=_recording_transport(requests, status=200), + ) + assert litellm_client.project_id() == "proj-under-test" + + host_requests: list[httpx.Request] = [] + host = Langfuse( + public_key="pk-shared-with-host", + secret_key="sk-host", + base_url="http://host.example", + httpx_client=_recording_transport(host_requests, status=200), + tracing_enabled=False, + ) + try: + assert host.auth_check() is True + finally: + host.shutdown() + + assert [request.url.host for request in requests] == ["litellm.example"] + assert host_requests[-1].url.host == "host.example" + assert host_requests[-1].headers["authorization"] == "Basic " + b64encode(b"pk-shared-with-host:sk-host").decode() + + +def test_rest_client_does_not_take_over_the_process_tracer_provider(): + provider_before = otel_trace.get_tracer_provider() + build_langfuse_client( + public_key="pk-sdk-global-test", secret_key="sk", base_url="http://127.0.0.1:1", httpx_client=None + ) + assert otel_trace.get_tracer_provider() is provider_before + + +def _finished_span(): + provider = TracerProvider() + span = provider.get_tracer("t").start_span("generation") + span.end() + return span + + +def _exporter_over(responses, *, delays=(0.5, 1.5), timeout=5.0): + """A LangfuseSpanExporter whose litellm HTTPHandler talks to a scripted transport instead of the network.""" + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + seen = [] + script = list(responses) + + def transport(request: httpx.Request) -> httpx.Response: + seen.append(request) + step = script.pop(0) + if isinstance(step, Exception): + raise step + return httpx.Response(step, request=request) + + handler = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(transport))) + exporter = LangfuseSpanExporter( + handler=handler, + endpoint="https://lf.internal.example/api/public/otel/v1/traces", + headers=MappingProxyType({"Authorization": "Basic cGs6c2s=", "Content-Type": "application/x-protobuf"}), + timeout=timeout, + delays=delays, + ) + return exporter, seen + + +def test_exporter_posts_the_otlp_batch_through_litellm_http_handler(monkeypatch): + """Traces travel through litellm's own handler, so litellm's TLS and proxy settings apply to them.""" + from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest + + slept = [] + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", slept.append) + exporter, seen = _exporter_over([200]) + span = _finished_span() + + assert exporter.export((span,)) is SpanExportResult.SUCCESS + + (request,) = seen + assert request.method == "POST" + assert str(request.url) == "https://lf.internal.example/api/public/otel/v1/traces" + assert request.headers["Authorization"] == "Basic cGs6c2s=" + assert request.headers["Content-Type"] == "application/x-protobuf" + decoded = ExportTraceServiceRequest() + decoded.ParseFromString(request.content) + exported = decoded.resource_spans[0].scope_spans[0].spans[0] + assert exported.name == "generation" + assert exported.span_id == span.context.span_id.to_bytes(8, "big") + assert slept == [] + + +@pytest.mark.parametrize( + "failure", + [httpx.ReadTimeout("stalled"), httpx.ConnectError("refused"), 503, 429, 408, 501, 507, 599], + ids=["read-timeout", "connect-error", "http-503", "http-429", "http-408", "http-501", "http-507", "http-599"], +) +def test_exporter_retries_a_failed_round_trip_and_then_succeeds(monkeypatch, failure): + """A stalled or restarting destination used to drop the batch outright; v2 backed off and re-sent every 5xx.""" + slept = [] + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", slept.append) + exporter, seen = _exporter_over([failure, failure, 200], delays=(0.5, 1.5, 2.5)) + + assert exporter.export((_finished_span(),)) is SpanExportResult.SUCCESS + assert len(seen) == 3 + assert len({request.content for request in seen}) == 1 + assert slept == [0.5, 1.5] + + +def test_exporter_gives_up_after_the_last_delay(monkeypatch): + slept = [] + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", slept.append) + exporter, seen = _exporter_over([httpx.ConnectError("refused")] * 3, delays=(1.0, 2.0)) + + assert exporter.export((_finished_span(),)) is SpanExportResult.FAILURE + assert len(seen) == 3 + assert slept == [1.0, 2.0] + + +def _exporter_with_body_cap(max_bytes: int, *, deliveries: list[int]): + """A destination that answers 413 to any body over ``max_bytes``, the way an ingress with a body limit does.""" + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + def transport(request: httpx.Request) -> httpx.Response: + if len(request.content) > max_bytes: + return httpx.Response(413, request=request) + deliveries.append(len(request.content)) + return httpx.Response(200, request=request) + + return LangfuseSpanExporter( + handler=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(transport))), + endpoint="https://lf.internal.example/api/public/otel/v1/traces", + headers=MappingProxyType({}), + timeout=5.0, + delays=(), + ) + + +def test_exporter_splits_a_batch_the_destination_finds_too_large(monkeypatch): + """One 413 used to drop every span in the batch; the v2 consumer sized its batches by bytes before posting.""" + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + spans = tuple(_finished_span() for _ in range(8)) + whole = _encode(spans) + assert whole is not None + deliveries: list[int] = [] + exporter = _exporter_with_body_cap(len(whole) // 2, deliveries=deliveries) + + assert exporter.export(spans) is SpanExportResult.SUCCESS + assert len(deliveries) >= 2 + assert all(size <= len(whole) // 2 for size in deliveries) + assert ( + sum(deliveries) >= len(whole) - 8 * 8 + ) # each half repeats the resource and scope envelope, spans are not lost + + +def test_exporter_drops_only_the_single_span_that_alone_exceeds_the_cap(monkeypatch, caplog): + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + provider = TracerProvider() + huge = provider.get_tracer("t").start_span("generation", attributes={"body": "x" * 4000}) + huge.end() + small = tuple(_finished_span() for _ in range(3)) + single_small = _encode(small[:1]) + assert single_small is not None + deliveries: list[int] = [] + exporter = _exporter_with_body_cap(len(single_small) * 3, deliveries=deliveries) + + with caplog.at_level(logging.ERROR, logger="LiteLLM"): + result = exporter.export((*small, huge)) + + assert result is SpanExportResult.FAILURE + assert len(deliveries) >= 1 and all(size <= len(single_small) * 3 for size in deliveries) + assert "single" in caplog.text and "too large" in caplog.text + + +def _decoded_attributes(body: bytes) -> dict[str, str]: + decoded = ExportTraceServiceRequest() + decoded.ParseFromString(body) + return { + attribute.key: attribute.value.string_value + for attribute in decoded.resource_spans[0].scope_spans[0].spans[0].attributes + } + + +def _generation_span(**attributes: str): + provider = TracerProvider() + span = provider.get_tracer("t").start_span("generation", attributes=attributes) + span.end() + return span + + +def test_exporter_truncates_a_single_oversized_span_the_way_v2_did_instead_of_dropping_it(monkeypatch, caplog): + """v2 replaced the largest of input, output and metadata with a marker and still delivered the observation; a + vision request over a self-hosted ingress cap used to lose the whole generation, model and usage included.""" + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + span = _generation_span( + **{ + "langfuse.observation.input": "data:image/png;base64," + "A" * 6000, + "langfuse.trace.input": "data:image/png;base64," + "A" * 200, + "langfuse.observation.output": "o" * 1000, + "langfuse.observation.metadata.team": "m" * 100, + "langfuse.observation.model.name": "gpt-4o", + } + ) + bodies: list[bytes] = [] + + def transport(request: httpx.Request) -> httpx.Response: + if len(request.content) > 2000: + return httpx.Response(413, request=request) + bodies.append(request.content) + return httpx.Response(200, request=request) + + exporter = LangfuseSpanExporter( + handler=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(transport))), + endpoint="https://lf.internal.example/api/public/otel/v1/traces", + headers=MappingProxyType({}), + timeout=5.0, + delays=(), + ) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + result = exporter.export((span,)) + + assert result is SpanExportResult.SUCCESS + delivered = _decoded_attributes(bodies[-1]) + assert delivered["langfuse.observation.input"] == "" + assert delivered["langfuse.trace.input"] == "" + assert delivered["langfuse.observation.output"] == "o" * 1000 + assert delivered["langfuse.observation.metadata.team"] == "m" * 100 + assert delivered["langfuse.observation.model.name"] == "gpt-4o" + assert "dropping it" not in caplog.text and "truncated" in caplog.text + + +def test_exporter_truncates_largest_first_and_drops_only_when_nothing_is_left(monkeypatch, caplog): + """Langfuse stores a bare ``langfuse.observation.metadata`` string as nothing, so the metadata marker travels + under a flattened key the way every other metadata value does.""" + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + span = _generation_span( + **{ + "langfuse.observation.input": "i" * 3000, + "langfuse.observation.output": "o" * 2000, + "langfuse.observation.metadata.a": "m" * 500, + "langfuse.trace.metadata.b": "m" * 500, + } + ) + posted: list[dict[str, str]] = [] + + def always_too_large(request: httpx.Request) -> httpx.Response: + posted.append(_decoded_attributes(request.content)) + return httpx.Response(413, request=request) + + exporter = LangfuseSpanExporter( + handler=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(always_too_large))), + endpoint="https://lf.internal.example/api/public/otel/v1/traces", + headers=MappingProxyType({}), + timeout=5.0, + delays=(), + ) + with caplog.at_level(logging.ERROR, logger="LiteLLM"): + assert exporter.export((span,)) is SpanExportResult.FAILURE + + marker = "" + assert [sorted(key for key, value in body.items() if value == marker) for body in posted] == [ + [], + ["langfuse.observation.input"], + ["langfuse.observation.input", "langfuse.observation.output"], + [ + "langfuse.observation.input", + "langfuse.observation.metadata.truncated", + "langfuse.observation.output", + "langfuse.trace.metadata.truncated", + ], + ] + assert "langfuse.observation.metadata.a" not in posted[-1] and "langfuse.trace.metadata.b" not in posted[-1] + assert "dropping it" in caplog.text + + +@pytest.mark.parametrize("status", [400, 401, 403, 404, 422, 499]) +def test_exporter_does_not_retry_a_rejected_batch(monkeypatch, status): + """Bad credentials or a bad payload will not get better on the next attempt, so retrying only delays the flush.""" + slept = [] + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", slept.append) + exporter, seen = _exporter_over([status, 200]) + + assert exporter.export((_finished_span(),)) is SpanExportResult.FAILURE + assert len(seen) == 1 + assert slept == [] + + +def test_exporter_names_the_server_floor_when_the_otlp_route_is_missing(monkeypatch, caplog): + """A Langfuse server too old to serve the OTLP route answers 404; a bare status leaves the operator guessing.""" + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + exporter, _ = _exporter_over([404]) + + with caplog.at_level(logging.ERROR, logger="LiteLLM"): + assert exporter.export((_finished_span(),)) is SpanExportResult.FAILURE + + assert "HTTP 404" in caplog.text and "3.63.0" in caplog.text + + +def _finished_span_named(name: object): + provider = TracerProvider() + span = provider.get_tracer("t").start_span("placeholder") + span._name = name # pyright: ignore[reportAttributeAccessIssue, reportPrivateUsage] # the SDK only stores str + span.end() + return span + + +def test_exporter_drops_a_span_the_encoder_rejects_and_still_posts_the_rest(monkeypatch, caplog): + """One span the OTLP encoder cannot serialize used to raise out of ``export`` and lose every span in the batch.""" + from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest + + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk.sleep", lambda _: None) + exporter, seen = _exporter_over([200]) + + with caplog.at_level(logging.ERROR, logger="LiteLLM"): + result = exporter.export((_finished_span(), _finished_span_named(12345), _finished_span())) + + assert result is SpanExportResult.SUCCESS + (request,) = seen + decoded = ExportTraceServiceRequest() + decoded.ParseFromString(request.content) + assert [span.name for span in decoded.resource_spans[0].scope_spans[0].spans] == ["generation", "generation"] + assert "dropped 1 span(s)" in caplog.text + + +def test_exporter_reports_failure_when_no_span_of_the_batch_can_be_encoded(monkeypatch): + exporter, seen = _exporter_over([200]) + + assert exporter.export((_finished_span_named(12345),)) is SpanExportResult.FAILURE + assert seen == [] + + +def test_built_exporter_uses_the_shared_litellm_handler_and_langfuse_headers(monkeypatch): + """No private requests session or TLS adapter: the channel is the same handler the rest of litellm uses.""" + from litellm.llms.custom_httpx.http_handler import _get_httpx_client + + monkeypatch.delenv("LANGFUSE_TIMEOUT", raising=False) + monkeypatch.delenv("LANGFUSE_MAX_RETRIES", raising=False) + default = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") + assert default.handler is _get_httpx_client() + assert default.timeout == 20 + assert len(default.delays) == 3 + + monkeypatch.setenv("LANGFUSE_TIMEOUT", "7.5") + monkeypatch.setenv("LANGFUSE_MAX_RETRIES", "1") + exporter = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") + assert exporter.endpoint == "https://lf.internal.example/api/public/otel/v1/traces" + assert exporter.timeout == 7.5 + assert exporter.delays == (1.0,) + assert exporter.headers["Authorization"] == "Basic " + b64encode(b"pk:sk").decode() + assert exporter.headers["x-langfuse-public-key"] == "pk" + assert exporter.headers["x-langfuse-sdk-version"] == installed_langfuse_version() + assert exporter.headers["x-langfuse-ingestion-version"] == "4" + + +def test_large_retry_count_builds_an_exporter_with_capped_backoff(monkeypatch): + """``LANGFUSE_MAX_RETRIES=1025`` constructed a v2 client; here ``2.0**1024`` would raise ``OverflowError`` + and take the whole callback down at init.""" + monkeypatch.setenv("LANGFUSE_MAX_RETRIES", "1025") + exporter = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") + assert 3 < len(exporter.delays) <= 1025 + assert exporter.delays[:4] == (1.0, 2.0, 4.0, 8.0) + assert max(exporter.delays) == exporter.delays[-1] <= 64.0 + + +def test_absurd_retry_count_is_clamped_instead_of_allocating_one_delay_per_retry(monkeypatch, caplog): + """A retry count with twelve digits must not turn callback init into a multi-gigabyte tuple allocation.""" + monkeypatch.setenv("LANGFUSE_MAX_RETRIES", "999999999999") + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + exporter = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") + assert 3 < len(exporter.delays) <= 1025 + assert exporter.delays[-1] <= 64.0 + assert any("LANGFUSE_MAX_RETRIES=999999999999" in record.getMessage() for record in caplog.records) + + caplog.clear() + monkeypatch.setenv("LANGFUSE_MAX_RETRIES", "5") + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + modest = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") + assert len(modest.delays) == 5 + assert not any("LANGFUSE_MAX_RETRIES" in record.getMessage() for record in caplog.records) + + +def test_enable_langfuse_debug_logging_makes_deliveries_visible_on_the_langfuse_logger(caplog): + """``LANGFUSE_DEBUG`` turned on the v2 SDK's own logger; it has to do the same for litellm's export channel.""" + exporter, _ = _exporter_over([200]) + langfuse_logger = logging.getLogger("langfuse") + level_before = langfuse_logger.level + try: + with caplog.at_level(logging.INFO, logger="langfuse"): + assert exporter.export((_finished_span(),)) is SpanExportResult.SUCCESS + assert "Exported" not in caplog.text + enable_langfuse_debug_logging() + assert langfuse_logger.level == logging.DEBUG + exporter_after, _ = _exporter_over([200]) + assert exporter_after.export((_finished_span(),)) is SpanExportResult.SUCCESS + assert "Exported" in caplog.text and "lf.internal.example" in caplog.text + finally: + langfuse_logger.setLevel(level_before) + + +@pytest.mark.parametrize( + ("base_url", "export_path", "expected"), + [ + ("https://lf.internal.example/", None, "https://lf.internal.example/api/public/otel/v1/traces"), + ("https://lf.internal.example", "/otel/traces", "https://lf.internal.example/otel/traces"), + ("https://lf.internal.example/", "/otel/traces", "https://lf.internal.example/otel/traces"), + ("https://lf.internal.example", "otel/traces", "https://lf.internal.example/otel/traces"), + ( + "https://lf.internal.example", + "//elsewhere.example/otel", + "https://lf.internal.example/elsewhere.example/otel", + ), + ( + "https://lf.internal.example", + "https://elsewhere.example/otel", + "https://lf.internal.example/https://elsewhere.example/otel", + ), + ], + ids=[ + "default", + "leading-slash", + "both-slashes", + "no-slash", + "scheme-relative-stays-on-host", + "absolute-stays-on-host", + ], +) +def test_export_endpoint_never_doubles_the_slash_or_leaves_the_configured_host( + monkeypatch, base_url, export_path, expected +): + if export_path is None: + monkeypatch.delenv("LANGFUSE_OTEL_TRACES_EXPORT_PATH", raising=False) + else: + monkeypatch.setenv("LANGFUSE_OTEL_TRACES_EXPORT_PATH", export_path) + + exporter = _build_span_exporter(public_key="pk", secret_key="sk", base_url=base_url) + + assert exporter.endpoint == expected + + +class _RecordingPromptsApi: + """Answers ``prompts.get`` with a text prompt that names the label it was asked for.""" + + def __init__(self) -> None: + self.prompts = self + self.requests: list[tuple[str, int | None, str | None]] = [] # mutable-ok: test-side call log + + def get(self, name: str, *, version: int | None, label: str | None, request_options: RequestOptions): + from langfuse.api import Prompt_Text + + assert request_options.get("max_retries") == 0, "a prompt fetch must not sleep through the client's retries" + self.requests.append((name, version, label)) + return Prompt_Text( + name=name, + version=version or 1, + config={}, + labels=[label or "production"], + tags=[], + prompt=f"label={label!r}", + ) + + +class _BlockingPromptsApi(_RecordingPromptsApi): + """Every fetch after the first blocks until the test releases it, and may be told to fail.""" + + def __init__(self) -> None: + super().__init__() + self.release = threading.Event() + self.fail_refresh = False + + def get(self, name: str, *, version: int | None, label: str | None, request_options: RequestOptions): + is_refresh = bool(self.requests) + prompt = super().get(name, version=version, label=label, request_options=request_options) + if is_refresh: + assert self.release.wait(5), "refresh was never released" + if self.fail_refresh: + raise RuntimeError("langfuse is down") + return prompt + + +def _wait_until(predicate, timeout: float = 5.0) -> None: + for _ in range(int(timeout / 0.01)): + if predicate(): + return + sleep(0.01) + raise AssertionError("condition not met in time") + + +def test_stale_prompt_is_served_at_once_while_the_refresh_runs_elsewhere(): + """``get_prompt`` runs on the proxy's event loop; a stale entry used to refetch inline and block every + request on the REST round trip. The stale prompt is returned immediately and refreshed off-thread.""" + api = _BlockingPromptsApi() + client = LangfuseApiClient(api, prompt_cache_ttl_seconds=0) # pyright: ignore[reportArgumentType] # duck-typed prompts API + + first = client.get_prompt("greeting") + started = monotonic() + stale = client.get_prompt("greeting") + + assert stale is first, "the stale prompt must come back without waiting on the refresh" + assert monotonic() - started < 1.0, "the stale read waited on the blocked refresh" + _wait_until(lambda: len(api.requests) == 2) + api.release.set() + _wait_until(lambda: client.get_prompt("greeting") is not first) + + +def test_a_failed_background_refresh_keeps_the_stale_prompt_in_service(caplog): + api = _BlockingPromptsApi() + api.fail_refresh = True + client = LangfuseApiClient(api, prompt_cache_ttl_seconds=0) # pyright: ignore[reportArgumentType] # duck-typed prompts API + + first = client.get_prompt("greeting") + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + started = monotonic() + assert client.get_prompt("greeting") is first + assert monotonic() - started < 1.0, "the stale read waited on the blocked refresh" + api.release.set() + _wait_until(lambda: "refresh failed" in caplog.text) + assert client.get_prompt("greeting") is first + + +def test_only_one_refresh_runs_for_a_stale_prompt_under_concurrent_reads(): + api = _BlockingPromptsApi() + client = LangfuseApiClient(api, prompt_cache_ttl_seconds=0.3) # pyright: ignore[reportArgumentType] # duck-typed prompts API + + first = client.get_prompt("greeting") + sleep(0.3) + for _ in range(20): + assert client.get_prompt("greeting") is first + _wait_until(lambda: len(api.requests) == 2) + api.release.set() + _wait_until(lambda: client.get_prompt("greeting") is not first) + assert len(api.requests) == 2 + + +def test_prompt_cache_keeps_a_missing_label_apart_from_the_label_named_none(): + """A prompt labelled ``"None"`` and the unlabelled default are different prompts in Langfuse + and must not answer each other's requests from the cache.""" + api = _RecordingPromptsApi() + client = LangfuseApiClient(api, prompt_cache_ttl_seconds=60) # pyright: ignore[reportArgumentType] # duck-typed prompts API + + unlabelled = client.get_prompt("greeting") + named_none = client.get_prompt("greeting", label="None") + cached_unlabelled = client.get_prompt("greeting") + + assert unlabelled.prompt == "label=None" + assert named_none.prompt == "label='None'" + assert cached_unlabelled is unlabelled + assert api.requests == [("greeting", None, None), ("greeting", None, "None")] diff --git a/tests/test_litellm/integrations/test_langfuse.py b/tests/test_litellm/integrations/test_langfuse.py index 37860ae8445..3e6e130cac5 100644 --- a/tests/test_litellm/integrations/test_langfuse.py +++ b/tests/test_litellm/integrations/test_langfuse.py @@ -1,6 +1,8 @@ import datetime import json -import sys +import logging +import threading +import time import types import unittest from typing import Final, Optional @@ -11,6 +13,7 @@ import pytest import litellm from litellm.integrations.langfuse import langfuse as langfuse_module from litellm.integrations.langfuse.langfuse import LangFuseLogger +from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id # Import LangfuseUsageDetails directly from the module where it's defined @@ -33,58 +36,20 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) self.env_patcher.start() - # Create mock objects - self.mock_langfuse_client = MagicMock() - # Mock the client attribute to prevent errors during logger initialization - self.mock_langfuse_client.client = MagicMock() - self.mock_langfuse_trace = MagicMock() - self.mock_langfuse_generation = MagicMock() - self.mock_langfuse_generation.trace_id = "test-trace-id" - - # Mock span method for trace (used by log_provider_specific_information_as_span and _log_guardrail_information_as_span) - self.mock_langfuse_span = MagicMock() - self.mock_langfuse_span.end = MagicMock() - self.mock_langfuse_trace.span.return_value = self.mock_langfuse_span - - # Setup the trace and generation chain - self.mock_langfuse_trace.generation.return_value = self.mock_langfuse_generation - self.last_trace_kwargs = {} - - def _trace_side_effect(*args, **kwargs): - self.last_trace_kwargs = kwargs - return self.mock_langfuse_trace - - self.mock_langfuse_client.trace.side_effect = _trace_side_effect - - # Mock the langfuse module that's imported locally in methods - self.langfuse_module_patcher = patch.dict( - "sys.modules", {"langfuse": MagicMock()} - ) - self.mock_langfuse_module = self.langfuse_module_patcher.start() - - # Create a mock for the langfuse module with version - self.mock_langfuse = MagicMock() - self.mock_langfuse.version = MagicMock() - self.mock_langfuse.version.__version__ = ( - "3.0.0" # Set a version that supports all features + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, ) - # Mock the Langfuse class - self.mock_langfuse_class = MagicMock() - self.mock_langfuse_class.return_value = self.mock_langfuse_client + self.span_exporter = InMemorySpanExporter() + self.real_provider = TracerProvider() + self.real_provider.add_span_processor(SimpleSpanProcessor(self.span_exporter)) - # Set up the sys.modules['langfuse'] mock - sys.modules["langfuse"] = self.mock_langfuse - sys.modules["langfuse"].Langfuse = self.mock_langfuse_class - - # Create a fresh logger instance for each test + # the host above is unreachable, so the REST client is cheap to build + # and each test swaps in the export channel it wants self.logger = LangFuseLogger() - # Explicitly set the Langfuse client to our mock - self.logger.Langfuse = self.mock_langfuse_client - # Ensure langfuse_sdk_version is set correctly for _supports_* methods - self.logger.langfuse_sdk_version = "3.0.0" - # Add the log_event_on_langfuse method to the instance def log_event_on_langfuse( self, @@ -113,23 +78,46 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) # Bind the method to the instance - self.logger.log_event_on_langfuse = types.MethodType( - log_event_on_langfuse, self.logger - ) + self.logger.log_event_on_langfuse = types.MethodType(log_event_on_langfuse, self.logger) def tearDown(self): # Clean up logger instance to prevent state leakage if hasattr(self, "logger"): - # Reset logger's Langfuse client to break any references - self.logger.Langfuse = None - # Delete logger instance to ensure complete cleanup del self.logger # Restore global Langfuse client counter to prevent cross-test pollution litellm.initialized_langfuse_clients = self._original_langfuse_clients_count self.env_patcher.stop() - self.langfuse_module_patcher.stop() # patch.dict automatically restores sys.modules + + def use_real_langfuse_client(self): + """Point the logger at an export channel whose spans land in memory.""" + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + + from litellm.integrations.langfuse.langfuse_sdk import build_langfuse_tracing + + self.span_exporter = InMemorySpanExporter() + self.logger.tracing = build_langfuse_tracing( + exporter=self.span_exporter, + environment=None, + release=None, + sample_rate=1.0, + flush_interval_millis=10, + ) + self.real_provider = self.logger.tracing.provider + return self.logger.tracing + + def exported_generation(self): + self.logger.tracing.flush() + spans = [s for s in self.span_exporter.get_finished_spans()] + assert spans, "no spans were exported" + return spans[-1] + + @staticmethod + def span_trace_id(span): + return format(span.context.trace_id, "032x") def test_langfuse_usage_details_type(self): """Test that LangfuseUsageDetails TypedDict is properly defined with the correct fields""" @@ -260,21 +248,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): Test that _log_langfuse_v2 correctly handles None values in the usage object by converting them to 0, preventing validation errors. """ - # Reset the mock to ensure clean state; clear side_effect so return_value takes effect - self.mock_langfuse_client.reset_mock(side_effect=True) - self.mock_langfuse_trace.reset_mock(side_effect=True) - self.mock_langfuse_generation.reset_mock(side_effect=True) - - # Re-setup the trace and generation chain with clean state - self.mock_langfuse_generation.trace_id = "test-trace-id" - mock_span = MagicMock() - mock_span.end = MagicMock() - self.mock_langfuse_trace.span.return_value = mock_span - self.mock_langfuse_trace.generation.return_value = self.mock_langfuse_generation - - # Ensure trace returns our mock - self.mock_langfuse_client.trace.return_value = self.mock_langfuse_trace - self.logger.Langfuse = self.mock_langfuse_client + self.use_real_langfuse_client() with ( patch( @@ -282,7 +256,6 @@ class TestLangfuseUsageDetails(unittest.TestCase): side_effect=lambda generation_params, **kwargs: generation_params, create=True, ) as mock_add_prompt_params, - patch.object(self.logger, "_supports_prompt", return_value=True), ): # Create a mock response object with usage information containing None values response_obj = MagicMock() @@ -332,29 +305,12 @@ class TestLangfuseUsageDetails(unittest.TestCase): except Exception as e: self.fail(f"_log_langfuse_v2 raised an exception: {e}") - # Verify that trace was called first - self.mock_langfuse_client.trace.assert_called() - - # Check the arguments passed to the mocked langfuse generation call - self.mock_langfuse_trace.generation.assert_called_once() - call_args, call_kwargs = self.mock_langfuse_trace.generation.call_args - - # Inspect the usage and usage_details dictionaries - usage_arg = call_kwargs.get("usage") - usage_details_arg = call_kwargs.get("usage_details") - - self.assertIsNotNone(usage_arg) - self.assertIsNotNone(usage_details_arg) - - # Verify that None values were converted to 0 - self.assertEqual(usage_arg["prompt_tokens"], 0) - self.assertEqual(usage_arg["completion_tokens"], 0) - - self.assertEqual(usage_details_arg["input"], 0) - self.assertEqual(usage_details_arg["output"], 0) - self.assertEqual(usage_details_arg["total"], 0) - self.assertEqual(usage_details_arg["cache_creation_input_tokens"], 0) - self.assertEqual(usage_details_arg["cache_read_input_tokens"], 0) + usage_details = json.loads(self.exported_generation().attributes["langfuse.observation.usage_details"]) + assert usage_details["input"] == 0 + assert usage_details["output"] == 0 + assert usage_details["total"] == 0 + assert usage_details["cache_creation_input_tokens"] == 0 + assert usage_details["cache_read_input_tokens"] == 0 mock_add_prompt_params.assert_called_once() @@ -407,7 +363,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): def test_log_langfuse_v2_uses_standard_trace_id_when_available(self): payload = self._build_standard_logging_payload(trace_id="std-trace-id") kwargs = self._build_langfuse_kwargs(payload) - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -429,12 +385,12 @@ class TestLangfuseUsageDetails(unittest.TestCase): litellm_call_id="call-id-xyz", ) - assert self.last_trace_kwargs.get("id") == "std-trace-id" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("std-trace-id") def test_log_langfuse_v2_defaults_to_call_id_without_standard_trace_id(self): payload = self._build_standard_logging_payload() kwargs = self._build_langfuse_kwargs(payload) - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -456,7 +412,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): litellm_call_id="call-id-xyz", ) - assert self.last_trace_kwargs.get("id") == "call-id-xyz" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("call-id-xyz") def test_log_langfuse_v2_uses_litellm_trace_id_fallback_over_call_id(self): """ @@ -468,7 +424,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): payload = self._build_standard_logging_payload() # no trace_id kwargs = self._build_langfuse_kwargs(payload) kwargs["litellm_trace_id"] = "trace-id-from-kwargs" - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -491,7 +447,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) # litellm_trace_id should be preferred over litellm_call_id - assert self.last_trace_kwargs.get("id") == "trace-id-from-kwargs" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("trace-id-from-kwargs") CANARY = "sk-lf-canary-SECRET-d4e5f6" @@ -521,24 +477,51 @@ class TestLangfuseUsageDetails(unittest.TestCase): } def _emitted_payload_text(self): - """Every blob this logger handed to the langfuse SDK, as one searchable string.""" + """Every attribute this logger exported to langfuse, as one searchable string.""" import json - blobs = [self.last_trace_kwargs] - if self.mock_langfuse_trace.generation.call_args is not None: - blobs.append(self.mock_langfuse_trace.generation.call_args.kwargs) - blobs.extend(call.kwargs for call in self.mock_langfuse_trace.span.call_args_list) - return json.dumps(blobs, default=repr) + self.logger.tracing.flush() + return json.dumps( + [dict(span.attributes or {}) for span in self.span_exporter.get_finished_spans()], + default=repr, + ) - def _drive_with_canary(self, extra_metadata=None, hidden_params=None): + def exported_generation_metadata(self): + """The generation's metadata as langfuse receives it, one attribute per key. + + v4 serializes each value onto the span, so they are decoded back here to + keep these assertions about what litellm emitted rather than about the + SDK's wire encoding. + """ + import json + + prefix = "langfuse.observation.metadata." + + def decoded(raw): + try: + return json.loads(raw) + except (TypeError, ValueError): + return raw + + return { + key[len(prefix) :]: decoded(value) + for key, value in (self.exported_generation().attributes or {}).items() + if key.startswith(prefix) + } + + def exported_spans_named(self, name): + self.logger.tracing.flush() + return [span for span in self.span_exporter.get_finished_spans() if span.name == name] + + def _drive_with_canary(self, extra_metadata=None, hidden_params=None, guardrail_information=None): metadata = {**self._canary_request_metadata(), **(extra_metadata or {})} payload = self._build_standard_logging_payload(trace_id="canary-trace-id") if hidden_params is not None: payload["hidden_params"] = hidden_params + if guardrail_information is not None: + payload["guardrail_information"] = guardrail_information kwargs = {**self._build_langfuse_kwargs(payload), "response_cost": 0.25} - self.last_trace_kwargs = {} - self.mock_langfuse_trace.generation.reset_mock() - self.mock_langfuse_trace.span.reset_mock() + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -559,7 +542,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): level="INFO", litellm_call_id="canary-call-id", ) - return self.mock_langfuse_trace.generation.call_args.kwargs["metadata"] + return self.exported_generation_metadata() def test_team_callback_credentials_never_reach_langfuse(self): """ @@ -583,10 +566,13 @@ class TestLangfuseUsageDetails(unittest.TestCase): debug_langfuse dumps request metadata into the trace as a second emit site. It must be sourced from the allowlisted payload too. """ - self._drive_with_canary(extra_metadata={"debug_langfuse": True}) + import json + + self._drive_with_canary(extra_metadata={"debug_langfuse": True}) + dumped = json.loads(self.exported_generation().attributes["langfuse.trace.metadata.metadata_passed_to_litellm"]) - dumped = self.last_trace_kwargs["metadata"]["metadata_passed_to_litellm"] assert "user_api_key_auth" not in dumped + assert dumped["first_custom"] == "keep-first" assert self.CANARY not in self._emitted_payload_text() def test_raw_request_metadata_reaches_the_emitted_blob_through_no_key(self): @@ -610,18 +596,68 @@ class TestLangfuseUsageDetails(unittest.TestCase): """ self._drive_with_canary(hidden_params={"vertex_ai_grounding_metadata": ["ground-a", "ground-b"]}) - span_inputs = [call.kwargs.get("input") for call in self.mock_langfuse_trace.span.call_args_list] + span_inputs = [ + span.attributes.get("langfuse.observation.input") + for span in self.exported_spans_named("vertex_ai_grounding_metadata") + ] assert span_inputs == ["ground-a", "ground-b"] assert self.CANARY not in self._emitted_payload_text() + def test_only_the_generation_claims_the_trace_root(self): + """ + Langfuse derives trace name and I/O from the root observation, and with several + roots the one with the latest start wins. A post-call guardrail starts after the + model call, so it must nest under the generation instead of being a root itself, or + the trace shows the guardrail's request instead of the model's. + """ + self._drive_with_canary( + hidden_params={"vertex_ai_grounding_metadata": ["ground-a"]}, + guardrail_information=[ + { + "guardrail_name": "pii-post", + "guardrail_mode": "post_call", + "guardrail_request": {"texts": ["post-call scan"]}, + "guardrail_response": {"flagged": False}, + "start_time": 1704110402.0, + "end_time": 1704110403.0, + } + ], + ) + + [generation] = [span for span in self.span_exporter.get_finished_spans() if span.name.startswith("litellm-")] + [guardrail] = self.exported_spans_named("guardrail") + [grounding] = self.exported_spans_named("vertex_ai_grounding_metadata") + assert generation.parent is None + assert generation.attributes["langfuse.trace.name"] == "canary-trace" + for child in (guardrail, grounding): + assert child.parent.span_id == generation.context.span_id + assert child.context.trace_id == generation.context.trace_id + assert "langfuse.trace.name" not in child.attributes + + def test_generation_is_exported_when_a_child_span_fails(self): + """v2 buffered the generation in one call, so a bad guardrail entry could not lose it; + the OTel generation is open until ``end()`` and must still be ended when a child raises.""" + self._drive_with_canary( + guardrail_information=[ + { + "guardrail_name": "pii-post", + "guardrail_mode": "post_call", + "start_time": "not-a-timestamp", + "end_time": 1704110403.0, + } + ], + ) + + [generation] = [span for span in self.span_exporter.get_finished_spans() if span.name.startswith("litellm-")] + assert generation.attributes["langfuse.trace.name"] == "canary-trace" + assert self.exported_spans_named("guardrail") == [] + def test_caller_cannot_spoof_an_allowlisted_identity_field(self): """ Request metadata never reaches the blob, so a caller naming user_api_key_alias cannot have their value emitted in place of the proxy-resolved one. """ - generation_metadata = self._drive_with_canary( - extra_metadata={"user_api_key_alias": "spoofed-by-caller"} - ) + generation_metadata = self._drive_with_canary(extra_metadata={"user_api_key_alias": "spoofed-by-caller"}) assert generation_metadata["user_api_key_alias"] == "canary-alias" @@ -636,7 +672,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): payload["metadata"]["requester_metadata"] = {"litellm_response_cost": "caller-value", "api_base": "caller"} kwargs = {**self._build_langfuse_kwargs(payload), "response_cost": 0.25} metadata = self._canary_request_metadata() - self.mock_langfuse_trace.generation.reset_mock() + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -658,10 +694,48 @@ class TestLangfuseUsageDetails(unittest.TestCase): litellm_call_id="canary-call-id", ) - generation_metadata = self.mock_langfuse_trace.generation.call_args.kwargs["metadata"] + generation_metadata = self.exported_generation_metadata() assert generation_metadata["litellm_response_cost"] == 0.25 assert generation_metadata["api_base"] == "https://real-api-base" + def test_generation_metadata_carries_the_call_id_and_response_id(self): + """ + v2's generation id was ``time-_``, so a generation could + be found from the provider response id. v4 hashes that string onto 16 hex chars, + which leaves nothing searchable unless both ids are emitted as metadata. + """ + payload = self._build_standard_logging_payload(trace_id="canary-trace-id") + kwargs = {**self._build_langfuse_kwargs(payload), "response_cost": 0.25} + metadata = self._canary_request_metadata() + self.use_real_langfuse_client() + + with patch( + "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", + side_effect=lambda generation_params, **kw: generation_params, + create=True, + ): + self.logger._log_langfuse_v2( + user_id="user-1", + metadata=metadata, + litellm_params={"metadata": metadata}, + output=None, + start_time=datetime.datetime(2024, 1, 1, 12, 0, 0), + end_time=datetime.datetime(2024, 1, 1, 12, 0, 1), + kwargs=kwargs, + optional_params={}, + input=None, + response_obj=litellm.ModelResponse( + id="chatcmpl-canary-response", choices=[{"message": {"role": "assistant", "content": "OK"}}] + ), + level="DEFAULT", + litellm_call_id="canary-call-id", + ) + + generation_metadata = self.exported_generation_metadata() + assert generation_metadata["litellm_call_id"] == "canary-call-id" + assert generation_metadata["response_id"] == "chatcmpl-canary-response" + assert "chatcmpl-canary-response" in self._emitted_payload_text() + def test_denied_steering_keys_and_enrichments(self): """ endpoint is a plain string, so without the deny-list it would ride the @@ -726,8 +800,9 @@ class TestLangfuseUsageDetails(unittest.TestCase): """ self._drive_with_canary() - assert self.last_trace_kwargs.get("session_id") == "canary-session" - assert self.last_trace_kwargs.get("name") == "canary-trace" + generation = self.exported_generation() + assert generation.attributes["session.id"] == "canary-session" + assert generation.attributes["langfuse.trace.name"] == "canary-trace" def test_failure_trace_survives_a_missing_standard_logging_object(self): """ @@ -746,8 +821,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): "messages": [], "litellm_trace_id": "trace-id-failure", } - self.last_trace_kwargs = {} - self.mock_langfuse_trace.generation.reset_mock() + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -771,9 +845,12 @@ class TestLangfuseUsageDetails(unittest.TestCase): import json - assert trace_id == "trace-id-failure" - assert self.last_trace_kwargs.get("id") == "trace-id-failure" - generation_metadata = self.mock_langfuse_trace.generation.call_args.kwargs["metadata"] + # Must use litellm_trace_id, not litellm_call_id. v4 addresses a trace by a + # 32-hex id, so the callback returns the resolved form, which is what makes + # the alerting deep link point at a trace langfuse can actually open + assert trace_id == resolve_trace_id("trace-id-failure") + assert self.span_trace_id(self.exported_generation()) == trace_id + generation_metadata = self.exported_generation_metadata() assert "user_api_key_auth" not in generation_metadata assert self.CANARY not in self._emitted_payload_text() assert "first_custom" not in generation_metadata @@ -790,7 +867,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): """ payload = self._build_standard_logging_payload(trace_id="std-trace-123") kwargs = self._build_langfuse_kwargs(payload) - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -813,9 +890,9 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) # session_id should be set for Langfuse session grouping - assert self.last_trace_kwargs.get("session_id") == "my-session-abc" + assert self.exported_generation().attributes["session.id"] == "my-session-abc" # trace_id should remain the standard trace_id, NOT the session_id - assert self.last_trace_kwargs.get("id") == "std-trace-123" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("std-trace-123") def test_log_langfuse_v2_session_id_preserved_for_error_level(self): """ @@ -825,7 +902,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): """ payload = self._build_standard_logging_payload(trace_id="std-trace-err") kwargs = self._build_langfuse_kwargs(payload) - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -848,11 +925,11 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) # session_id must be preserved even for ERROR level logs - assert self.last_trace_kwargs.get("session_id") == "error-session-xyz" + assert self.exported_generation().attributes["session.id"] == "error-session-xyz" # trace_id should be the standard trace_id, not the session_id - assert self.last_trace_kwargs.get("id") == "std-trace-err" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("std-trace-err") # status_message should be set for error traces - assert self.last_trace_kwargs.get("status_message") is not None + assert self.exported_generation().attributes["langfuse.observation.level"] == "ERROR" def test_log_langfuse_v2_explicit_trace_id_takes_priority_over_session_id(self): """ @@ -861,7 +938,7 @@ class TestLangfuseUsageDetails(unittest.TestCase): """ payload = self._build_standard_logging_payload() kwargs = self._build_langfuse_kwargs(payload) - self.last_trace_kwargs = {} + self.use_real_langfuse_client() with patch( "litellm.integrations.langfuse.langfuse._add_prompt_to_generation_params", @@ -892,9 +969,9 @@ class TestLangfuseUsageDetails(unittest.TestCase): ) # Explicit trace_id must take priority - assert self.last_trace_kwargs.get("id") == "explicit-trace-id-777" + assert self.span_trace_id(self.exported_generation()) == resolve_trace_id("explicit-trace-id-777") # session_id must still be set for session grouping - assert self.last_trace_kwargs.get("session_id") == "session-999" + assert self.exported_generation().attributes["session.id"] == "session-999" def test_failure_handler_langfuse_kwargs_excludes_original_response(): @@ -942,12 +1019,10 @@ def test_failure_handler_langfuse_kwargs_excludes_original_response(): try: # Mock LangFuseHandler to return our capturing mock logger - with patch( - "litellm.litellm_core_utils.litellm_logging.LangFuseHandler" - ) as mock_handler_class: - mock_handler_class.get_langfuse_logger_for_request.return_value = ( - mock_langfuse_logger - ) + with ( + patch("litellm.litellm_core_utils.litellm_logging.LangFuseHandler") as mock_handler_class + ): # test-quality-ok: route the request to the capturing logger; the real handler builds live clients + mock_handler_class.get_langfuse_logger_for_request.return_value = mock_langfuse_logger # Call the actual failure_handler test_exception = Exception("TestError: model not found") @@ -959,23 +1034,19 @@ def test_failure_handler_langfuse_kwargs_excludes_original_response(): ) # Verify log_event_on_langfuse was actually called - assert ( - mock_langfuse_logger.log_event_on_langfuse.called - ), "log_event_on_langfuse was not called" + assert mock_langfuse_logger.log_event_on_langfuse.called, "log_event_on_langfuse was not called" # Verify original_response is NOT in the kwargs passed to Langfuse langfuse_kwargs = captured_kwargs.get("kwargs", {}) - assert ( - "original_response" not in langfuse_kwargs - ), "original_response should be excluded from kwargs passed to Langfuse" + assert "original_response" not in langfuse_kwargs, ( + "original_response should be excluded from kwargs passed to Langfuse" + ) # Verify session_id metadata is preserved in the kwargs - langfuse_metadata = langfuse_kwargs.get("litellm_params", {}).get( - "metadata", {} + langfuse_metadata = langfuse_kwargs.get("litellm_params", {}).get("metadata", {}) + assert langfuse_metadata.get("session_id") == "test-session-failure", ( + "session_id should be preserved in kwargs passed to Langfuse" ) - assert ( - langfuse_metadata.get("session_id") == "test-session-failure" - ), "session_id should be preserved in kwargs passed to Langfuse" # Verify level is ERROR assert captured_kwargs.get("level") == "ERROR" @@ -1017,9 +1088,9 @@ async def test_async_log_failure_event_logs_to_langfuse(): "generation_id": "mock-gen", } - with patch( - "litellm.integrations.langfuse.langfuse_prompt_management.LangFuseHandler" - ) as mock_handler: + with ( + patch("litellm.integrations.langfuse.langfuse_prompt_management.LangFuseHandler") as mock_handler + ): # test-quality-ok: route the request to the capturing logger; the real handler builds live clients mock_handler.get_langfuse_logger_for_request.return_value = mock_logger kwargs = { @@ -1044,9 +1115,7 @@ async def test_async_log_failure_event_logs_to_langfuse(): ) # Verify log_event_on_langfuse was called - assert ( - mock_logger.log_event_on_langfuse.called - ), "log_event_on_langfuse was not called for failure event" + assert mock_logger.log_event_on_langfuse.called, "log_event_on_langfuse was not called for failure event" call_kwargs = mock_logger.log_event_on_langfuse.call_args[1] assert call_kwargs["level"] == "ERROR" assert call_kwargs["status_message"] == "API error: model not found" @@ -1086,9 +1155,9 @@ async def test_async_log_failure_event_works_without_standard_logging_object(): "generation_id": "mock-gen", } - with patch( - "litellm.integrations.langfuse.langfuse_prompt_management.LangFuseHandler" - ) as mock_handler: + with ( + patch("litellm.integrations.langfuse.langfuse_prompt_management.LangFuseHandler") as mock_handler + ): # test-quality-ok: route the request to the capturing logger; the real handler builds live clients mock_handler.get_langfuse_logger_for_request.return_value = mock_logger kwargs = { @@ -1119,6 +1188,77 @@ async def test_async_log_failure_event_works_without_standard_logging_object(): assert "InternalServerError" in call_kwargs["status_message"] +class _OtlpReceiver: + """A local HTTP server that records the paths of every POST it gets, standing in for Langfuse.""" + + def __init__(self) -> None: + from http.server import BaseHTTPRequestHandler, HTTPServer + + self.received: list[str] = [] + received = self.received + + class _Handler(BaseHTTPRequestHandler): + def do_POST(self): + received.append(self.path) + self.rfile.read(int(self.headers.get("Content-Length") or 0)) + self.send_response(200) + self.send_header("Content-Length", "0") + self.end_headers() + + def log_message(self, *args): + pass + + self.server = HTTPServer(("127.0.0.1", 0), _Handler) + threading.Thread(target=self.server.serve_forever, daemon=True).start() + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self.server.server_port}" + + def close(self) -> None: + self.server.shutdown() + + +def _log_one_completion(logger: LangFuseLogger) -> None: + now = datetime.datetime.now() + logger.log_event_on_langfuse( + kwargs={ + "call_type": "completion", + "litellm_params": {"metadata": {}, "proxy_server_request": {"headers": {}}}, + "messages": [{"role": "user", "content": "hi"}], + "optional_params": {}, + }, + response_obj=litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "yo"}}]), + start_time=now, + end_time=now, + ) + logger.flush() + + +def test_mock_mode_makes_no_network_calls(monkeypatch): + """LANGFUSE_MOCK promises full execution without egress. + + The mock intercepts httpx, but v4 ships observations over its own OTLP + exporter, so nothing stops a real request to the configured host without an + exporter that drops them. + """ + receiver = _OtlpReceiver() + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_HOST", receiver.url) + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-mock-egress") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-mock-egress") + + try: + logger = LangFuseLogger() + assert logger.is_mock_mode is True + _log_one_completion(logger) + time.sleep(1) + finally: + receiver.close() + + assert receiver.received == [], f"mock mode sent real requests: {receiver.received}" + + def test_max_langfuse_clients_limit(): """ Test that the max langfuse clients limit is respected when initializing multiple clients @@ -1154,7 +1294,7 @@ def test_max_langfuse_clients_limit(): assert litellm.initialized_langfuse_clients == 2 # Third client should fail with exception - with pytest.raises(Exception, match='Max langfuse clients reached') as exc_info: + with pytest.raises(Exception, match="Max langfuse clients reached") as exc_info: logger3 = LangFuseLogger( langfuse_public_key="test_key_3", langfuse_secret="test_secret_3", @@ -1170,73 +1310,76 @@ def test_max_langfuse_clients_limit(): litellm.initialized_langfuse_clients = original_initialized_langfuse_clients -class _RecordingLangfuse: - last_parameters: Optional[dict] = None - - def __init__(self, environment=None, **parameters): - type(self).last_parameters = {"environment": environment, **parameters} - self.client = MagicMock() +_UNREACHABLE_HOST: Final = "http://127.0.0.1:1" -class _RecordingLangfuseWithoutEnvironment: - last_parameters: Optional[dict] = None - - def __init__(self, **parameters): - type(self).last_parameters = parameters - self.client = MagicMock() - - -def _build_langfuse_logger(monkeypatch) -> LangFuseLogger: +def _build_langfuse_logger(monkeypatch, **overrides) -> LangFuseLogger: monkeypatch.setenv("LANGFUSE_MOCK", "false") monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuse): - return LangFuseLogger( - langfuse_public_key="pk-lit5228", - langfuse_secret="sk-lit5228", - langfuse_host="https://test.langfuse.com", - ) + return LangFuseLogger( + **{ + "langfuse_public_key": "pk-lit5228", + "langfuse_secret": "sk-lit5228", + "langfuse_host": _UNREACHABLE_HOST, + **overrides, + } + ) -def test_langfuse_environment_is_passed_to_sdk_client(monkeypatch): - monkeypatch.setenv("LANGFUSE_MOCK", "false") +def _exported_environment(logger: LangFuseLogger): + from langfuse import LangfuseOtelSpanAttributes + + return logger.tracing.provider.resource.attributes.get(LangfuseOtelSpanAttributes.ENVIRONMENT) + + +def test_langfuse_environment_lands_on_every_exported_span(monkeypatch): monkeypatch.delenv("LANGFUSE_TRACING_ENVIRONMENT", raising=False) - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuse): - logger = LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - langfuse_environment="staging", - ) + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-env", langfuse_environment="staging") assert logger.langfuse_environment == "staging" - assert _RecordingLangfuse.last_parameters["environment"] == "staging" + assert _exported_environment(logger) == "staging" def test_langfuse_environment_falls_back_to_deployment_env_var(monkeypatch): - monkeypatch.setenv("LANGFUSE_MOCK", "false") monkeypatch.setenv("LANGFUSE_TRACING_ENVIRONMENT", "deployment-wide") - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuse): - logger = LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - ) + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-env") assert logger.langfuse_environment == "deployment-wide" - assert _RecordingLangfuse.last_parameters["environment"] == "deployment-wide" + assert _exported_environment(logger) == "deployment-wide" -def test_langfuse_environment_omitted_for_old_sdk_versions(monkeypatch): - monkeypatch.setenv("LANGFUSE_MOCK", "false") - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuseWithoutEnvironment): - LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - langfuse_environment="staging", - ) - assert "environment" not in _RecordingLangfuseWithoutEnvironment.last_parameters +def _exported_release(logger: LangFuseLogger): + from langfuse import LangfuseOtelSpanAttributes + + return logger.tracing.provider.resource.attributes.get(LangfuseOtelSpanAttributes.RELEASE) + + +@pytest.mark.parametrize("platform_var", ["GITHUB_SHA", "CI_COMMIT_SHA", "RENDER_GIT_COMMIT", "SOURCE_VERSION"]) +def test_release_falls_back_to_the_deploy_platforms_commit_variable(monkeypatch, platform_var): + """Deployments that never set ``LANGFUSE_RELEASE`` still got a release on every trace from the v2 SDK, which + read the CI or hosting platform's commit variable; dropping that silently blanked their release filter.""" + from litellm.integrations.langfuse.langfuse_sdk import _COMMON_RELEASE_ENVS + + for name in ("LANGFUSE_RELEASE", *_COMMON_RELEASE_ENVS): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv(platform_var, "deadbeef") + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key=f"pk-release-{platform_var}") + assert logger.langfuse_release == "deadbeef" + assert _exported_release(logger) == "deadbeef" + + +def test_explicit_langfuse_release_wins_over_the_platform_commit(monkeypatch): + monkeypatch.setenv("LANGFUSE_RELEASE", "v9") + monkeypatch.setenv("GITHUB_SHA", "deadbeef") + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-release-explicit") + assert _exported_release(logger) == "v9" + + +def test_non_string_generation_name_is_exported_as_its_text(monkeypatch): + """v2 coerced ``generation_name`` through pydantic; a raw int would now fail OTLP encoding and lose the batch.""" + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"generation_name": 12345}) + + assert span.name == "12345" def test_dynamic_langfuse_environment_triggers_dynamic_logger(): @@ -1247,13 +1390,11 @@ def test_dynamic_langfuse_environment_triggers_dynamic_logger(): assert LangFuseHandler._dynamic_langfuse_credentials_are_passed(params) is True - config = LangFuseHandler.get_dynamic_langfuse_logging_config( - standard_callback_dynamic_params=params - ) + config = LangFuseHandler.get_dynamic_langfuse_logging_config(standard_callback_dynamic_params=params) assert config["langfuse_environment"] == "team-a-env" -def test_langfuse_sdk_client_survives_httpx_cache_eviction(monkeypatch): +def test_langfuse_rest_client_survives_httpx_cache_eviction(monkeypatch): import gc import weakref @@ -1263,21 +1404,20 @@ def test_langfuse_sdk_client_survives_httpx_cache_eviction(monkeypatch): monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) logger = _build_langfuse_logger(monkeypatch) - sdk_client = _RecordingLangfuse.last_parameters["httpx_client"] cached_handler = _get_httpx_client() handler_ref = weakref.ref(cached_handler) - assert sdk_client is logger.langfuse_client - assert sdk_client is cached_handler.client + assert logger.langfuse_client is cached_handler.client litellm.in_memory_llm_clients_cache = LLMClientCache() del cached_handler gc.collect() assert litellm.in_memory_llm_clients_cache.get_cache("httpx_client") is None - assert handler_ref() is not None, "logger must keep the handler that owns the client it handed the SDK" - assert not sdk_client.is_closed + assert handler_ref() is not None, "logger must keep the handler that owns the client behind its REST API" + assert not logger.langfuse_client.is_closed + assert logger.api_client.auth_check() is not None def test_langfuse_logger_reuses_the_shared_cached_client(monkeypatch): @@ -1301,20 +1441,100 @@ def test_langfuse_logger_reuses_the_shared_cached_client(monkeypatch): _LANGFUSE_REDACTED = "redacted-by-litellm" -def _steering_logger() -> LangFuseLogger: - """``__new__`` skips the SDK and network setup in ``__init__``.""" - logger = LangFuseLogger.__new__(LangFuseLogger) - logger.Langfuse = MagicMock() - logger.langfuse_sdk_version = "2.60.0" - return logger - - -def _emit(logger: LangFuseLogger, *, metadata=None, headers=None): - """``log_event_on_langfuse`` is the entry point that folds ``langfuse_*`` headers into metadata.""" - now = datetime.datetime.now() - response_obj = litellm.ModelResponse( - choices=[{"message": {"role": "assistant", "content": "the-output"}}] +def _steering_logger(): + """``__new__`` skips the network setup in ``__init__``; spans land in memory.""" + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, ) + + from litellm.integrations.langfuse.langfuse import installed_langfuse_version + from litellm.integrations.langfuse.langfuse_sdk import build_langfuse_client, build_langfuse_tracing + + exporter = InMemorySpanExporter() + logger = LangFuseLogger.__new__(LangFuseLogger) + logger.tracing = build_langfuse_tracing( + exporter=exporter, environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + logger.api_client = build_langfuse_client( + public_key="pk-steering-test", secret_key="sk-steering-test", base_url=_UNREACHABLE_HOST, httpx_client=None + ) + logger.langfuse_sdk_version = installed_langfuse_version() + return logger, exporter + + +def test_log_event_keeps_exporting_after_the_dynamic_cache_evicts_the_logger(): + """Per-key loggers are evicted from ``DynamicLoggingCache`` while a callback may still hold them. + + v2 lost that callback's events to a shut-down client; the export channel is shared per + credential set and outlives any one logger, so the events still land. + """ + from litellm.litellm_core_utils.specialty_caches.dynamic_logging_cache import LangfuseInMemoryCache + + logger, exporter = _steering_logger() + cache = LangfuseInMemoryCache() + cache.set_cache("langfuse-evicted", logger) + litellm.initialized_langfuse_clients += 1 + before = litellm.initialized_langfuse_clients + cache._remove_key("langfuse-evicted") + + now = datetime.datetime.now() + returned = logger.log_event_on_langfuse( + kwargs={ + "call_type": "completion", + "litellm_params": {"metadata": {}}, + "messages": [{"role": "user", "content": "the-input"}], + "optional_params": {}, + }, + response_obj=litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "the-output"}}]), + start_time=now, + end_time=now, + ) + + assert litellm.initialized_langfuse_clients == before - 1 + assert _span_trace_id(_exported_span(logger, exporter)) == returned["trace_id"] + + +def _exported_span(logger, exporter): + logger.flush() + return exporter.get_finished_spans()[-1] + + +_TRACE_FIELD_KEYS = { + "user.id": "user_id", + "session.id": "session_id", + "langfuse.version": "version", + "langfuse.release": "release", +} + + +def _trace_params(span): + """The trace-level fields of the exported span, keyed as v2's ``trace_params`` were.""" + prefix = "langfuse.trace." + attributes = span.attributes or {} + return { + **{ + key[len(prefix) :]: value + for key, value in attributes.items() + if key.startswith(prefix) and not key.startswith(prefix + "metadata.") + }, + **{name: attributes[key] for key, name in _TRACE_FIELD_KEYS.items() if key in attributes}, + } + + +def _span_trace_id(span): + return format(span.context.trace_id, "032x") + + +def _emit(rig, *, metadata=None, headers=None): + """``log_event_on_langfuse`` is the entry point that folds ``langfuse_*`` headers into metadata. + + Both the trace-level and the observation fields are read back off the span litellm exported. + """ + logger, exporter = rig + exporter.clear() + + now = datetime.datetime.now() + response_obj = litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "the-output"}}]) logger.log_event_on_langfuse( kwargs={ "call_type": "completion", @@ -1329,10 +1549,14 @@ def _emit(logger: LangFuseLogger, *, metadata=None, headers=None): start_time=now, end_time=now, ) - return ( - logger.Langfuse.trace.call_args.kwargs, - logger.Langfuse.trace.return_value.generation.call_args.kwargs, - ) + prefix = "langfuse.observation." + span = _exported_span(logger, exporter) + generation_params = { + key[len(prefix) :]: value + for key, value in (span.attributes or {}).items() + if key.startswith(prefix) and not key.startswith(prefix + "metadata.") + } + return _trace_params(span), generation_params, span @pytest.mark.parametrize("level", ["DEFAULT", "ERROR"]) @@ -1458,8 +1682,9 @@ def test_session_header_trace_provenance(headers, metadata, expected_id, level): redact_credential_headers, ) - logger: Final = _steering_logger() + logger, exporter = _steering_logger() for turn in range(2): + exporter.clear() call_id = f"call-{turn}" request_headers = Headers(headers) data = LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( @@ -1489,17 +1714,19 @@ def test_session_header_trace_provenance(headers, metadata, expected_id, level): level=level, status_message="provider error" if level == "ERROR" else None, ) - trace_params = logger.Langfuse.trace.call_args.kwargs - assert trace_params["id"] == (call_id if expected_id == "call" else expected_id) - assert result["trace_id"] == trace_params["id"] + span = _exported_span(logger, exporter) + assert _span_trace_id(span) == resolve_trace_id(call_id if expected_id == "call" else expected_id) + assert result["trace_id"] == _span_trace_id(span) if expected_id != "existing-trace": - assert trace_params["session_id"] == headers.get("langfuse_session_id", original_metadata.get("session_id")) + assert span.attributes.get("session.id") == headers.get( + "langfuse_session_id", original_metadata.get("session_id") + ) steering = {key[len("langfuse_") :]: value for key, value in headers.items() if key.startswith("langfuse_")} assert data["metadata"] == {**original_metadata, **steering} def test_session_header_trace_without_call_id_keeps_session_alias(): - logger: Final = _steering_logger() + logger, exporter = _steering_logger() now: Final = datetime.datetime.now() result: Final = logger.log_event_on_langfuse( @@ -1518,8 +1745,8 @@ def test_session_header_trace_without_call_id_keeps_session_alias(): end_time=now, ) - assert logger.Langfuse.trace.call_args.kwargs["id"] == "session-7125" - assert result["trace_id"] == "session-7125" + assert _span_trace_id(_exported_span(logger, exporter)) == resolve_trace_id("session-7125") + assert result["trace_id"] == resolve_trace_id("session-7125") def test_every_proxy_session_header_shape_is_classified_as_a_session_alias(): @@ -1553,7 +1780,7 @@ def test_every_proxy_session_header_shape_is_classified_as_a_session_alias(): ) def test_sdk_caller_without_request_headers_keeps_its_trace(proxy_server_request): """A direct SDK caller has no request headers, so a session-shaped trace id stays the caller's.""" - logger: Final = _steering_logger() + logger, exporter = _steering_logger() now: Final = datetime.datetime.now() result: Final = logger.log_event_on_langfuse( @@ -1572,8 +1799,8 @@ def test_sdk_caller_without_request_headers_keeps_its_trace(proxy_server_request end_time=now, ) - assert logger.Langfuse.trace.call_args.kwargs["id"] == "session-7125" - assert result["trace_id"] == "session-7125" + assert _span_trace_id(_exported_span(logger, exporter)) == resolve_trace_id("session-7125") + assert result["trace_id"] == resolve_trace_id("session-7125") def test_session_header_classifier_survives_non_string_header_keys(): @@ -1587,38 +1814,38 @@ def test_session_header_classifier_survives_non_string_header_keys(): def test_mask_input_header_false_keeps_the_prompt(): - logger = _steering_logger() + rig = _steering_logger() - trace_params, generation_params = _emit(logger, headers={"langfuse_mask_input": "false"}) + trace_params, generation_params, _ = _emit(rig, headers={"langfuse_mask_input": "false"}) - assert trace_params["input"] == {"messages": [{"role": "user", "content": "the-input"}]} - assert generation_params["input"] == {"messages": [{"role": "user", "content": "the-input"}]} + assert "input" not in trace_params + assert json.loads(generation_params["input"]) == {"messages": [{"role": "user", "content": "the-input"}]} def test_mask_input_header_true_redacts_the_prompt(): - logger = _steering_logger() + rig = _steering_logger() - trace_params, generation_params = _emit(logger, headers={"langfuse_mask_input": "true"}) + trace_params, generation_params, _ = _emit(rig, headers={"langfuse_mask_input": "true"}) - assert trace_params["input"] == _LANGFUSE_REDACTED + assert "input" not in trace_params assert generation_params["input"] == _LANGFUSE_REDACTED def test_mask_output_header_false_keeps_the_completion(): - logger = _steering_logger() + rig = _steering_logger() - trace_params, generation_params = _emit(logger, headers={"langfuse_mask_output": "false"}) + trace_params, generation_params, _ = _emit(rig, headers={"langfuse_mask_output": "false"}) - assert trace_params["output"] != _LANGFUSE_REDACTED - assert generation_params["output"] != _LANGFUSE_REDACTED + assert "output" not in trace_params + assert "the-output" in generation_params["output"] def test_mask_output_header_true_redacts_the_completion(): - logger = _steering_logger() + rig = _steering_logger() - trace_params, generation_params = _emit(logger, headers={"langfuse_mask_output": "true"}) + trace_params, generation_params, _ = _emit(rig, headers={"langfuse_mask_output": "true"}) - assert trace_params["output"] == _LANGFUSE_REDACTED + assert "output" not in trace_params assert generation_params["output"] == _LANGFUSE_REDACTED @@ -1632,30 +1859,31 @@ def test_mask_output_header_true_redacts_the_completion(): ], ) def test_mask_input_from_the_request_body_is_unchanged(mask_input, expect_redacted): - logger = _steering_logger() + rig = _steering_logger() - trace_params, _ = _emit(logger, metadata={"mask_input": mask_input}) + _, generation_params, _ = _emit(rig, metadata={"mask_input": mask_input}) - assert (trace_params["input"] == _LANGFUSE_REDACTED) is expect_redacted + assert (generation_params["input"] == _LANGFUSE_REDACTED) is expect_redacted @pytest.mark.parametrize("flag", [True, "true"]) -def test_update_trace_keys_header_applies_every_key_when_enabled(flag): - logger = _steering_logger() +def test_update_trace_keys_header_applies_every_key_when_enabled(flag, monkeypatch): + rig = _steering_logger() - with patch.object(litellm, "langfuse_enable_update_trace_keys", flag): - trace_params, _ = _emit( - logger, - headers={ - "langfuse_existing_trace_id": "trace-1", - "langfuse_update_trace_keys": "trace_release, trace_tail", - "langfuse_trace_release": "v1.2.3", - "langfuse_trace_tail": "last", - }, - ) + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", flag) + trace_params, _, span = _emit( + rig, + headers={ + "langfuse_existing_trace_id": "trace-1", + "langfuse_update_trace_keys": "trace_release, trace_tail", + "langfuse_trace_release": "v1.2.3", + "langfuse_trace_tail": "last", + }, + ) assert trace_params["release"] == "v1.2.3" - assert trace_params["tail"] == "last" + assert span.attributes["langfuse.release"] == "v1.2.3" + assert not [key for key in span.attributes if key.endswith("tail")] def test_update_trace_keys_is_off_by_default(): @@ -1664,10 +1892,10 @@ def test_update_trace_keys_is_off_by_default(): user_api_key_auth and have the resolved auth object, including team callback credentials, serialized onto the trace. It stays inert until an operator opts in. """ - logger = _steering_logger() + rig = _steering_logger() - trace_params, _ = _emit( - logger, + trace_params, _, span = _emit( + rig, metadata={ "existing_trace_id": "trace-1", "update_trace_keys": ["user_api_key_auth", "trace_release"], @@ -1678,41 +1906,185 @@ def test_update_trace_keys_is_off_by_default(): assert "user_api_key_auth" not in trace_params assert "release" not in trace_params - assert "sk-canary" not in json.dumps(trace_params, default=repr) + assert "sk-canary" not in json.dumps(dict(span.attributes or {}), default=repr) -def test_update_trace_keys_input_and_output_are_gated_too(): - logger = _steering_logger() +def test_update_trace_keys_input_and_output_are_gated_too(monkeypatch): + rig = _steering_logger() - off, _ = _emit(logger, metadata={"existing_trace_id": "trace-1", "update_trace_keys": ["input", "output"]}) - with patch.object(litellm, "langfuse_enable_update_trace_keys", True): - on, _ = _emit(logger, metadata={"existing_trace_id": "trace-1", "update_trace_keys": ["input", "output"]}) + off, _, _ = _emit(rig, metadata={"existing_trace_id": "trace-1", "update_trace_keys": ["input", "output"]}) + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", True) + on, _, _ = _emit(rig, metadata={"existing_trace_id": "trace-1", "update_trace_keys": ["input", "output"]}) assert "input" not in off and "output" not in off assert "input" in on and "output" in on -def test_update_trace_keys_from_the_request_body_list_applies_when_enabled(): - logger = _steering_logger() +def test_update_trace_keys_input_output_reach_the_trace_even_under_a_parent(monkeypatch): + """With a real parent the generation is not the trace root, so trace-level + I/O must be stamped explicitly; v2 updated the trace object directly.""" + rig = _steering_logger() - with patch.object(litellm, "langfuse_enable_update_trace_keys", True): - trace_params, _ = _emit( - logger, - metadata={ - "existing_trace_id": "trace-1", - "update_trace_keys": ["trace_release"], - "trace_release": "v1.2.3", - }, - ) + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", True) + _, _, span = _emit( + rig, + metadata={ + "existing_trace_id": "trace-1", + "parent_observation_id": "b" * 16, + "update_trace_keys": ["input", "output"], + }, + ) + + assert "the-input" in str(span.attributes["langfuse.trace.input"]) + assert "the-output" in str(span.attributes["langfuse.trace.output"]) + + +def test_a_fresh_trace_under_a_callers_parent_still_carries_its_own_input_and_output(): + """Langfuse copies I/O onto a trace only from its root observation; a caller's ``parent_observation_id`` + makes the generation a child, so the trace-level fields v2 set on ``trace(...)`` must be stamped.""" + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"parent_observation_id": "0123456789abcdef"}) + + assert span.parent is not None + assert "the-input" in str(span.attributes["langfuse.trace.input"]) + assert "the-output" in str(span.attributes["langfuse.trace.output"]) + + +def test_a_failed_call_under_a_callers_parent_stamps_the_error_as_the_trace_output(): + """The ERROR branch used to write a trace-level ``status_message``, a field the v4 trace schema does not + have, and skip ``output``; the generation's parent is the caller's, so nothing else fills the trace.""" + logger, exporter = _steering_logger() + now = datetime.datetime.now() + + logger.log_event_on_langfuse( + kwargs={ + "call_type": "completion", + "litellm_params": {"metadata": {"parent_observation_id": "0123456789abcdef"}}, + "messages": [{"role": "user", "content": "the-input"}], + "optional_params": {}, + }, + response_obj=None, + start_time=now, + end_time=now, + level="ERROR", + status_message="provider said no", + ) + span = _exported_span(logger, exporter) + + assert span.parent is not None + assert "provider said no" in str(span.attributes["langfuse.trace.output"]) + assert span.attributes["langfuse.observation.status_message"] == "provider said no" + + +def test_a_fresh_trace_root_leaves_the_duplicate_io_to_langfuse(): + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"trace_id": "a" * 32}) + + assert span.parent is None + assert "langfuse.trace.input" not in (span.attributes or {}) + assert "the-input" in str(span.attributes["langfuse.observation.input"]) + + +def test_existing_trace_id_appends_without_claiming_trace_root(): + """Langfuse copies a root observation's name and I/O onto the trace, so a + continuation that claimed root would rename the trace after every request; + v2 only ever touched the keys in ``update_trace_keys``.""" + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"existing_trace_id": "trace-1", "trace_name": "second-call"}) + + assert span.parent is not None + assert "langfuse.trace.name" not in (span.attributes or {}) + + +def test_a_fresh_trace_still_claims_root_so_its_generation_names_it(): + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"trace_id": "a" * 32, "trace_name": "first-call"}) + + assert span.parent is None + assert span.attributes["langfuse.trace.name"] == "first-call" + + +def test_trace_io_is_not_stamped_when_update_trace_keys_does_not_ask(monkeypatch): + rig = _steering_logger() + + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", True) + _, _, span = _emit( + rig, + metadata={ + "existing_trace_id": "trace-1", + "parent_observation_id": "b" * 16, + "update_trace_keys": ["trace_release"], + }, + ) + + assert "langfuse.trace.input" not in (span.attributes or {}) + assert "langfuse.trace.output" not in (span.attributes or {}) + + +def test_update_trace_keys_from_the_request_body_list_applies_when_enabled(monkeypatch): + rig = _steering_logger() + + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", True) + trace_params, _, span = _emit( + rig, + metadata={ + "existing_trace_id": "trace-1", + "update_trace_keys": ["trace_release"], + "trace_release": "v1.2.3", + }, + ) assert trace_params["release"] == "v1.2.3" + assert span.attributes["langfuse.release"] == "v1.2.3" + + +def test_update_trace_keys_trace_metadata_reaches_the_trace_and_stays_off_the_generation(monkeypatch): + rig = _steering_logger() + + monkeypatch.setattr(litellm, "langfuse_enable_update_trace_keys", True) + _, _, span = _emit( + rig, + metadata={ + "existing_trace_id": "trace-1", + "parent_observation_id": "b" * 16, + "update_trace_keys": ["trace_metadata"], + "trace_metadata": {"step": 2, "note": "x" * 300}, + }, + ) + + assert span.attributes["langfuse.trace.metadata.step"] == 2 + assert span.attributes["langfuse.trace.metadata.note"] == "x" * 300 + assert "langfuse.observation.metadata.step" not in span.attributes + + +def test_non_mapping_trace_metadata_does_not_lose_the_event(): + """A caller who passes ``trace_metadata`` as a string still gets a generation, and the string is not spread.""" + rig = _steering_logger() + + trace_params, generation_params, span = _emit(rig, metadata={"trace_metadata": "just-a-note"}) + + assert json.loads(generation_params["output"])["content"] == "the-output" + assert trace_params["name"] == "litellm-completion" + assert not any(key.startswith("langfuse.trace.metadata.") for key in span.attributes or {}) + + +def test_trace_metadata_is_not_propagated_when_absent(): + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"trace_name": "plain"}) + + assert not any(key.startswith("langfuse.trace.metadata.") for key in span.attributes or {}) def test_update_trace_keys_matches_whole_keys_not_substrings(): - logger = _steering_logger() + rig = _steering_logger() - trace_params, _ = _emit( - logger, + trace_params, _, _ = _emit( + rig, headers={"langfuse_existing_trace_id": "trace-1", "langfuse_update_trace_keys": "my_input"}, ) @@ -1720,25 +2092,12 @@ def test_update_trace_keys_matches_whole_keys_not_substrings(): def test_langfuse_environment_is_coerced_and_validated(monkeypatch): - monkeypatch.setenv("LANGFUSE_MOCK", "false") monkeypatch.delenv("LANGFUSE_TRACING_ENVIRONMENT", raising=False) - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuse): - logger = LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - langfuse_environment=123, # non-string: must coerce, not crash - ) + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-env", langfuse_environment=123) assert logger.langfuse_environment == "123" with pytest.raises(ValueError, match="langfuse_environment"): - LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - langfuse_environment="Production", - ) + _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-env", langfuse_environment="Production") def test_langfuse_empty_environment_falls_back_and_is_not_dynamic(monkeypatch): @@ -1748,15 +2107,7 @@ def test_langfuse_empty_environment_falls_back_and_is_not_dynamic(monkeypatch): monkeypatch.setenv("LANGFUSE_TRACING_ENVIRONMENT", "production") # '' falls back to the deployment env var at init - monkeypatch.setenv("LANGFUSE_MOCK", "false") - monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) - with patch("langfuse.Langfuse", _RecordingLangfuse): - logger = LangFuseLogger( - langfuse_public_key="pk-env", - langfuse_secret="sk-env", - langfuse_host="https://test.langfuse.com", - langfuse_environment="", - ) + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-env", langfuse_environment="") assert logger.langfuse_environment == "production" # env-only params that add nothing do not select a dynamic logger @@ -1799,3 +2150,322 @@ def test_langfuse_deployment_environment_fallback_never_raises(monkeypatch, env_ langfuse_host="https://test.langfuse.com", ) assert logger.langfuse_environment == expected + + +def test_continued_trace_keeps_the_generation_version(): + """v2 set ``version`` on the generation even when the trace was not being updated.""" + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"existing_trace_id": "b" * 32, "version": "gen-7"}) + + assert span.attributes["langfuse.version"] == "gen-7" + + +def test_new_trace_version_takes_precedence_over_the_generation_version(): + """v4 has one ``langfuse.version`` per span, so unlike v2's separate trace and generation fields only one + value can survive; ``trace_version`` wins, matching the v4 SDK, whose propagated attributes overwrite a span's own.""" + rig = _steering_logger() + + captured_trace_params, _, span = _emit(rig, metadata={"trace_version": "trace-1", "version": "gen-7"}) + + assert captured_trace_params["version"] == "trace-1" + assert span.attributes["langfuse.version"] == "trace-1" + + +def test_log_event_returns_the_v2_dict_shape_for_the_alerting_trace_id_cache(): + """litellm_logging only caches the langfuse trace id off a dict with a ``trace_id`` key. + + Slack alerting builds its trace URL from that cache, so a different return + shape silently breaks alert links. + """ + rig = _steering_logger() + logger, _ = rig + + returned = logger.log_event_on_langfuse( + kwargs={ + "call_type": "completion", + "litellm_params": {"metadata": {"trace_id": "c" * 32}}, + "messages": [{"role": "user", "content": "the-input"}], + "optional_params": {}, + }, + response_obj=litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "the-output"}}]), + start_time=datetime.datetime.now(), + end_time=datetime.datetime.now(), + ) + + assert isinstance(returned, dict) + assert returned["trace_id"] == "c" * 32 + assert returned["generation_id"] + + +def test_parse_langfuse_debug_only_enables_on_true_strings(): + """v4 treats any truthy value as debug=on, so the raw env string "false" would enable debug.""" + assert langfuse_module.parse_langfuse_debug("true") is True + assert langfuse_module.parse_langfuse_debug("True") is True + assert langfuse_module.parse_langfuse_debug("1") is True + assert langfuse_module.parse_langfuse_debug("false") is False + assert langfuse_module.parse_langfuse_debug("False") is False + assert langfuse_module.parse_langfuse_debug("") is False + assert langfuse_module.parse_langfuse_debug(None) is False + + +@pytest.mark.parametrize( + ("raw", "expected"), + [(None, 1), ("", 1), ("3", 3), ("0", 1), ("-5", 1), ("abc", 1)], + ids=["unset", "empty", "valid", "zero", "negative", "text"], +) +def test_flush_interval_env_falls_back_instead_of_failing_the_first_request(monkeypatch, raw, expected, caplog): + """The batch scheduler rejects a non-positive delay; v2's consumer thread accepted 0, so the value must + not raise out of the lazily built logger and take Langfuse logging down for the worker.""" + if raw is None: + monkeypatch.delenv("LANGFUSE_FLUSH_INTERVAL", raising=False) + else: + monkeypatch.setenv("LANGFUSE_FLUSH_INTERVAL", raw) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert LangFuseLogger._get_langfuse_flush_interval(1) == expected # pyright: ignore[reportPrivateUsage] # the parser under test + assert ("LANGFUSE_FLUSH_INTERVAL" in caplog.text) is (raw in ("0", "-5", "abc")) + + +def test_zero_flush_interval_still_builds_a_working_export_channel(monkeypatch): + receiver = _OtlpReceiver() + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-flush-zero-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-flush-zero-test") + monkeypatch.delenv("LANGFUSE_MOCK", raising=False) + monkeypatch.setenv("LANGFUSE_FLUSH_INTERVAL", "0") + monkeypatch.setattr(litellm, "initialized_langfuse_clients", litellm.initialized_langfuse_clients) + + try: + logger = LangFuseLogger(langfuse_host=receiver.url) + _log_one_completion(logger) + finally: + receiver.close() + + assert receiver.received == ["/api/public/otel/v1/traces"] + + +def test_langfuse_debug_env_string_false_stays_off(monkeypatch): + """LANGFUSE_DEBUG=false must not reach the v4 client as a truthy string. + + The v4 client does ``if debug:`` and then mutates root logging via + ``logging.basicConfig``, so the unparsed string "false" turns debug ON. + """ + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-debug-parse-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-debug-parse-test") + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_DEBUG", "false") + monkeypatch.setattr(litellm, "initialized_langfuse_clients", litellm.initialized_langfuse_clients) + + assert LangFuseLogger().langfuse_debug is False + + +def test_langfuse_debug_env_true_turns_on_the_langfuse_logger(monkeypatch): + """``LANGFUSE_DEBUG=true`` reached the v2 client as ``debug=`` and switched the SDK's logger to DEBUG; + a parsed flag that nothing reads would make the variable a silent no-op.""" + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-debug-wire-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-debug-wire-test") + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("LANGFUSE_DEBUG", "true") + monkeypatch.setattr(litellm, "initialized_langfuse_clients", litellm.initialized_langfuse_clients) + langfuse_logger = logging.getLogger("langfuse") + level_before = langfuse_logger.level + langfuse_logger.setLevel(logging.WARNING) + try: + assert LangFuseLogger().langfuse_debug is True + assert langfuse_logger.level == logging.DEBUG + finally: + langfuse_logger.setLevel(level_before) + + +def test_explicit_langfuse_host_beats_the_v4_base_url_env(monkeypatch): + """Per-key/per-team ``langfuse_host`` must win over LANGFUSE_BASE_URL. + + v4 resolves ``base_url or $LANGFUSE_BASE_URL or host``, so a stray env var + could silently redirect every tenant's traces to one server. The proof is a + real round trip: the observation lands on the configured host. + """ + receiver = _OtlpReceiver() + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-base-url-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-base-url-test") + monkeypatch.delenv("LANGFUSE_MOCK", raising=False) + monkeypatch.setenv("LANGFUSE_BASE_URL", "http://127.0.0.1:1") + monkeypatch.setenv("LANGFUSE_FLUSH_INTERVAL", "1") + monkeypatch.setattr(litellm, "initialized_langfuse_clients", litellm.initialized_langfuse_clients) + + try: + logger = LangFuseLogger(langfuse_host=receiver.url) + _log_one_completion(logger) + finally: + receiver.close() + + assert logger.langfuse_host == receiver.url + assert receiver.received == ["/api/public/otel/v1/traces"] + + +def test_resolve_credentials_falls_back_to_langfuse_base_url(monkeypatch): + """v4's canonical env var works when LANGFUSE_HOST is unset, but never beats it.""" + monkeypatch.setenv("LANGFUSE_BASE_URL", "https://from-base-url.example") + monkeypatch.delenv("LANGFUSE_HOST", raising=False) + + _, _, host = langfuse_module.resolve_langfuse_credentials() + assert host == "https://from-base-url.example" + + monkeypatch.setenv("LANGFUSE_HOST", "https://from-host.example") + _, _, host = langfuse_module.resolve_langfuse_credentials() + assert host == "https://from-host.example" + + _, _, host = langfuse_module.resolve_langfuse_credentials(langfuse_host="https://explicit.example") + assert host == "https://explicit.example" + + +def test_version_gate_rejects_v5_prereleases(): + """ "5.0.0rc1" sorts below "5", so a plain version comparison would admit it.""" + langfuse_module.raise_if_unsupported_langfuse_version("4.7") + with pytest.raises(ImportError): + langfuse_module.raise_if_unsupported_langfuse_version("5.0.0rc1") + with pytest.raises(ImportError): + langfuse_module.raise_if_unsupported_langfuse_version("5.0.0") + + +def test_old_sdk_fails_with_the_upgrade_message_before_the_otel_module_is_imported(monkeypatch): + """On a v2 install `langfuse_sdk` itself fails to import, so the version gate must run first + or the caller is told the package is missing when it only needs upgrading.""" + import sys + + monkeypatch.setattr(langfuse_module, "installed_langfuse_version", lambda: "2.59.7") + monkeypatch.setitem(sys.modules, "litellm.integrations.langfuse.langfuse_sdk", None) + + with pytest.raises(ImportError) as raised: + _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-old-sdk") + + assert "2.59.7" in str(raised.value) + assert "langfuse_otel" in str(raised.value) + assert "not installed" not in str(raised.value) + + +def test_missing_sdk_is_reported_as_not_installed(monkeypatch): + from importlib.metadata import PackageNotFoundError + + def not_installed() -> str: + raise PackageNotFoundError("langfuse") + + monkeypatch.setattr(langfuse_module, "installed_langfuse_version", not_installed) + + with pytest.raises(Exception, match="Langfuse not installed"): + _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-no-sdk") + + +@pytest.mark.parametrize("raw", ["abc", "2.5", ""], ids=["text", "fraction", "empty"]) +def test_prompt_cache_ttl_typo_is_named_before_the_sdk_is_imported(monkeypatch, raw): + """The v4 SDK evaluates ``int(LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS)`` at import, so without this + gate every request failed with a bare ``invalid literal for int()`` that never named the variable.""" + import sys + + monkeypatch.setenv("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS", raw) + monkeypatch.setitem(sys.modules, "litellm.integrations.langfuse.langfuse_sdk", None) + + with pytest.raises(ValueError, match="LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS") as raised: + _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-ttl-typo") + + assert repr(raw) in str(raised.value) + + +@pytest.mark.parametrize("raw", ["5", " -3 ", "+0"], ids=["whole", "negative", "signed-zero"]) +def test_whole_second_prompt_cache_ttl_passes_the_gate(monkeypatch, raw): + monkeypatch.setenv("LANGFUSE_PROMPT_CACHE_DEFAULT_TTL_SECONDS", raw) + assert langfuse_module.raise_if_unusable_prompt_cache_ttl() is None + + +def test_stopped_logger_hands_its_export_channel_back(monkeypatch): + """`DynamicLoggingCache` calls `stop()` on expiry; the channel must be retired once every + logger that held it has stopped, or each credential rotation leaks a batch export thread.""" + from litellm.integrations.langfuse.langfuse_sdk import acquire_langfuse_tracing, release_langfuse_tracing + + logger = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-stop-releases") + + def acquire_same_credentials(): + return acquire_langfuse_tracing( + public_key="pk-stop-releases", + secret_key="sk-lit5228", + base_url=_UNREACHABLE_HOST, + environment=logger.langfuse_environment, + release=logger.langfuse_release, + flush_interval=logger.langfuse_flush_interval, + mock_mode=False, + ) + + logger.stop() + reacquired = acquire_same_credentials() + assert reacquired is logger.tracing, "the channel stays up while another logger still holds it" + + release_langfuse_tracing(reacquired, grace_seconds=0.0) + assert acquire_same_credentials() is not logger.tracing, "stop() did not give the logger's hold back" + + +def test_logger_that_fails_to_build_takes_no_slot_and_no_channel(monkeypatch): + """Each failed retry for the same dynamic credentials would otherwise eat a client slot and a + holder on the channel, so fixing the configuration could not bring Langfuse logging back.""" + from litellm.integrations.langfuse.langfuse_sdk import acquire_langfuse_tracing, release_langfuse_tracing + + monkeypatch.setenv("LANGFUSE_TIMEOUT", "5.5") + probe = _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-failed-build") + + def acquire_same_credentials(): + return acquire_langfuse_tracing( + public_key="pk-failed-build", + secret_key="sk-lit5228", + base_url=_UNREACHABLE_HOST, + environment=probe.langfuse_environment, + release=probe.langfuse_release, + flush_interval=probe.langfuse_flush_interval, + mock_mode=False, + ) + + monkeypatch.setenv("LANGFUSE_TIMEOUT", "not-a-number") + with pytest.raises(ValueError, match="not-a-number"): + _build_langfuse_logger(monkeypatch, langfuse_public_key="pk-failed-build") + assert litellm.initialized_langfuse_clients == 0 + + monkeypatch.setenv("LANGFUSE_TIMEOUT", "5.5") + release_langfuse_tracing(probe.tracing, grace_seconds=0.0) + assert acquire_same_credentials() is not probe.tracing, "the failed build left a holder on the channel" + + +def test_int_steering_values_reach_langfuse_as_strings(): + """Langfuse models user, session and version as strings; v2's pydantic coerced ints for the caller.""" + rig = _steering_logger() + + _, _, span = _emit(rig, metadata={"trace_user_id": 12345, "session_id": 67, "trace_version": 3}) + + assert span.attributes["user.id"] == "12345" + assert span.attributes["session.id"] == "67" + assert span.attributes["langfuse.version"] == "3" + + +def test_long_steering_values_are_neither_capped_nor_dropped(): + """v2 sent ids of any length; the SDK's 200 character rule belongs to baggage propagation, which litellm no longer uses.""" + rig = _steering_logger() + long_user: Final = "u" * 250 + + _, _, span = _emit(rig, metadata={"trace_user_id": long_user}) + + assert span.attributes["user.id"] == long_user + + +def test_returned_generation_id_names_the_exported_observation(): + """v4 derives observation ids from the OTel span, so a pre-computed id would name nothing.""" + logger, exporter = _steering_logger() + + returned = logger.log_event_on_langfuse( + kwargs={ + "call_type": "completion", + "litellm_params": {"metadata": {"trace_id": "d" * 32}}, + "messages": [{"role": "user", "content": "the-input"}], + "optional_params": {}, + }, + response_obj=litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "the-output"}}]), + start_time=datetime.datetime.now(), + end_time=datetime.datetime.now(), + ) + + span = _exported_span(logger, exporter) + assert returned["generation_id"] == format(span.context.span_id, "016x") 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/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 4322662cfcb..26124ac24de 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -3847,3 +3847,51 @@ def test_convert_to_anthropic_tool_invoke_keeps_paired_server_tool_use(): }, server_result, ] + + +def test_anthropic_messages_pt_keeps_system_role_after_user_turn(): + """Models flagged supports_mid_conversation_system accept role=system inside + messages; the converter must emit it as a system message with its text + blocks and cache_control intact instead of rejecting the role.""" + messages = [ + {"role": "user", "content": "First question"}, + { + "role": "system", + "content": [{"type": "text", "text": "Answer in one word.", "cache_control": {"type": "ephemeral"}}], + }, + {"role": "assistant", "content": "Yes"}, + {"role": "user", "content": "Second question"}, + ] + + result = anthropic_messages_pt(messages=messages, model="claude-opus-4-8", llm_provider="anthropic") + + assert [m["role"] for m in result] == ["user", "system", "assistant", "user"] + assert result[1] == { + "role": "system", + "content": [{"type": "text", "text": "Answer in one word.", "cache_control": {"type": "ephemeral"}}], + } + + +def test_anthropic_messages_pt_system_string_content_becomes_text_block(): + messages = [ + {"role": "user", "content": "First question"}, + {"role": "system", "content": "Answer in one word."}, + ] + + result = anthropic_messages_pt(messages=messages, model="claude-opus-4-8", llm_provider="anthropic") + + assert result[1] == {"role": "system", "content": [{"type": "text", "text": "Answer in one word."}]} + + +def test_anthropic_messages_pt_drops_a_system_message_with_no_text(): + """Anthropic rejects empty text blocks, so a text-less system message must + vanish rather than reach the wire as an empty system turn.""" + messages = [ + {"role": "user", "content": "First question"}, + {"role": "system", "content": ""}, + {"role": "assistant", "content": "Yes"}, + ] + + result = anthropic_messages_pt(messages=messages, model="claude-opus-4-8", llm_provider="anthropic") + + assert [m["role"] for m in result] == ["user", "assistant"] diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py new file mode 100644 index 00000000000..d1e23a17747 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py @@ -0,0 +1,392 @@ +"""Placement policy for mid-conversation ``role: "system"`` messages on the chat path. + +The provider-facing behaviour is covered through ``transform_request`` in the +Anthropic, Vertex, Azure AI and Bedrock Invoke transformation tests; these pin +the pure placement rules on the OpenAI-format message list. +""" + +import pytest + +import litellm +from litellm.litellm_core_utils.prompt_templates.common_utils import encrypted_reasoning_signature +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import ( + CONVERTED_SYSTEM_NOTE, + place_mid_conversation_system, + split_leading_system_run, +) + + +def _roles(messages: object) -> list[str]: + return [m["role"] if isinstance(m, dict) else m.role for m in messages] + + +def _texts(message: dict) -> list[str]: + return [block["text"] for block in message["content"]] + + +SENDS_NOTHING = pytest.mark.parametrize( + "empty_content", + [[], None, [{"type": "input_audio", "input_audio": {"data": "AAAA", "format": "wav"}}]], + ids=["empty-list", "none", "unsupported-part-only"], +) + + +def test_split_leading_system_run_keeps_later_system_messages_in_the_conversation(): + messages = [ + {"role": "system", "content": "one"}, + {"role": "system", "content": "two"}, + {"role": "user", "content": "q"}, + {"role": "system", "content": "reminder"}, + ] + + leading, later = split_leading_system_run(messages) + + assert [m["content"] for m in leading] == ["one", "two"] + assert _roles(later) == ["user", "system"] + + +def test_flagged_placement_moves_a_system_run_after_the_user_turn_that_follows_it(): + messages = [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1"}, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": "q2"}, + {"role": "assistant", "content": "a2"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=True) + + assert _roles(placed) == ["user", "assistant", "user", "system", "assistant"] + + +def test_flagged_placement_pushes_a_system_between_two_user_turns_after_both(): + """Two user turns collapse into one on the wire, and a system message must + be followed by an assistant turn or nothing.""" + messages = [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": "q2"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=True) + + assert _roles(placed) == ["user", "user", "system"] + + +def test_flagged_placement_keeps_a_system_after_tool_results(): + messages = [ + {"role": "user", "content": "q1"}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "c1", "content": "r"}, + {"role": "system", "content": "reminder"}, + {"role": "assistant", "content": "a2"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=True) + + assert _roles(placed) == ["user", "assistant", "tool", "system", "assistant"] + + +def test_flagged_placement_drops_a_system_message_with_no_text(): + messages = [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": ""}, + {"role": "assistant", "content": "a1"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=True) + + assert _roles(placed) == ["user", "assistant"] + + +def test_placement_reads_roles_off_pydantic_messages_in_the_history(): + """Callers routinely append the previous ``litellm.Message`` object straight + into the history; placement must read its role without assuming a dict and + hand the object through untouched.""" + assistant = litellm.Message(role="assistant", content="a1") + messages = [ + {"role": "user", "content": "q1"}, + assistant, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": "q2"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=True) + + assert _roles(placed) == ["user", "assistant", "user", "system"] + assert placed[1] is assistant + + +def test_unflagged_conversion_keeps_the_client_order_when_no_tool_result_follows(): + messages = [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1"}, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": "q2"}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=False) + + assert _roles(placed) == ["user", "assistant", "user", "user"] + assert _texts(placed[2]) == [CONVERTED_SYSTEM_NOTE, "reminder"] + + +@pytest.mark.parametrize( + "cache_control, expected", + [ + ({"type": "ephemeral", "ttl": "1h"}, {"type": "ephemeral", "ttl": "1h"}), + ({"type": "ephemeral", "ttl": "5m"}, {"type": "ephemeral", "ttl": "5m"}), + ({"type": "ephemeral", "ttl": "2h"}, {"type": "ephemeral"}), + ], +) +def test_unflagged_conversion_rebuilds_cache_control_on_the_converted_block(cache_control, expected): + """Only the shapes Anthropic accepts survive: ephemeral with a 5m or 1h ttl, or no ttl.""" + messages = [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder", "cache_control": cache_control}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=False) + + assert placed[1]["content"][1] == {"type": "text", "text": "reminder", "cache_control": expected} + + +def test_unflagged_conversion_drops_a_cache_control_that_is_not_ephemeral(): + messages = [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder", "cache_control": {"type": "persistent"}}, + ] + + placed = place_mid_conversation_system(messages, supports_mid_conversation_system=False) + + assert placed[1]["content"][1] == {"type": "text", "text": "reminder"} + + +def test_placement_is_a_no_op_without_later_system_messages(): + messages = [{"role": "user", "content": "q1"}, {"role": "assistant", "content": "a1"}] + + assert place_mid_conversation_system(messages, supports_mid_conversation_system=False) == tuple(messages) + assert place_mid_conversation_system(messages, supports_mid_conversation_system=True) == tuple(messages) + + +def test_flagged_placement_converts_a_run_followed_by_an_assistant_turn_in_place(): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1"}, + {"role": "system", "content": "reminder"}, + {"role": "assistant", "content": "a2"}, + {"role": "user", "content": "q2"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "assistant", "user", "assistant", "user"] + assert _texts(placed[2]) == [CONVERTED_SYSTEM_NOTE, "reminder"] + + +def test_flagged_placement_of_an_earlier_run_does_not_move_when_later_turns_are_appended(): + turn_n = [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1"}, + {"role": "system", "content": "reminder"}, + ] + turn_n_plus_one = [ + *turn_n, + {"role": "assistant", "content": "a2"}, + {"role": "user", "content": "q2"}, + ] + + placed_n = place_mid_conversation_system(turn_n, supports_mid_conversation_system=True) + placed_n_plus_one = place_mid_conversation_system(turn_n_plus_one, supports_mid_conversation_system=True) + + assert placed_n_plus_one[: len(placed_n)] == placed_n + assert _roles(placed_n_plus_one) == ["user", "assistant", "user", "assistant", "user"] + + +@SENDS_NOTHING +def test_flagged_placement_converts_a_run_whose_preceding_user_turn_sends_nothing(empty_content): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": empty_content}, + {"role": "system", "content": "reminder"}, + {"role": "assistant", "content": "a1"}, + {"role": "user", "content": "q2"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "user", "assistant", "user"] + assert _texts(placed[1]) == [CONVERTED_SYSTEM_NOTE, "reminder"] + + +@SENDS_NOTHING +def test_flagged_placement_converts_a_run_whose_following_user_turn_sends_nothing(empty_content): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "assistant", "content": "a1"}, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": empty_content}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "assistant", "user", "user"] + assert _texts(placed[2]) == [CONVERTED_SYSTEM_NOTE, "reminder"] + + +def test_flagged_placement_keeps_a_system_behind_a_user_turn_merged_with_an_empty_one(): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "user", "content": []}, + {"role": "system", "content": "reminder"}, + {"role": "assistant", "content": "a1"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "user", "system", "assistant"] + + +EMPTY_ASSISTANT = pytest.mark.parametrize( + "empty_assistant", + [ + {"role": "assistant", "content": None}, + {"role": "assistant", "content": []}, + {"role": "assistant", "content": [{"type": "thinking", "thinking": "unsigned"}]}, + { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "bridged", "signature": encrypted_reasoning_signature("abc")}], + }, + { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "redacted_thinking", "data": encrypted_reasoning_signature("abc")}], + }, + { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "unsigned"}], + "thinking_blocks": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, + { + "role": "assistant", + "content": [{"type": "redacted_thinking", "data": "x"}], + "thinking_blocks": [{"type": "redacted_thinking", "data": "x"}], + }, + litellm.Message(role="assistant", content=None), + ], + ids=[ + "none", + "empty-list", + "unsigned-thinking-part", + "encrypted-thinking-part", + "encrypted-redacted-thinking-block", + "unsigned-inline-part-hides-signed-block", + "inline-redacted-part-hides-redacted-block", + "pydantic-none", + ], +) + + +@EMPTY_ASSISTANT +def test_flagged_placement_converts_a_run_when_the_assistant_turn_after_its_anchor_sends_nothing(empty_assistant): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder"}, + empty_assistant, + {"role": "user", "content": "q2"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "user", "assistant", "user"] + assert _texts(placed[1]) == [CONVERTED_SYSTEM_NOTE, "reminder"] + assert placed[2] is empty_assistant + + +@EMPTY_ASSISTANT +def test_flagged_placement_keeps_a_system_whose_empty_assistant_follower_ends_the_array(empty_assistant): + placed = place_mid_conversation_system( + [{"role": "user", "content": "q1"}, {"role": "system", "content": "reminder"}, empty_assistant], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "system", "assistant"] + + +@pytest.mark.parametrize( + "assistant_turn", + [ + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "toolu_1", "type": "function", "function": {"name": "f", "arguments": "{}"}}], + }, + { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, + {"role": "assistant", "content": None, "thinking_blocks": [{"type": "redacted_thinking", "data": "x"}]}, + {"role": "assistant", "content": [{"type": "thinking", "thinking": "hm", "signature": "s"}]}, + {"role": "assistant", "content": None, "function_call": {"name": "f", "arguments": "{}"}}, + litellm.Message( + role="assistant", + content="", + tool_calls=[{"id": "toolu_1", "type": "function", "function": {"name": "f", "arguments": "{}"}}], + ), + ], + ids=[ + "tool-calls", + "signed-thinking-block", + "redacted-thinking-block", + "signed-thinking-part", + "function-call", + "pydantic-tool-calls", + ], +) +def test_flagged_placement_keeps_a_system_before_an_assistant_turn_that_renders_without_text(assistant_turn): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder"}, + assistant_turn, + {"role": "user", "content": "q2"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "system", "assistant", "user"] + + +@pytest.mark.parametrize( + "padded_assistant", + [ + {"role": "assistant", "content": ""}, + {"role": "assistant", "content": " "}, + {"role": "assistant", "content": [{"type": "text", "text": ""}]}, + litellm.Message(role="assistant", content=""), + ], + ids=["empty-string", "whitespace-string", "empty-text-part", "pydantic-empty-string"], +) +def test_flagged_placement_keeps_a_system_before_an_assistant_turn_whose_empty_text_the_converter_pads( + padded_assistant, +): + placed = place_mid_conversation_system( + [ + {"role": "user", "content": "q1"}, + {"role": "system", "content": "reminder"}, + padded_assistant, + {"role": "user", "content": "q2"}, + ], + supports_mid_conversation_system=True, + ) + + assert _roles(placed) == ["user", "system", "assistant", "user"] diff --git a/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py b/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py index 5ed9dca68fd..3ab710d8023 100644 --- a/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py +++ b/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py @@ -24,10 +24,7 @@ class TestLangfuseInMemoryCache: # Create a mock LangFuseLogger class class MockLangFuseLogger: - def __init__(self): - self.Langfuse = MagicMock() - self.Langfuse.flush = MagicMock() - self.Langfuse.shutdown = MagicMock() + pass mock_logger = MockLangFuseLogger() @@ -50,29 +47,72 @@ class TestLangfuseInMemoryCache: assert litellm.initialized_langfuse_clients == initial_count - 1 @patch("litellm.initialized_langfuse_clients", 3) - def test_langfuse_client_shutdown_called_on_eviction(self): - """Test that langfuse client shutdown is called to close the thread.""" + def test_evicted_logger_releases_its_hold_on_the_shared_export_channel(self): + """Export channels are shared per credential set: eviction gives this logger's hold back + while a sibling logger keeps exporting, and the channel is retired once the last hold goes.""" + from litellm.integrations.langfuse.langfuse import LangFuseLogger + from litellm.integrations.langfuse.langfuse_sdk import acquire_langfuse_tracing, release_langfuse_tracing - # Create a mock LangFuseLogger class - class MockLangFuseLogger: - def __init__(self): - self.Langfuse = MagicMock() - self.Langfuse.flush = MagicMock() - self.Langfuse.shutdown = MagicMock() + def acquire(): + return acquire_langfuse_tracing( + public_key="pk-eviction-test", + secret_key="sk", + base_url="http://127.0.0.1:1", + environment=None, + release=None, + flush_interval=1.0, + mock_mode=True, + ) - mock_logger = MockLangFuseLogger() + logger = LangFuseLogger.__new__(LangFuseLogger) + logger.api_client = MagicMock() + logger.api_client.get_prompt.return_value = "prompt-after-eviction" + logger.tracing = acquire() + sibling = acquire() + self.cache.cache_dict["test_key"] = logger + self.cache.ttl_dict["test_key"] = time.time() + 100 - # Patch the LangFuseLogger import to return our mock class - with patch( - "litellm.integrations.langfuse.langfuse.LangFuseLogger", MockLangFuseLogger - ): - # Add the mock logger to cache - self.cache.cache_dict["test_key"] = mock_logger - self.cache.ttl_dict["test_key"] = time.time() + 100 + self.cache._remove_key("test_key") - # Remove the key (this should trigger cleanup) - self.cache._remove_key("test_key") + assert litellm.initialized_langfuse_clients == 2 + assert logger.api_client.get_prompt("greeting") == "prompt-after-eviction" + with sibling.tracer.start_as_current_span("still-open"): + pass + assert sibling.flush(1000) is True - # Verify flush and shutdown were called - mock_logger.Langfuse.flush.assert_called_once() - mock_logger.Langfuse.shutdown.assert_called_once() + release_langfuse_tracing(sibling, grace_seconds=0.0) + assert acquire() is not logger.tracing, "eviction did not release the evicted logger's hold" + + @patch("litellm.initialized_langfuse_clients", 3) + def test_second_evictor_of_the_same_entry_releases_nothing(self): + """Two callers can expire the same entry at once (a request thread and the reaper). Only the one that + claims the entry may give its slot and channel hold back, or a sibling logger loses its channel.""" + from litellm.integrations.langfuse.langfuse import LangFuseLogger + from litellm.integrations.langfuse.langfuse_sdk import acquire_langfuse_tracing, release_langfuse_tracing + + def acquire(): + return acquire_langfuse_tracing( + public_key="pk-double-eviction-test", + secret_key="sk", + base_url="http://127.0.0.1:1", + environment=None, + release=None, + flush_interval=1.0, + mock_mode=True, + ) + + logger = LangFuseLogger.__new__(LangFuseLogger) + logger.api_client = MagicMock() + logger.tracing = acquire() + sibling = acquire() + self.cache.cache_dict["test_key"] = logger + self.cache.ttl_dict["test_key"] = time.time() + 100 + + self.cache._remove_key("test_key") + self.cache._remove_key("test_key") + + assert litellm.initialized_langfuse_clients == 2 + assert "test_key" not in self.cache.cache_dict and "test_key" not in self.cache.ttl_dict + assert acquire() is sibling, "the second evictor took the sibling logger's hold on the channel" + release_langfuse_tracing(sibling) + release_langfuse_tracing(sibling, grace_seconds=0.0) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 6c0e2095300..102038cad30 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -1,5 +1,6 @@ import asyncio import contextlib +import copy import datetime import json import logging @@ -8833,3 +8834,41 @@ def test_extract_response_obj_and_hidden_params_reads_binary_content_hidden_para assert hidden_params == {"headers": {"x-request-id": "req_tts"}} assert response_obj["object"] == "binary" + + +def _preserved_thinking_client_turns() -> tuple[list[dict], list[dict]]: + turn_n = [{"role": "user", "content": "First question"}] + reply = { + "role": "assistant", + "content": "First answer", + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": "sig-1"}], + } + return turn_n, [*turn_n, reply, {"role": "user", "content": "Second question"}] + + +@pytest.mark.asyncio +async def test_prompt_management_with_unchanged_variables_replays_a_byte_identical_prefix(logging_obj, tmp_path): + """A prompt template rendered with the same variables on every turn must prepend the + same messages, or the signed thinking blocks in the history lose their binding.""" + from litellm.integrations.dotprompt.dotprompt_manager import DotpromptManager + + (tmp_path / "greeting.prompt").write_text( + "---\nmodel: claude-fable-5-1\n---\nSystem: You are a {{persona}}. Answer in one sentence.\n" + ) + manager = DotpromptManager(prompt_directory=str(tmp_path)) + compiled = [ + await logging_obj.async_get_chat_completion_prompt( + model="claude-fable-5-1", + messages=copy.deepcopy(turn), + non_default_params={}, + prompt_variables={"persona": "pirate"}, + prompt_id="greeting", + prompt_management_logger=manager, + ) + for turn in _preserved_thinking_client_turns() + ] + (_, messages_n, _), (_, messages_n_plus_one, _) = compiled + + assert json.dumps(messages_n_plus_one[: len(messages_n)], sort_keys=True) == json.dumps(messages_n, sort_keys=True) + assert messages_n[0] == {"role": "system", "content": "You are a pirate. Answer in one sentence."} + assert len(messages_n_plus_one) == len(messages_n) + 2 diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 7167a67d80d..729f46ec57f 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -1,4 +1,4 @@ - +import copy import json from typing import Final from unittest.mock import MagicMock, patch @@ -17,10 +17,18 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, RESPONSE_FORMAT_TOOL_NAME, ) +from litellm.litellm_core_utils.prompt_templates.common_utils import encrypted_reasoning_signature from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, ) +from litellm.llms.azure_ai.anthropic.transformation import AzureAnthropicConfig +from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( + AmazonAnthropicClaudeConfig, +) +from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import ( + VertexAIAnthropicConfig, +) from litellm.types.llms.anthropic import ANTHROPIC_BETA_HEADER_VALUES from litellm.types.utils import ServerToolUse, Usage @@ -37,13 +45,9 @@ def test_response_format_transformation_unit_test(): "additionalProperties": False, } - result = config._create_json_tool_call_for_response_format( - json_schema=response_format_json_schema - ) + result = config._create_json_tool_call_for_response_format(json_schema=response_format_json_schema) - assert result["input_schema"]["properties"] == { - "agent_doing": {"title": "Agent Doing", "type": "string"} - } + assert result["input_schema"]["properties"] == {"agent_doing": {"title": "Agent Doing", "type": "string"}} print(result) @@ -128,7 +132,9 @@ def test_calculate_usage_prefers_served_speed_from_response_usage(): assert no_response_speed.speed == "fast" -@pytest.mark.parametrize("input_update, expected_fresh", [({}, 1000), ({"input_tokens": 0}, 0), ({"input_tokens": 2000}, 2000)]) +@pytest.mark.parametrize( + "input_update, expected_fresh", [({}, 1000), ({"input_tokens": 0}, 0), ({"input_tokens": 2000}, 2000)] +) def test_streaming_iterator_persists_cumulative_usage_across_partial_chunks(input_update, expected_fresh): """ Omitted input/cache/pricing fields retain their last cumulative values; @@ -138,11 +144,17 @@ def test_streaming_iterator_persists_cumulative_usage_across_partial_chunks(inpu iterator = ModelResponseIterator(None, sync_stream=True, speed="fast") - start_usage = iterator._handle_usage({ - "input_tokens": 1000, "output_tokens": 1, "speed": "standard", "inference_geo": "us", - "cache_creation_input_tokens": 3000, "cache_read_input_tokens": 2000, - "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 3000}, - }) + start_usage = iterator._handle_usage( + { + "input_tokens": 1000, + "output_tokens": 1, + "speed": "standard", + "inference_geo": "us", + "cache_creation_input_tokens": 3000, + "cache_read_input_tokens": 2000, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 3000}, + } + ) delta_usage = iterator._handle_usage({"output_tokens": 5, **input_update}) assert start_usage.speed == "standard" @@ -564,9 +576,7 @@ def test_extract_response_content_with_citations(): }, } - _, citations, _, _, _, _, _, _ = config.extract_response_content( - completion_response - ) + _, citations, _, _, _, _, _, _ = config.extract_response_content(completion_response) assert citations == [ [ { @@ -639,12 +649,8 @@ def test_web_search_tool_transformation(): assert anthropic_web_search_tool["user_location"]["city"] == "San Francisco" -@pytest.mark.parametrize( - "search_context_size, expected_max_uses", [("low", 1), ("medium", 5), ("high", 10)] -) -def test_web_search_tool_transformation_with_search_context_size( - search_context_size, expected_max_uses -): +@pytest.mark.parametrize("search_context_size, expected_max_uses", [("low", 1), ("medium", 5), ("high", 10)]) +def test_web_search_tool_transformation_with_search_context_size(search_context_size, expected_max_uses): from litellm.types.llms.openai import OpenAIWebSearchOptions config = AnthropicConfig() @@ -819,10 +825,7 @@ def test_web_search_tool_result_in_provider_specific_fields(): assert "web_search_results" in provider_fields assert len(provider_fields["web_search_results"]) == 1 assert provider_fields["web_search_results"][0]["type"] == "web_search_tool_result" - assert ( - provider_fields["web_search_results"][0]["tool_use_id"] - == "srvtoolu_provider_test" - ) + assert provider_fields["web_search_results"][0]["tool_use_id"] == "srvtoolu_provider_test" def test_multiple_web_search_tool_results(): @@ -1046,10 +1049,7 @@ def test_transform_response_with_prefix_prompt(): ) assert result is not None - assert ( - result.choices[0].message.content - == "You are a helpful assistant. The grass is green." - ) + assert result.choices[0].message.content == "You are a helpful assistant. The grass is green." def test_get_supported_params_thinking(): @@ -1164,18 +1164,12 @@ def test_anthropic_beta_header_merging_with_output_format(): } } - result_headers = config.update_headers_with_optional_anthropic_beta( - headers, optional_params - ) + result_headers = config.update_headers_with_optional_anthropic_beta(headers, optional_params) # Both beta headers should be present beta_value = result_headers["anthropic-beta"] - assert ( - "context-1m-2025-08-07" in beta_value - ), f"User's context-1m beta header missing from: {beta_value}" - assert ( - "structured-outputs-2025-11-13" in beta_value - ), f"Structured output beta header missing from: {beta_value}" + assert "context-1m-2025-08-07" in beta_value, f"User's context-1m beta header missing from: {beta_value}" + assert "structured-outputs-2025-11-13" in beta_value, f"Structured output beta header missing from: {beta_value}" def test_anthropic_beta_header_merging_with_multiple_features(): @@ -1197,9 +1191,7 @@ def test_anthropic_beta_header_merging_with_multiple_features(): "tools": [{"type": "web_fetch_20250910", "name": "web_fetch"}], } - result_headers = config.update_headers_with_optional_anthropic_beta( - headers, optional_params - ) + result_headers = config.update_headers_with_optional_anthropic_beta(headers, optional_params) beta_value = result_headers["anthropic-beta"] @@ -1242,9 +1234,7 @@ def test_anthropic_structured_output_beta_header(): "strict": True, "schema": { "description": 'Progress report for the thinking process\n\nThis model represents a snapshot of the agent\'s current progress during\nthe thinking process, providing a brief description of the current activity.\n\nAttributes:\n agent_doing: Brief description of what the agent is currently doing.\n Should be kept under 10 words. Example: "Learning about home automation"', - "properties": { - "agent_doing": {"title": "Agent Doing", "type": "string"} - }, + "properties": {"agent_doing": {"title": "Agent Doing", "type": "string"}}, "required": ["agent_doing"], "title": "ThinkingStep", "type": "object", @@ -1258,10 +1248,7 @@ def test_anthropic_structured_output_beta_header(): assert response is not None print(f"response: {response}") print(f"raw_request_headers: {response['raw_request_headers']}") - assert ( - "structured-outputs-2025-11-13" - in response["raw_request_headers"]["anthropic-beta"] - ) + assert "structured-outputs-2025-11-13" in response["raw_request_headers"]["anthropic-beta"] @pytest.mark.parametrize( @@ -1397,9 +1384,7 @@ def test_tool_search_regex_detection(): config = AnthropicModelInfo() # Test with tool search regex tool - tools = [ - {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"} - ] + tools = [{"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}] assert config.is_tool_search_used(tools) is True # Test without tool search @@ -1414,9 +1399,7 @@ def test_tool_search_bm25_detection(): config = AnthropicModelInfo() # Test with tool search BM25 tool - tools = [ - {"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"} - ] + tools = [{"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"}] assert config.is_tool_search_used(tools) is True @@ -1608,9 +1591,7 @@ def test_tool_search_complete_response_parsing(): "tool_use_id": "srvtoolu_015i6aVA2niwzv4RG4DtnxDJ", "content": { "type": "tool_search_tool_search_result", - "tool_references": [ - {"type": "tool_reference", "tool_name": "get_weather"} - ], + "tool_references": [{"type": "tool_reference", "tool_name": "get_weather"}], }, }, {"type": "text", "text": "Great! I found a weather tool."}, @@ -1661,9 +1642,7 @@ def test_tool_search_complete_response_parsing(): assert usage.server_tool_use is not None assert usage.server_tool_use.web_search_requests == 0 - assert ( - usage.server_tool_use.tool_search_requests == 1 - ) # Counted from server_tool_use blocks + assert usage.server_tool_use.tool_search_requests == 1 # Counted from server_tool_use blocks def test_allowed_callers_field_preservation(): @@ -1715,9 +1694,7 @@ def test_programmatic_tool_calling_beta_header(): assert is_programmatic is True # Test header generation - headers = model_info.get_anthropic_headers( - api_key="test-key", programmatic_tool_calling_used=True - ) + headers = model_info.get_anthropic_headers(api_key="test-key", programmatic_tool_calling_used=True) assert "anthropic-beta" in headers assert "advanced-tool-use-2025-11-20" in headers["anthropic-beta"] @@ -1861,9 +1838,7 @@ def test_input_examples_beta_header(): assert is_examples_used is True # Test header generation - headers = model_info.get_anthropic_headers( - api_key="test-key", input_examples_used=True - ) + headers = model_info.get_anthropic_headers(api_key="test-key", input_examples_used=True) assert "anthropic-beta" in headers assert "advanced-tool-use-2025-11-20" in headers["anthropic-beta"] @@ -1949,10 +1924,7 @@ def test_input_examples_empty_list_not_added(): transformed_tool, _ = config._map_tool_helper(tool) assert transformed_tool is not None # Empty list should not be added - assert ( - "input_examples" not in transformed_tool - or len(transformed_tool.get("input_examples", [])) == 0 - ) + assert "input_examples" not in transformed_tool or len(transformed_tool.get("input_examples", [])) == 0 # ============ Effort Parameter Tests ============ @@ -2012,9 +1984,7 @@ def test_effort_beta_header_injection(): effort_used = model_info.is_effort_used(optional_params=optional_params, custom_llm_provider="anthropic") assert effort_used is True - headers = model_info.get_anthropic_headers( - api_key="test-key", effort_used=effort_used - ) + headers = model_info.get_anthropic_headers(api_key="test-key", effort_used=effort_used) assert "anthropic-beta" in headers assert "effort-2025-11-24" in headers["anthropic-beta"] @@ -2040,9 +2010,7 @@ def test_effort_validation(): optional_params = {"output_config": {"effort": "invalid"}} - with pytest.raises( - litellm.exceptions.BadRequestError, match="Invalid effort value" - ): + with pytest.raises(litellm.exceptions.BadRequestError, match="Invalid effort value"): config.transform_request( model="claude-opus-4-5-20251101", messages=messages, @@ -2278,16 +2246,8 @@ def test_anthropic_model_supports_speed_param_rejects_non_anthropic_providers( ): """Fast mode is direct-Anthropic-only. Vertex/Azure/Bedrock strip their prefix before the shared transform runs, so the bare Opus id must still be rejected.""" - assert ( - AnthropicConfig._model_supports_speed_param( - "claude-opus-4-8", custom_llm_provider - ) - is False - ) - assert ( - AnthropicConfig._model_supports_speed_param("claude-opus-4-8", "anthropic") - is True - ) + assert AnthropicConfig._model_supports_speed_param("claude-opus-4-8", custom_llm_provider) is False + assert AnthropicConfig._model_supports_speed_param("claude-opus-4-8", "anthropic") is True def test_vertex_anthropic_drops_speed_for_opus_with_drop_params(monkeypatch): @@ -2571,9 +2531,7 @@ def test_supports_effort_level_handles_provider_prefixes(model, level, expected) ("claude-opus-4-5-20251101", None, False), ], ) -def test_validate_effort_for_model_centralises_per_model_gating( - model, effort, expect_error -): +def test_validate_effort_for_model_centralises_per_model_gating(model, effort, expect_error): err = AnthropicConfig._validate_effort_for_model(model, effort, "anthropic") if expect_error: assert err is not None @@ -2622,11 +2580,7 @@ def test_transform_request_injects_dummy_tool_without_tools_param(): litellm.modify_params = prev_modify_params assert "tools" in result - names = [ - t.get("name") - for t in result["tools"] - if isinstance(t, dict) and t.get("name") is not None - ] + names = [t.get("name") for t in result["tools"] if isinstance(t, dict) and t.get("name") is not None] assert "dummy_tool" in names @@ -2692,13 +2646,9 @@ def test_calculate_usage_completion_tokens_details_with_reasoning(): "output_tokens": 500, } # Simulating reasoning content that would count as ~50 tokens - reasoning_content = ( - "Let me think about this step by step. " * 10 - ) # Roughly 50 tokens + reasoning_content = "Let me think about this step by step. " * 10 # Roughly 50 tokens - usage = config.calculate_usage( - usage_object=usage_object, reasoning_content=reasoning_content - ) + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=reasoning_content) # completion_tokens_details should be populated with both reasoning and text tokens assert usage.completion_tokens_details is not None @@ -2749,9 +2699,7 @@ def test_reasoning_effort_maps_to_adaptive_thinking_for_claude_4_6_models(): # reasoning_effort should not be in the result (it's transformed to thinking) assert "reasoning_effort" not in result # Should set output_config with the mapped effort value - assert ( - "output_config" in result - ), f"output_config missing for {model} with effort={effort}" + assert "output_config" in result, f"output_config missing for {model} with effort={effort}" assert result["output_config"]["effort"] == effort_map[effort] @@ -2852,9 +2800,7 @@ def test_raw_adaptive_thinking_untouched_for_46_plus_model(): ("gpt-4o", False), ], ) -def test_is_adaptive_thinking_model_is_sourced_from_cost_map( - local_model_cost_map, model, expected -): +def test_is_adaptive_thinking_model_is_sourced_from_cost_map(local_model_cost_map, model, expected): """Adaptive thinking resolves from the cost map first (an explicit supports_adaptive_thinking entry, or the anthropic-claude fallback rule for unmapped future Claudes), then from a date-safe opus/sonnet/haiku >= 4.6 name version as a @@ -2970,9 +2916,7 @@ def test_reasoning_effort_sets_output_config_for_46_models(): drop_params=False, ) - assert ( - "output_config" in result - ), f"output_config missing for {model} with effort={effort}" + assert "output_config" in result, f"output_config missing for {model} with effort={effort}" assert result["output_config"]["effort"] == effort @@ -3011,9 +2955,7 @@ def test_reasoning_effort_does_not_set_output_config_for_older_models(): drop_params=False, ) - assert ( - "output_config" not in result - ), f"output_config should not be set for {model}" + assert "output_config" not in result, f"output_config should not be set for {model}" @pytest.mark.parametrize( @@ -3053,14 +2995,10 @@ def test_reasoning_effort_accepts_dict_shape_for_adaptive_model(reasoning_effort ) # thinking must be set (adaptive for 4.6+) - assert ( - "thinking" in result - ), f"thinking missing for reasoning_effort={reasoning_effort_value!r}" + assert "thinking" in result, f"thinking missing for reasoning_effort={reasoning_effort_value!r}" assert result["thinking"]["type"] == "adaptive" # output_config must carry the mapped effort - assert ( - "output_config" in result - ), f"output_config missing for reasoning_effort={reasoning_effort_value!r}" + assert "output_config" in result, f"output_config missing for reasoning_effort={reasoning_effort_value!r}" assert result["output_config"]["effort"] == "low" @@ -3089,16 +3027,13 @@ def test_reasoning_effort_accepts_dict_shape_for_non_adaptive_model( drop_params=False, ) - assert ( - "thinking" in result - ), f"thinking missing for reasoning_effort={reasoning_effort_value!r}" + assert "thinking" in result, f"thinking missing for reasoning_effort={reasoning_effort_value!r}" assert result["thinking"]["type"] == "enabled" assert "budget_tokens" in result["thinking"] assert result["thinking"]["budget_tokens"] > 0 # Older models must not get adaptive-thinking output_config assert "output_config" not in result, ( - f"output_config should not be set for non-adaptive model " - f"(reasoning_effort={reasoning_effort_value!r})" + f"output_config should not be set for non-adaptive model (reasoning_effort={reasoning_effort_value!r})" ) @@ -3149,12 +3084,8 @@ def test_reasoning_effort_unparseable_dict_is_dropped(bad_value): model="claude-sonnet-4-6-20260219", drop_params=False, ) - assert ( - "thinking" not in result - ), f"thinking should not be set for bad value {bad_value!r}" - assert ( - "output_config" not in result - ), f"output_config should not be set for bad value {bad_value!r}" + assert "thinking" not in result, f"thinking should not be set for bad value {bad_value!r}" + assert "output_config" not in result, f"output_config should not be set for bad value {bad_value!r}" @pytest.mark.parametrize( @@ -3285,9 +3216,7 @@ def test_reasoning_effort_garbage_raises_bad_request(effort): ("max", DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET), ], ) -def test_reasoning_effort_xhigh_max_maps_to_budget_on_budget_model( - effort, expected_budget -): +def test_reasoning_effort_xhigh_max_maps_to_budget_on_budget_model(effort, expected_budget): """``xhigh`` / ``max`` extend the budget_tokens progression on budget-mode models.""" config = AnthropicConfig() @@ -3434,17 +3363,11 @@ def test_code_execution_tool_results_extraction(): # Verify first tool call assert transformed_response.choices[0].message.tool_calls[0].id == "srvtoolu_01ABC" - assert ( - transformed_response.choices[0].message.tool_calls[0].function.name - == "bash_code_execution" - ) + assert transformed_response.choices[0].message.tool_calls[0].function.name == "bash_code_execution" # Verify second tool call assert transformed_response.choices[0].message.tool_calls[1].id == "srvtoolu_01DEF" - assert ( - transformed_response.choices[0].message.tool_calls[1].function.name - == "text_editor_code_execution" - ) + assert transformed_response.choices[0].message.tool_calls[1].function.name == "text_editor_code_execution" # Verify tool results are in provider_specific_fields provider_fields = transformed_response.choices[0].message.provider_specific_fields @@ -3467,10 +3390,7 @@ def test_code_execution_tool_results_extraction(): assert editor_result["content"]["is_file_update"] is False # Verify text content is properly concatenated - assert ( - "I'll calculate that for you." - in transformed_response.choices[0].message.content - ) + assert "I'll calculate that for you." in transformed_response.choices[0].message.content assert "Done!" in transformed_response.choices[0].message.content @@ -3538,10 +3458,7 @@ def test_code_execution_tool_results_in_hidden_params(): assert "provider_specific_fields" in hidden assert "tool_results" in hidden["provider_specific_fields"] assert len(hidden["provider_specific_fields"]["tool_results"]) == 1 - assert ( - hidden["provider_specific_fields"]["tool_results"][0]["content"]["stdout"] - == "hello\n" - ) + assert hidden["provider_specific_fields"]["tool_results"][0]["content"]["stdout"] == "hello\n" def test_tool_search_tool_result_not_in_tool_results(): @@ -3737,10 +3654,7 @@ def test_compaction_block_in_provider_specific_fields(): assert "compaction_blocks" in provider_fields assert len(provider_fields["compaction_blocks"]) == 1 assert provider_fields["compaction_blocks"][0]["type"] == "compaction" - assert ( - "Summary of the conversation" - in provider_fields["compaction_blocks"][0]["content"] - ) + assert "Summary of the conversation" in provider_fields["compaction_blocks"][0]["content"] def test_multiple_compaction_blocks(): @@ -3775,12 +3689,22 @@ def test_multiple_compaction_blocks(): assert compaction_blocks[1]["content"] == "Second summary..." -@pytest.mark.parametrize("messages_api,gateway,native_endpoint", [ - (False, False, False), (True, False, False), (False, True, False), (True, True, False), (True, True, True), -]) +@pytest.mark.parametrize( + "messages_api,gateway,native_endpoint", + [ + (False, False, False), + (True, False, False), + (False, True, False), + (True, True, False), + (True, True, True), + ], +) async def test_native_compaction_wire_roundtrip( - messages_api: bool, gateway: bool, native_endpoint: bool, - monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, + messages_api: bool, + gateway: bool, + native_endpoint: bool, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, ) -> None: monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True") @@ -3788,8 +3712,11 @@ async def test_native_compaction_wire_roundtrip( monkeypatch.setattr(litellm, "use_chat_completions_url_for_anthropic_messages", False) block: Final = {"type": "compaction", "content": "Exact summary", "signature": "opaque-signature"} operation: Final = {"type": "summarize", "instructions": "Keep identifiers"} - usage: Final = {"input_tokens": 0, "output_tokens": 0, - "iterations": [{"type": "compaction", "input_tokens": 103, "output_tokens": 165}]} + usage: Final = { + "input_tokens": 0, + "output_tokens": 0, + "iterations": [{"type": "compaction", "input_tokens": 103, "output_tokens": 165}], + } chat_wire: Final = gateway and not native_endpoint base: Final = "https://gateway.test/v1" if gateway else "https://api.anthropic.com/v1" route: Final = respx_mock.post(f"{base}/{'chat/completions' if chat_wire else 'messages'}") @@ -3798,27 +3725,51 @@ async def test_native_compaction_wire_roundtrip( payload: Final = json.loads(request.content) assert len(request.headers.get_list("anthropic-beta")) == 1 assert {value.strip() for value in request.headers["anthropic-beta"].split(",")} == { - "compact-2026-09-04", "interleaved-thinking-2025-05-14", + "compact-2026-09-04", + "interleaved-thinking-2025-05-14", } if "compaction" in payload: assert payload["compaction"] == operation else: assert payload["messages"][0] == {"role": "assistant", "content": [block]} body: Final = ( - {"id": "chatcmpl_compact", "object": "chat.completion", "created": 1, "model": "claude-sonnet-5", - "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "", - "provider_specific_fields": {"compaction_blocks": [block]}}}], - "usage": {"prompt_tokens": 103, "completion_tokens": 165, "total_tokens": 268}} - if chat_wire else - {"id": "msg_compact", "type": "message", "role": "assistant", "model": "claude-sonnet-5", - "content": [block], "stop_reason": "compaction", "usage": usage} + { + "id": "chatcmpl_compact", + "object": "chat.completion", + "created": 1, + "model": "claude-sonnet-5", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": "", + "provider_specific_fields": {"compaction_blocks": [block]}, + }, + } + ], + "usage": {"prompt_tokens": 103, "completion_tokens": 165, "total_tokens": 268}, + } + if chat_wire + else { + "id": "msg_compact", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5", + "content": [block], + "stop_reason": "compaction", + "usage": usage, + } ) return httpx.Response(200, json=body) route.mock(side_effect=respond) call: Final = litellm.anthropic.messages.acreate if messages_api else litellm.acompletion params: Final = dict( - model=f"{'openai/' if gateway else ''}anthropic/claude-sonnet-5", api_key="test", max_tokens=512, + model=f"{'openai/' if gateway else ''}anthropic/claude-sonnet-5", + api_key="test", + max_tokens=512, api_base=base if gateway else "https://api.anthropic.com", extra_headers={"Anthropic-Beta": f"interleaved-thinking-2025-05-14{',compact-2026-09-04' if gateway else ''}"}, model_info={"supported_endpoints": ["/v1/messages"]} if native_endpoint else {}, @@ -3852,9 +3803,7 @@ def test_compaction_block_request_transformation(): {"role": "user", "content": "What is the weather in San Francisco?"}, { "role": "assistant", - "content": [ - {"type": "text", "text": "I don't have access to real-time data."} - ], + "content": [{"type": "text", "text": "I don't have access to real-time data."}], "provider_specific_fields": { "compaction_blocks": [ { @@ -3867,9 +3816,7 @@ def test_compaction_block_request_transformation(): {"role": "user", "content": "What about New York?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-opus-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-opus-4-6", llm_provider="anthropic") # Find the assistant message assistant_message = None @@ -3983,9 +3930,7 @@ def test_map_openai_context_management_to_anthropic(): "instructions": "Focus on preserving code snippets", } ] - result = config.map_openai_context_management_to_anthropic( - openai_format_with_instructions - ) + result = config.map_openai_context_management_to_anthropic(openai_format_with_instructions) assert result is not None assert result["edits"][0]["trigger"]["value"] == 150000 @@ -4012,9 +3957,7 @@ def test_map_openai_params_with_context_management(): config = AnthropicConfig() # Test with OpenAI list format - non_default_params = { - "context_management": [{"type": "compaction", "compact_threshold": 200000}] - } + non_default_params = {"context_management": [{"type": "compaction", "compact_threshold": 200000}]} optional_params = {} result = config.map_openai_params( @@ -4051,10 +3994,7 @@ def test_map_openai_params_with_context_management(): ) assert "context_management" in result - assert ( - result["context_management"] - == non_default_params_anthropic["context_management"] - ) + assert result["context_management"] == non_default_params_anthropic["context_management"] def test_cache_control_in_supported_params(): @@ -4165,10 +4105,7 @@ def test_compaction_block_empty_list_not_added(): # Verify compaction_blocks is not in provider_specific_fields when there are none provider_fields = result.choices[0].message.provider_specific_fields if provider_fields: - assert ( - "compaction_blocks" not in provider_fields - or provider_fields.get("compaction_blocks") is None - ) + assert "compaction_blocks" not in provider_fields or provider_fields.get("compaction_blocks") is None def test_fast_mode_beta_header(): @@ -4217,9 +4154,7 @@ def test_fast_mode_usage_calculation(): "output_tokens": 500, } - usage = config.calculate_usage( - usage_object=usage_object, reasoning_content=None, speed="fast" - ) + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None, speed="fast") assert usage.prompt_tokens == 1000 assert usage.completion_tokens == 500 @@ -4240,9 +4175,7 @@ def test_fast_mode_cost_calculation(): base_completion = 0.025 with ( - patch( - "litellm.llms.anthropic.cost_calculation.generic_cost_per_token" - ) as mock_cost, + patch("litellm.llms.anthropic.cost_calculation.generic_cost_per_token") as mock_cost, patch("litellm.get_model_info") as mock_info, ): mock_cost.return_value = (base_prompt, base_completion) @@ -4282,9 +4215,7 @@ def test_fast_mode_with_inference_geo(): base_completion = 0.025 with ( - patch( - "litellm.llms.anthropic.cost_calculation.generic_cost_per_token" - ) as mock_cost, + patch("litellm.llms.anthropic.cost_calculation.generic_cost_per_token") as mock_cost, patch("litellm.get_model_info") as mock_info, ): mock_cost.return_value = (base_prompt, base_completion) @@ -4475,9 +4406,7 @@ def test_map_tool_helper_enforces_object_type_when_missing(): "name": "search_code", "description": "Search for code patterns", "parameters": { - "properties": { - "query": {"type": "string", "description": "Search query"} - }, + "properties": {"query": {"type": "string", "description": "Search query"}}, "required": ["query"], }, }, @@ -4490,9 +4419,9 @@ def test_map_tool_helper_enforces_object_type_when_missing(): assert "properties" in result["input_schema"] assert "query" in result["input_schema"]["properties"] # Original parameters dict must not be modified in place - assert ( - tool["function"]["parameters"] == original_params - ), "parameters dict was mutated; _map_tool_helper should not modify caller data" + assert tool["function"]["parameters"] == original_params, ( + "parameters dict was mutated; _map_tool_helper should not modify caller data" + ) def test_map_tool_helper_enforces_object_type_when_wrong_type(): @@ -4518,13 +4447,13 @@ def test_map_tool_helper_enforces_object_type_when_wrong_type(): result, _ = config._map_tool_helper(tool) assert result is not None assert result["input_schema"]["type"] == "object" - assert ( - result["input_schema"].get("properties") == {} - ), "properties should be injected as {} when schema has non-object type and no properties key" + assert result["input_schema"].get("properties") == {}, ( + "properties should be injected as {} when schema has non-object type and no properties key" + ) # Original parameters dict must not be modified in place - assert ( - tool["function"]["parameters"] == original_params - ), "parameters dict was mutated; _map_tool_helper should not modify caller data" + assert tool["function"]["parameters"] == original_params, ( + "parameters dict was mutated; _map_tool_helper should not modify caller data" + ) def test_map_tool_helper_preserves_valid_object_schema(): @@ -4591,12 +4520,8 @@ def test_extract_response_content_thinking_block_null_thinking(): {"type": "text", "text": "Hello"}, ] } - text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content( - completion_response_null - ) - assert ( - thinking_blocks is not None - ), "thinking blocks should not be None when thinking=null" + text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content(completion_response_null) + assert thinking_blocks is not None, "thinking blocks should not be None when thinking=null" assert len(thinking_blocks) == 1 assert "Hello" in text @@ -4607,12 +4532,8 @@ def test_extract_response_content_thinking_block_null_thinking(): {"type": "text", "text": "World"}, ] } - text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content( - completion_response_missing - ) - assert ( - thinking_blocks is not None - ), "thinking blocks should not be None when thinking key is absent" + text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content(completion_response_missing) + assert thinking_blocks is not None, "thinking blocks should not be None when thinking key is absent" assert len(thinking_blocks) == 1 assert "World" in text @@ -4623,9 +4544,7 @@ def test_extract_response_content_thinking_block_null_thinking(): {"type": "text", "text": "Done"}, ] } - text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content( - completion_response_text - ) + text, _, thinking_blocks, _, _, _, _, _ = config.extract_response_content(completion_response_text) assert thinking_blocks is not None assert len(thinking_blocks) == 1 assert thinking_blocks[0]["thinking"] == "Let me think..." @@ -4684,12 +4603,8 @@ def test_advisor_beta_header_injected(): } ] } - result = config.update_headers_with_optional_anthropic_beta( - headers, optional_params - ) - assert ANTHROPIC_BETA_HEADER_VALUES.ADVISOR_TOOL_2026_03_01.value in result.get( - "anthropic-beta", "" - ) + result = config.update_headers_with_optional_anthropic_beta(headers, optional_params) + assert ANTHROPIC_BETA_HEADER_VALUES.ADVISOR_TOOL_2026_03_01.value in result.get("anthropic-beta", "") def test_advisor_beta_header_not_injected_without_tool(): @@ -4697,9 +4612,7 @@ def test_advisor_beta_header_not_injected_without_tool(): config = AnthropicConfig() headers: dict = {} optional_params: dict = {"tools": []} - result = config.update_headers_with_optional_anthropic_beta( - headers, optional_params - ) + result = config.update_headers_with_optional_anthropic_beta(headers, optional_params) assert "advisor-tool-2026-03-01" not in result.get("anthropic-beta", "") @@ -4726,9 +4639,7 @@ def test_advisor_tool_result_preserved_in_response(): {"type": "text", "text": "Here is the implementation."}, ] } - text, _, _, _, tool_calls, _, tool_results, _ = config.extract_response_content( - completion_response - ) + text, _, _, _, tool_calls, _, tool_results, _ = config.extract_response_content(completion_response) assert "Consulting advisor." in text assert "Here is the implementation." in text # server_tool_use (advisor) should be a tool_call @@ -4843,9 +4754,7 @@ def test_basic_sanitize_anthropic_tool_name_replaces_invalid_chars(): ) assert ( - _basic_sanitize_anthropic_tool_name( - "github_openapi_mcp-actions/download-job-logs-for-workflow-run" - ) + _basic_sanitize_anthropic_tool_name("github_openapi_mcp-actions/download-job-logs-for-workflow-run") == "github_openapi_mcp-actions_download-job-logs-for-workflow-run" ) # other punctuation @@ -4874,9 +4783,7 @@ def test_build_anthropic_tool_name_maps_no_collisions(): ] ) assert forward == { - "actions/download-job-logs-for-workflow-run": ( - "actions_download-job-logs-for-workflow-run" - ), + "actions/download-job-logs-for-workflow-run": ("actions_download-job-logs-for-workflow-run"), "pulls/list-files": "pulls_list-files", } assert reverse == {v: k for k, v in forward.items()} @@ -4927,9 +4834,7 @@ def test_build_anthropic_tool_name_maps_three_way_collision(): _build_anthropic_tool_name_maps, ) - forward, reverse = _build_anthropic_tool_name_maps( - ["foo_bar", "foo/bar", "foo.bar"] - ) + forward, reverse = _build_anthropic_tool_name_maps(["foo_bar", "foo/bar", "foo.bar"]) assert "foo_bar" not in forward # untouched assert forward["foo/bar"] == "foo_bar_2" assert forward["foo.bar"] == "foo_bar_3" @@ -5002,16 +4907,13 @@ def test_map_openai_params_does_not_pollute_optional_params_with_internal_keys() ) # No internal keys may appear in optional_params for ANY input. for key in optional_params: - assert not key.startswith( - "_anthropic_tool_name" - ), f"optional_params leaked internal key {key!r}: {optional_params}" + assert not key.startswith("_anthropic_tool_name"), ( + f"optional_params leaked internal key {key!r}: {optional_params}" + ) # And no key starting with `_` either; optional_params should only # contain documented Anthropic Messages API parameters. for key in optional_params: - assert not key.startswith("_"), ( - f"optional_params leaked underscore-prefixed key {key!r}: " - f"{optional_params}" - ) + assert not key.startswith("_"), f"optional_params leaked underscore-prefixed key {key!r}: {optional_params}" def test_map_openai_params_no_maps_when_all_names_already_valid(): @@ -5040,11 +4942,7 @@ def test_map_openai_params_no_maps_when_all_names_already_valid(): def test_rewrite_tool_names_in_messages_uses_forward_map(): config = AnthropicConfig() - forward_map = { - "actions/download-job-logs-for-workflow-run": ( - "actions_download-job-logs-for-workflow-run" - ) - } + forward_map = {"actions/download-job-logs-for-workflow-run": ("actions_download-job-logs-for-workflow-run")} messages = [ {"role": "user", "content": "go"}, { @@ -5067,15 +4965,9 @@ def test_rewrite_tool_names_in_messages_uses_forward_map(): out = config._rewrite_tool_names_in_messages(messages, forward_map) # input list must not be mutated - assert ( - messages[1]["tool_calls"][0]["function"]["name"] - == "actions/download-job-logs-for-workflow-run" - ) + assert messages[1]["tool_calls"][0]["function"]["name"] == "actions/download-job-logs-for-workflow-run" # output rewritten according to forward map - assert ( - out[1]["tool_calls"][0]["function"]["name"] - == "actions_download-job-logs-for-workflow-run" - ) + assert out[1]["tool_calls"][0]["function"]["name"] == "actions_download-job-logs-for-workflow-run" # non-tool-call messages pass through unchanged (same object) assert out[0] is messages[0] assert out[2] is messages[2] @@ -5151,9 +5043,7 @@ def test_sanitize_tool_names_in_request_does_not_mutate_caller_tool_dicts(): caller_tools = [caller_tool] optional_params: dict = {"tools": caller_tools} - forward, reverse = config._sanitize_tool_names_in_request( - optional_params=optional_params - ) + forward, reverse = config._sanitize_tool_names_in_request(optional_params=optional_params) assert forward.get(original_name) sanitized = forward[original_name] @@ -5302,10 +5192,7 @@ def test_streaming_iterator_reverse_maps_tool_use_name(): parsed = iterator.chunk_parser(chunk=chunk) tool_calls = parsed.choices[0].delta.tool_calls assert tool_calls is not None and len(tool_calls) == 1 - assert ( - tool_calls[0]["function"]["name"] - == "actions/download-job-logs-for-workflow-run" - ) + assert tool_calls[0]["function"]["name"] == "actions/download-job-logs-for-workflow-run" def test_streaming_iterator_passthrough_when_name_not_in_map(): @@ -5401,9 +5288,9 @@ def test_transform_request_does_not_leak_internal_keys_into_body(): for tool in data.get("tools", []): name = tool.get("name") assert isinstance(name, str) - assert _re.fullmatch( - r"[a-zA-Z0-9_-]{1,128}", name - ), f"sanitized tool name {name!r} still violates Anthropic regex" + assert _re.fullmatch(r"[a-zA-Z0-9_-]{1,128}", name), ( + f"sanitized tool name {name!r} still violates Anthropic regex" + ) # Sent name for the bad tool is the disambiguated form, valid name passes through. sent_names = {t["name"] for t in data["tools"]} @@ -5539,9 +5426,7 @@ def test_transform_request_rewrites_tool_names_in_history(): for block in content: if isinstance(block, dict) and block.get("type") == "tool_use": tool_use_names.append(block.get("name")) - assert ( - tool_use_names - ), "expected at least one tool_use block in transformed messages" + assert tool_use_names, "expected at least one tool_use block in transformed messages" for name in tool_use_names: assert name == "actions_download-job-logs-for-workflow-run", ( f"history tool_use.name {name!r} not rewritten -- Anthropic will " @@ -5565,19 +5450,12 @@ def test_sanitize_tool_names_in_request_skips_hosted_tools(): } forward, reverse = AnthropicConfig._sanitize_tool_names_in_request(optional_params) # Only the custom tool was rewritten. - assert forward == { - "actions/download-job-logs-for-workflow-run": "actions_download-job-logs-for-workflow-run" - } - assert reverse == { - "actions_download-job-logs-for-workflow-run": "actions/download-job-logs-for-workflow-run" - } + assert forward == {"actions/download-job-logs-for-workflow-run": "actions_download-job-logs-for-workflow-run"} + assert reverse == {"actions_download-job-logs-for-workflow-run": "actions/download-job-logs-for-workflow-run"} # Hosted tool's name unchanged. assert optional_params["tools"][0]["name"] == "web_search" # Custom tool's name updated in place. - assert ( - optional_params["tools"][1]["name"] - == "actions_download-job-logs-for-workflow-run" - ) + assert optional_params["tools"][1]["name"] == "actions_download-job-logs-for-workflow-run" def test_sanitize_tool_names_in_request_no_tools_is_noop(): @@ -5811,9 +5689,7 @@ def test_translate_system_message_keeps_billing_header_for_first_party_anthropic assert config.should_strip_billing_metadata() is False result = config.translate_system_message( - messages=_system_with_billing_header( - "You are Claude Code, Anthropic's official CLI for Claude." - ) + messages=_system_with_billing_header("You are Claude Code, Anthropic's official CLI for Claude.") ) texts = [block["text"] for block in result] @@ -5829,9 +5705,7 @@ def test_translate_system_message_strips_billing_header_for_bedrock(): config = BedrockClaudePlatformConfig() assert config.should_strip_billing_metadata() is True - result = config.translate_system_message( - messages=_system_with_billing_header("real system prompt") - ) + result = config.translate_system_message(messages=_system_with_billing_header("real system prompt")) texts = [block["text"] for block in result] assert all(not t.startswith("x-anthropic-billing-header:") for t in texts) @@ -5897,9 +5771,7 @@ def test_translate_system_message_strips_billing_header_for_bedrock_invoke(): config = AmazonAnthropicClaudeConfig() assert config.should_strip_billing_metadata() is True - result = config.translate_system_message( - messages=_system_with_billing_header("real system prompt") - ) + result = config.translate_system_message(messages=_system_with_billing_header("real system prompt")) texts = [block["text"] for block in result] assert all(not t.startswith("x-anthropic-billing-header:") for t in texts) @@ -5953,9 +5825,7 @@ def test_translate_system_message_strips_billing_header_for_bedrock_invoke(): ), ], ) -def test_should_strip_billing_metadata_by_provider( - module_path, class_name, expected_strip -): +def test_should_strip_billing_metadata_by_provider(module_path, class_name, expected_strip): import importlib config_cls = getattr(importlib.import_module(module_path), class_name) @@ -6127,12 +5997,8 @@ def test_sampling_param_gating_driven_by_model_map_flag(monkeypatch): """The drop/raise decision must come from ``supports_sampling_params`` in the model map, not just name matching: a flagged entry gates a model whose name says nothing, and an explicit ``true`` overrides the name fallback.""" - monkeypatch.setitem( - litellm.model_cost, "claude-zeta-9", {"supports_sampling_params": False} - ) - monkeypatch.setitem( - litellm.model_cost, "claude-fable-5-test", {"supports_sampling_params": True} - ) + monkeypatch.setitem(litellm.model_cost, "claude-zeta-9", {"supports_sampling_params": False}) + monkeypatch.setitem(litellm.model_cost, "claude-fable-5-test", {"supports_sampling_params": True}) config = AnthropicConfig() flagged_off = config.map_openai_params( @@ -6252,9 +6118,7 @@ def test_is_anthropic_usage_object_rejects_responses_api_usage(): ("claude-sonnet-4-5-20250929", False), ], ) -def test_disabled_thinking_omitted_only_for_always_on_models( - local_model_cost_map, model, expected_dropped -): +def test_disabled_thinking_omitted_only_for_always_on_models(local_model_cost_map, model, expected_dropped): """``thinking={"type": "disabled"}`` is omitted for always-on-thinking models (Fable/Mythos, which 400 on it: the API remedy is to omit the param) and is forwarded verbatim for every model that accepts it.""" @@ -6300,9 +6164,7 @@ def test_forced_tool_choice_raises_clean_error_on_fable_5_1_without_drop_params( "tool_choice", ["required", {"type": "required"}, {"type": "function", "function": {"name": "get_weather"}}], ) -def test_forced_tool_choice_downgraded_to_auto_on_fable_5_1_with_drop_params( - local_model_cost_map, tool_choice -): +def test_forced_tool_choice_downgraded_to_auto_on_fable_5_1_with_drop_params(local_model_cost_map, tool_choice): config = AnthropicConfig() result = config.map_openai_params( @@ -6329,9 +6191,7 @@ def test_forced_tool_choice_downgrade_keeps_parallel_tool_calls_flag(local_model @pytest.mark.parametrize("tool_choice, expected_type", [("auto", "auto"), ("none", "none")]) -def test_unforced_tool_choice_forwarded_on_fable_5_1( - local_model_cost_map, tool_choice, expected_type, monkeypatch -): +def test_unforced_tool_choice_forwarded_on_fable_5_1(local_model_cost_map, tool_choice, expected_type, monkeypatch): monkeypatch.setattr(litellm, "drop_params", False) config = AnthropicConfig() @@ -6346,9 +6206,7 @@ def test_unforced_tool_choice_forwarded_on_fable_5_1( @pytest.mark.parametrize("model", ["claude-fable-5", "claude-opus-5", "claude-sonnet-5"]) -def test_forced_tool_choice_forwarded_on_models_that_support_it( - local_model_cost_map, model, monkeypatch -): +def test_forced_tool_choice_forwarded_on_models_that_support_it(local_model_cost_map, model, monkeypatch): monkeypatch.setattr(litellm, "drop_params", False) config = AnthropicConfig() @@ -6519,3 +6377,424 @@ def test_eager_input_streaming_reaches_anthropic_request_tools(): assert result["tools"][0]["eager_input_streaming"] is True assert result["tools"][0]["name"] == "write_file" + + +# --------------------------------------------------------------------------- +# Mid-conversation ``role: "system"`` on the chat completions path. +# +# Hoisting a later system message into the top-level ``system`` block rewrites +# the cached prefix and re-bills the whole conversation at cache-write pricing +# on every reminder (#36559). The chat path must keep the prefix stable: leading +# system messages still become the ``system`` param, later ones stay in place as +# ``role: "system"`` on models flagged ``supports_mid_conversation_system`` and +# become a user turn on models that reject the role inside ``messages``. +# --------------------------------------------------------------------------- + +UNFLAGGED_CLAUDE = "claude-opus-4-7" +FLAGGED_CLAUDE = "claude-opus-4-8" +CONVERTED_SYSTEM_NOTE = ( + "Operator note (not from the user): the following was originally a mid-conversation system-role reminder." +) +REMINDER_TEXT = "Answer with exactly one word." +CACHED_SYSTEM_BLOCK = {"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}} + + +def _chat_request(config: AnthropicConfig, model: str, messages: list[dict]) -> dict: + return config.transform_request( + model=model, + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + ) + + +def _reminder_conversation() -> list[dict]: + """The shape Claude Code sends mid-session: cached system prompt, turns, a + reminder right after a user turn, an assistant turn, a fresh user turn.""" + return [ + {"role": "system", "content": [dict(CACHED_SYSTEM_BLOCK)]}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + {"role": "system", "content": REMINDER_TEXT}, + {"role": "assistant", "content": "Second answer"}, + {"role": "user", "content": "Third question"}, + ] + + +def _texts(message: dict) -> list[str]: + return [block["text"] for block in message["content"] if block.get("type") == "text"] + + +def test_chat_unflagged_model_converts_mid_conversation_system_to_user_turn(local_model_cost_map): + result = _chat_request(AnthropicConfig(), UNFLAGGED_CLAUDE, _reminder_conversation()) + + assert result["system"] == [CACHED_SYSTEM_BLOCK] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "assistant", "user"] + assert _texts(result["messages"][2]) == ["Second question", CONVERTED_SYSTEM_NOTE, REMINDER_TEXT] + + +def test_chat_flagged_model_keeps_mid_conversation_system_in_messages(local_model_cost_map): + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, _reminder_conversation()) + + assert result["system"] == [CACHED_SYSTEM_BLOCK] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] + assert result["messages"][3] == {"role": "system", "content": [{"type": "text", "text": REMINDER_TEXT}]} + + +def test_chat_flagged_model_keeps_cache_control_on_mid_conversation_system(local_model_cost_map): + messages = _reminder_conversation() + messages[4] = { + "role": "system", + "content": [{"type": "text", "text": REMINDER_TEXT, "cache_control": {"type": "ephemeral"}}], + } + + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, messages) + + assert result["messages"][3]["content"] == [ + {"type": "text", "text": REMINDER_TEXT, "cache_control": {"type": "ephemeral"}} + ] + + +def test_chat_flagged_model_moves_system_after_the_user_turn_it_precedes(local_model_cost_map): + """Anthropic only accepts role=system directly after a user turn; an + OpenAI-shaped client that puts the reminder before its next question gets a + placement-valid request without the reminder leaving ``messages``.""" + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "system", "content": REMINDER_TEXT}, + {"role": "user", "content": "Second question"}, + ] + + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, messages) + + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "system"] + assert _texts(result["messages"][2]) == ["Second question"] + assert _texts(result["messages"][3]) == [REMINDER_TEXT] + + +def test_chat_flagged_model_converts_system_with_no_following_user_turn(local_model_cost_map): + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "system", "content": REMINDER_TEXT}, + ] + + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, messages) + + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user"] + assert _texts(result["messages"][2]) == [CONVERTED_SYSTEM_NOTE, REMINDER_TEXT] + + +@pytest.mark.parametrize( + "empty_content", + [[], None, [{"type": "input_audio", "input_audio": {"data": "AAAA", "format": "wav"}}]], + ids=["empty-list", "none", "unsupported-part-only"], +) +def test_chat_flagged_model_converts_a_system_behind_a_user_turn_that_sends_nothing( + local_model_cost_map, empty_content +): + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": empty_content}, + {"role": "system", "content": REMINDER_TEXT}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + ] + + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, messages) + + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user"] + assert _texts(result["messages"][0]) == [CONVERTED_SYSTEM_NOTE, REMINDER_TEXT] + + +USER_PART_BY_TYPE = { + "text": {"type": "text", "text": "hello"}, + "image_url": {"type": "image_url", "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}}, + "document": {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": "hello"}}, + "file": {"type": "file", "file": {"file_data": "data:text/plain;base64,aGVsbG8=", "filename": "hello.txt"}}, + "input_audio": {"type": "input_audio", "input_audio": {"data": "AAAA", "format": "wav"}}, + "video_url": {"type": "video_url", "video_url": {"url": "https://example.com/clip.mp4"}}, +} + + +@pytest.mark.parametrize("part_type", sorted(USER_PART_BY_TYPE)) +def test_chat_flagged_model_anchors_a_system_on_a_user_turn_exactly_when_that_turn_reaches_the_wire( + local_model_cost_map, part_type +): + part_only_turn = {"role": "user", "content": [USER_PART_BY_TYPE[part_type]]} + tail = [{"role": "assistant", "content": "First answer"}, {"role": "user", "content": "Second question"}] + + without_reminder = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, [part_only_turn, *tail]) + with_reminder = _chat_request( + AnthropicConfig(), FLAGGED_CLAUDE, [part_only_turn, {"role": "system", "content": REMINDER_TEXT}, *tail] + ) + + turn_reaches_wire = [m["role"] for m in without_reminder["messages"]] == ["user", "assistant", "user"] + expected_roles = ["user", "system", "assistant", "user"] if turn_reaches_wire else ["user", "assistant", "user"] + assert [m["role"] for m in with_reminder["messages"]] == expected_roles + + +ASSISTANT_TURN_BY_SHAPE = { + "text": {"role": "assistant", "content": "First answer"}, + "tool-calls": { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "toolu_1", "type": "function", "function": {"name": "f", "arguments": "{}"}}], + }, + "signed-thinking-part": { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, + "empty-string": {"role": "assistant", "content": ""}, + "whitespace-string": {"role": "assistant", "content": " "}, + "empty-text-part": {"role": "assistant", "content": [{"type": "text", "text": ""}]}, + "none": {"role": "assistant", "content": None}, + "empty-list": {"role": "assistant", "content": []}, + "unsigned-thinking-part": {"role": "assistant", "content": [{"type": "thinking", "thinking": "hm"}]}, + "signed-thinking-block": { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, + "redacted-thinking-block": { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "redacted_thinking", "data": "x"}], + }, + "encrypted-thinking-part": { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "hm", "signature": encrypted_reasoning_signature("abc")}], + }, + "encrypted-redacted-thinking-block": { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "redacted_thinking", "data": encrypted_reasoning_signature("abc")}], + }, + "unsigned-inline-part-hides-signed-block": { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "hm"}], + "thinking_blocks": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, + "inline-redacted-part-hides-redacted-block": { + "role": "assistant", + "content": [{"type": "redacted_thinking", "data": "x"}], + "thinking_blocks": [{"type": "redacted_thinking", "data": "x"}], + }, + "text-part-beside-signed-block": { + "role": "assistant", + "content": [{"type": "text", "text": "First answer"}], + "thinking_blocks": [{"type": "thinking", "thinking": "hm", "signature": "s"}], + }, +} + + +@pytest.mark.parametrize("shape", sorted(ASSISTANT_TURN_BY_SHAPE)) +def test_chat_flagged_model_keeps_a_system_exactly_when_the_assistant_turn_after_it_reaches_the_wire( + local_model_cost_map, shape +): + first_turn = {"role": "user", "content": "First question"} + tail = [ASSISTANT_TURN_BY_SHAPE[shape], {"role": "user", "content": "Second question"}] + + without_reminder = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, [first_turn, *tail]) + with_reminder = _chat_request( + AnthropicConfig(), FLAGGED_CLAUDE, [first_turn, {"role": "system", "content": REMINDER_TEXT}, *tail] + ) + + turn_reaches_wire = [m["role"] for m in without_reminder["messages"]] == ["user", "assistant", "user"] + expected_roles = ["user", "system", "assistant", "user"] if turn_reaches_wire else ["user", "user"] + assert [m["role"] for m in with_reminder["messages"]] == expected_roles + if not turn_reaches_wire: + assert _texts(with_reminder["messages"][0]) == ["First question", CONVERTED_SYSTEM_NOTE, REMINDER_TEXT] + + +def test_chat_flagged_model_merges_adjacent_system_messages(local_model_cost_map): + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "First question"}, + {"role": "system", "content": "Reminder one."}, + {"role": "system", "content": "Reminder two."}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + ] + + result = _chat_request(AnthropicConfig(), FLAGGED_CLAUDE, messages) + + assert [m["role"] for m in result["messages"]] == ["user", "system", "assistant", "user"] + assert _texts(result["messages"][1]) == ["Reminder one.", "Reminder two."] + + +def test_chat_unflagged_model_keeps_tool_result_first_when_system_precedes_tool_message(local_model_cost_map): + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "Weather?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": "{}"}, + } + ], + }, + {"role": "system", "content": REMINDER_TEXT}, + {"role": "tool", "tool_call_id": "call_1", "content": "sunny"}, + {"role": "user", "content": "Thanks"}, + ] + + result = _chat_request(AnthropicConfig(), UNFLAGGED_CLAUDE, messages) + + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user"] + blocks = result["messages"][2]["content"] + assert blocks[0]["type"] == "tool_result" + assert blocks[0]["tool_use_id"] == "call_1" + assert _texts(result["messages"][2]) == [CONVERTED_SYSTEM_NOTE, REMINDER_TEXT, "Thanks"] + + +def test_chat_transform_request_does_not_mutate_caller_messages(local_model_cost_map): + messages = _reminder_conversation() + snapshot = copy.deepcopy(messages) + + _chat_request(AnthropicConfig(), UNFLAGGED_CLAUDE, messages) + + assert messages == snapshot + + +_CHAT_CONFIGS = [ + pytest.param(AnthropicConfig, UNFLAGGED_CLAUDE, id="anthropic-unflagged"), + pytest.param(AnthropicConfig, FLAGGED_CLAUDE, id="anthropic-flagged"), + pytest.param(VertexAIAnthropicConfig, UNFLAGGED_CLAUDE, id="vertex_ai-unflagged"), + pytest.param(VertexAIAnthropicConfig, FLAGGED_CLAUDE, id="vertex_ai-flagged"), + pytest.param(AzureAnthropicConfig, UNFLAGGED_CLAUDE, id="azure_ai-unflagged"), + pytest.param(AzureAnthropicConfig, FLAGGED_CLAUDE, id="azure_ai-flagged"), + pytest.param(AmazonAnthropicClaudeConfig, "invoke/us.anthropic.claude-opus-4-7", id="bedrock_invoke-unflagged"), + pytest.param(AmazonAnthropicClaudeConfig, "invoke/us.anthropic.claude-opus-4-8", id="bedrock_invoke-flagged"), +] + + +@pytest.mark.parametrize("config_cls, model", _CHAT_CONFIGS) +def test_chat_mid_conversation_system_keeps_earlier_turns_a_prefix_of_the_next_request( + local_model_cost_map, config_cls, model +): + """The provider-side prompt cache is a prefix match over ``system`` + + ``messages``. Whatever the policy for the reminder, turn N's request must + stay a prefix of turn N+1's request or the whole conversation is re-billed. + + Anthropic combines consecutive same-role messages into one turn, so the + cache-relevant sequence is ``(role, content block)`` pairs, not the message + list: a reminder that joins the preceding user turn still extends the prefix. + """ + conversation = _reminder_conversation() + + earlier = _chat_request(config_cls(), model, copy.deepcopy(conversation[:4])) + later = _chat_request(config_cls(), model, copy.deepcopy(conversation)) + + assert later["system"] == earlier["system"] + earlier_blocks = _role_block_pairs(earlier["messages"]) + later_blocks = _role_block_pairs(later["messages"]) + assert later_blocks[: len(earlier_blocks)] == earlier_blocks + assert len(later_blocks) > len(earlier_blocks) + + +def _role_block_pairs(messages: list[dict]) -> list[tuple[str, object]]: + return [ + (message["role"], block) + for message in messages + for block in (message["content"] if isinstance(message["content"], list) else [message["content"]]) + ] + + +def _thinking_reply(text: str) -> dict: + return { + "role": "assistant", + "content": text, + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": f"sig-{text}"}], + } + + +def _preserved_thinking_turns(reminder_after_user: bool) -> tuple[list[dict], list[dict], list[dict]]: + turn_n = [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "First question"}] + reminder = {"role": "system", "content": REMINDER_TEXT} + second_question = {"role": "user", "content": "Second question"} + second_turn = [second_question, reminder] if reminder_after_user else [reminder, second_question] + turn_n_plus_one = [*turn_n, _thinking_reply("First answer"), *second_turn] + turn_n_plus_two = [ + *turn_n_plus_one, + _thinking_reply("Second answer"), + {"role": "user", "content": "Third question"}, + ] + return turn_n, turn_n_plus_one, turn_n_plus_two + + +def _replayed_prefix(request: dict, message_count: int) -> str: + replayed = { + "system": request.get("system"), + "tools": request.get("tools"), + "messages": request["messages"][:message_count], + } + return json.dumps(replayed, sort_keys=True) + + +def _assert_prefix_stable(requests: list[dict]) -> None: + for earlier, later in zip(requests, requests[1:]): + count = len(earlier["messages"]) + assert _replayed_prefix(later, count) == _replayed_prefix(earlier, count) + + +@pytest.mark.parametrize("reminder_after_user", [True, False]) +def test_chat_flagged_model_replays_a_byte_identical_prefix_around_a_mid_conversation_reminder( + local_model_cost_map, reminder_after_user +): + """Preserved thinking binds each signed block to the request prefix it was created + under (``system``, ``tools`` and the earlier messages), so turn N's transformed + request must be a byte-identical prefix of turn N+1's or the block is dropped.""" + requests = [ + AnthropicConfig().transform_request( + model="claude-fable-5-1", messages=copy.deepcopy(turn), optional_params={}, litellm_params={}, headers={} + ) + for turn in _preserved_thinking_turns(reminder_after_user) + ] + + _assert_prefix_stable(requests) + assert [m["role"] for m in requests[1]["messages"]] == ["user", "assistant", "user", "system"] + assert [m["role"] for m in requests[2]["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] + + +def test_chat_dummy_tool_result_for_an_orphaned_tool_call_replays_a_byte_identical_prefix( + local_model_cost_map, monkeypatch +): + monkeypatch.setattr(litellm, "modify_params", True) + tools = [ + {"name": "lookup", "description": "Look something up", "input_schema": {"type": "object", "properties": {}}} + ] + orphaned_call = { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}], + } + turn_n = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "First question"}, + orphaned_call, + ] + turn_n_plus_one = [*turn_n, _thinking_reply("First answer"), {"role": "user", "content": "Second question"}] + requests = [ + AnthropicConfig().transform_request( + model="claude-fable-5-1", + messages=copy.deepcopy(turn), + optional_params={"tools": copy.deepcopy(tools)}, + litellm_params={}, + headers={}, + ) + for turn in (turn_n, turn_n_plus_one) + ] + + _assert_prefix_stable(requests) + assert [m["role"] for m in requests[0]["messages"]] == ["user", "assistant", "user"] + assert requests[0]["messages"][2]["content"][0]["type"] == "tool_result" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index d03174bc2c6..2fe22ba2620 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1046,11 +1046,26 @@ def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_azure(model: st assert openai_request["prompt_cache_key"] == "session-abc" +@pytest.mark.parametrize( + "model", + [ + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/moonshotai/kimi-k2-thinking-maas", + "vertex_ai/xai/grok-4.1-fast-non-reasoning", + ], +) +def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_vertex_maas_models(model: str): + openai_request = _translate_with_metadata(model, {"user_id": CLAUDE_CODE_USER_ID}, "vertex_ai") + assert openai_request["prompt_cache_key"] == "session-abc" + + @pytest.mark.parametrize( "model, custom_llm_provider", [ ("gemini/gemini-2.5-pro", "gemini"), ("vertex_ai/gemini-2.5-pro", "vertex_ai"), + ("vertex_ai/gemma/gemma-2-2b-it", "vertex_ai"), + ("vertex_ai/openai/mg-endpoint-lit8592", "vertex_ai"), ("anthropic/claude-sonnet-4-5", "anthropic"), ("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", "bedrock"), ("no-such-model-lit5875", "no-such-provider-lit5875"), diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py index 9e2bfb08852..ddbc168589a 100644 --- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py +++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py @@ -437,3 +437,51 @@ class TestAzureAnthropicConfig: assert "anthropic-beta" in headers assert "compact-2026-01-12" in headers["anthropic-beta"] assert "context-management-2025-06-27" in headers["anthropic-beta"] + + +def _mid_conversation_system_conversation() -> list[dict]: + return [ + {"role": "system", "content": [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}]}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + {"role": "system", "content": "Answer with exactly one word."}, + {"role": "assistant", "content": "Second answer"}, + {"role": "user", "content": "Third question"}, + ] + + +def test_chat_unflagged_model_converts_mid_conversation_system_instead_of_hoisting(local_model_cost_map): + """A hoisted reminder rewrites the top-level system block and invalidates the + prompt cache for the whole conversation (#36559).""" + result = AzureAnthropicConfig().transform_request( + model="claude-opus-4-7", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "assistant", "user"] + texts = [b["text"] for b in result["messages"][2]["content"] if b.get("type") == "text"] + assert texts[0] == "Second question" + assert texts[-1] == "Answer with exactly one word." + + +def test_chat_flagged_model_keeps_mid_conversation_system_role_in_place(local_model_cost_map): + result = AzureAnthropicConfig().transform_request( + model="claude-opus-4-8", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] + assert result["messages"][3] == { + "role": "system", + "content": [{"type": "text", "text": "Answer with exactly one word."}], + } + diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index ad77f9d4d1b..499096621c5 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -1,3 +1,4 @@ +import copy import json import os from typing import Final @@ -8,6 +9,7 @@ import pytest import litellm from litellm import ModelResponse +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import CONVERTED_SYSTEM_NOTE from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig from litellm.types.llms.bedrock import ConverseTokenUsageBlock @@ -7584,6 +7586,246 @@ def test_eager_input_streaming_non_boolean_is_a_bad_request(): ) +def test_mid_conversation_system_after_multiple_tool_results(): + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": "calling tools", + "tool_calls": [ + { + "id": "call_a", + "type": "function", + "function": {"name": "f", "arguments": "{}"}, + } + ], + }, + {"role": "system", "content": "reminder"}, + {"role": "tool", "tool_call_id": "call_a", "content": "r1"}, + {"role": "tool", "tool_call_id": "call_b", "content": "r2"}, + {"role": "user", "content": "done"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert system_blocks == [] + assert [m["role"] for m in out_messages] == [ + "user", + "assistant", + "tool", + "tool", + "user", + "user", + ] + assert out_messages[2]["content"] == "r1" + assert out_messages[3]["content"] == "r2" + # Reminder lands after ALL tool results, not between them. + assert out_messages[4]["content"][1]["text"] == "reminder" + assert out_messages[5]["content"] == "done" + + +def test_mid_conversation_system_reorders_around_a_pydantic_assistant_tool_call(): + config = AmazonConverseConfig() + assistant = litellm.Message( + role="assistant", + content="calling tools", + tool_calls=[{"id": "call_a", "type": "function", "function": {"name": "f", "arguments": "{}"}}], + ) + messages = [ + {"role": "user", "content": "hi"}, + assistant, + {"role": "system", "content": "reminder"}, + {"role": "tool", "tool_call_id": "call_a", "content": "r1"}, + {"role": "user", "content": "done"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert system_blocks == [] + assert [m["role"] for m in out_messages] == ["user", "assistant", "tool", "user", "user"] + assert out_messages[1] is assistant + assert out_messages[3]["content"][1]["text"] == "reminder" + + +def test_mid_conversation_multi_system_run_after_multiple_tool_results(): + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": "calling tools", + "tool_calls": [ + { + "id": "call_a", + "type": "function", + "function": {"name": "f", "arguments": "{}"}, + } + ], + }, + {"role": "system", "content": "reminder 1"}, + {"role": "system", "content": "reminder 2"}, + {"role": "tool", "tool_call_id": "call_a", "content": "r1"}, + {"role": "tool", "tool_call_id": "call_b", "content": "r2"}, + {"role": "user", "content": "done"}, + ] + out_messages, _ = config._transform_system_message(messages) + assert [m["role"] for m in out_messages] == [ + "user", + "assistant", + "tool", + "tool", + "user", + "user", + "user", + ] + assert out_messages[4]["content"][1]["text"] == "reminder 1" + assert out_messages[5]["content"][1]["text"] == "reminder 2" + + +def test_opens_with_tool_result_rejects_non_dict(): + config = AmazonConverseConfig() + assert config._opens_with_tool_result("not-a-dict") is False + assert config._opens_with_tool_result(None) is False + assert config._opens_with_tool_result([{"role": "tool"}]) is False + + +def test_mid_conversation_system_without_tools_stays_in_place(): + config = AmazonConverseConfig() + messages = [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "hello"}, + {"role": "system", "content": "reminder"}, + {"role": "user", "content": "thanks"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert [b["text"] for b in system_blocks if "text" in b] == ["You are helpful."] + assert [m["role"] for m in out_messages] == ["user", "assistant", "user", "user"] + assert out_messages[2]["content"][1]["text"] == "reminder" + assert out_messages[3]["content"] == "thanks" + + +def test_mid_conversation_system_str_with_cache_control(): + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "system", + "content": "reminder", + "cache_control": {"type": "ephemeral"}, + }, + {"role": "user", "content": "done"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert system_blocks == [] + assert out_messages[1]["role"] == "user" + assert out_messages[1]["content"][1] == { + "type": "text", + "text": "reminder", + "cache_control": {"type": "ephemeral"}, + } + + +def test_mid_conversation_system_list_content_with_cache_control(): + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "system", + "content": [ + {"type": "text", "text": "keep this", "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "plain"}, + {"type": "text", "text": ""}, + {"type": "image", "source": "x"}, + "raw-string", + ], + }, + {"role": "user", "content": "done"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert system_blocks == [] + blocks = out_messages[1]["content"] + assert blocks[0]["text"] == CONVERTED_SYSTEM_NOTE + assert blocks[1] == { + "type": "text", + "text": "keep this", + "cache_control": {"type": "ephemeral"}, + } + assert blocks[2] == {"type": "text", "text": "plain"} + assert len(blocks) == 3 + + +@pytest.mark.parametrize( + "empty_content", + ["", [], None, [{"type": "image", "source": "x"}, {"type": "text", "text": ""}]], + ids=["empty-string", "empty-list", "none", "no-text-parts"], +) +def test_mid_conversation_system_entry_without_text_is_dropped(empty_content): + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "hi"}, + {"role": "system", "content": empty_content}, + {"role": "user", "content": "done"}, + ] + out_messages, system_blocks = config._transform_system_message(messages) + assert system_blocks == [] + assert out_messages == [{"role": "user", "content": "hi"}, {"role": "user", "content": "done"}] + + +def _thinking_reply(text: str) -> dict: + return { + "role": "assistant", + "content": text, + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": f"sig-{text}"}], + } + + +def _preserved_thinking_turns(reminder_after_user: bool) -> tuple[list[dict], list[dict], list[dict]]: + turn_n = [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "First question"}] + reminder = {"role": "system", "content": "Answer with exactly one word."} + second_question = {"role": "user", "content": "Second question"} + second_turn = [second_question, reminder] if reminder_after_user else [reminder, second_question] + turn_n_plus_one = [*turn_n, _thinking_reply("First answer"), *second_turn] + turn_n_plus_two = [*turn_n_plus_one, _thinking_reply("Second answer"), {"role": "user", "content": "Third question"}] + return turn_n, turn_n_plus_one, turn_n_plus_two + + +def _replayed_prefix(request: dict, message_count: int) -> str: + replayed = { + "system": request.get("system"), + "toolConfig": request.get("toolConfig"), + "messages": request["messages"][:message_count], + } + return json.dumps(replayed, sort_keys=True) + + +def _assert_prefix_stable(requests: list[dict]) -> None: + for earlier, later in zip(requests, requests[1:]): + count = len(earlier["messages"]) + assert _replayed_prefix(later, count) == _replayed_prefix(earlier, count) + + +@pytest.mark.parametrize("reminder_after_user", [True, False]) +def test_flagged_model_replays_a_byte_identical_prefix_around_a_mid_conversation_reminder( + local_model_cost_map, reminder_after_user +): + """Converse rejects ``role: system`` inside ``messages``, so the reminder becomes a + user turn in place; hoisting it into ``system`` would change the prefix every + signed thinking block in the history is bound to.""" + requests = [ + AmazonConverseConfig().transform_request( + model="bedrock/us.anthropic.claude-fable-5-1", + messages=copy.deepcopy(turn), + optional_params={}, + litellm_params={}, + headers={}, + ) + for turn in _preserved_thinking_turns(reminder_after_user) + ] + + _assert_prefix_stable(requests) + assert requests[1]["system"] == [{"text": "You are terse."}] + assert [m["role"] for m in requests[1]["messages"]] == ["user", "assistant", "user"] + assert [m["role"] for m in requests[2]["messages"]] == ["user", "assistant", "user", "assistant", "user"] + + @pytest.mark.parametrize("model", ("anthropic.claude-opus-4-7", "us.anthropic.claude-opus-4-7")) def test_converse_accepts_anthropic_default_temperature(model: str) -> None: result: Final = litellm.utils.get_optional_params( diff --git a/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py b/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py index ac7cd24766d..6014e8a514e 100644 --- a/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py +++ b/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py @@ -15,6 +15,7 @@ from litellm.llms.fal_ai.image_generation import ( get_fal_ai_image_generation_config, ) from litellm.types.utils import ImageObject, ImageResponse +from litellm.utils import get_optional_params_image_gen @pytest.mark.parametrize( @@ -23,6 +24,8 @@ from litellm.types.utils import ImageObject, ImageResponse "fal-ai/nano-banana", "nano-banana", "fal-ai/gemini-25-flash-image", + "fal-ai/nano-banana-2", + "fal-ai/nano-banana-pro", ], ) def test_nano_banana_config_selected(model): @@ -145,3 +148,21 @@ def test_transform_request_includes_prompt_and_mapped_params(): } +@pytest.mark.parametrize("model", ["fal-ai/nano-banana-2", "fal-ai/nano-banana-pro"]) +def test_resolution_extra_param_is_forwarded_to_fal(model): + optional_params = get_optional_params_image_gen( + model=model, + n=1, + size="1024x1024", + custom_llm_provider="fal_ai", + provider_config=FalAINanoBananaConfig(), + resolution="4K", + ) + request = FalAINanoBananaConfig().transform_image_generation_request( + model=model, + prompt="a cat", + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + assert request == {"prompt": "a cat", "num_images": 1, "aspect_ratio": "1:1", "resolution": "4K"} diff --git a/tests/test_litellm/llms/fal_ai/test_cost_calculator.py b/tests/test_litellm/llms/fal_ai/test_cost_calculator.py index 3c6e6aea090..35eec247f7e 100644 --- a/tests/test_litellm/llms/fal_ai/test_cost_calculator.py +++ b/tests/test_litellm/llms/fal_ai/test_cost_calculator.py @@ -242,3 +242,77 @@ def test_passthrough_cost_is_none_only_when_no_price_applies_to_the_request(monk assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {}) is None assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {"resolution": 1024}) is None assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {"resolution": "512"}) == 0.02 + + +NANO_BANANA_RESOLUTION_MODELS: Final = ("fal-ai/nano-banana-2", "fal-ai/nano-banana-pro") + + +@pytest.mark.parametrize("model", NANO_BANANA_RESOLUTION_MODELS) +def test_nano_banana_default_request_charges_the_1k_rate_per_image(model): + entry: Final = litellm.model_cost[f"fal_ai/{model}"] + cost: Final = cost_calculator( + model=f"fal_ai/{model}", + image_response=_image_response(num_images=2), + optional_params={"num_images": 2, "aspect_ratio": "1:1"}, + ) + assert cost == 2 * entry["output_cost_per_image"] == 2 * entry["output_cost_per_image_1K"] > 0 + + +@pytest.mark.parametrize("model", NANO_BANANA_RESOLUTION_MODELS) +def test_nano_banana_4k_request_charges_the_4k_rate_above_1k(model): + one_k: Final = cost_calculator( + model=f"fal_ai/{model}", image_response=_image_response(), optional_params={"resolution": "1K"} + ) + four_k: Final = cost_calculator( + model=f"fal_ai/{model}", image_response=_image_response(), optional_params={"resolution": "4K"} + ) + assert four_k == litellm.model_cost[f"fal_ai/{model}"]["output_cost_per_image_4K"] + assert four_k > one_k > 0 + + +@pytest.mark.parametrize("model", NANO_BANANA_RESOLUTION_MODELS) +@pytest.mark.parametrize("resolution", ("1K", "2K", "4K")) +def test_nano_banana_images_generations_and_passthrough_charge_the_same_tier(model, resolution): + images_generations_cost: Final = cost_calculator( + model=f"fal_ai/{model}", image_response=_image_response(), optional_params={"resolution": resolution} + ) + assert images_generations_cost == fal_ai_passthrough_cost(model, {"resolution": resolution}) > 0 + + +def test_nano_banana_2_resolution_tiers_are_monotonic(): + costs: Final = tuple( + cost_calculator( + model="fal_ai/fal-ai/nano-banana-2", + image_response=_image_response(), + optional_params={"resolution": resolution}, + ) + for resolution in ("0.5K", "1K", "2K", "4K") + ) + assert costs == tuple(sorted(costs)) and len(set(costs)) == len(costs) + + +@pytest.mark.parametrize("model", NANO_BANANA_RESOLUTION_MODELS) +def test_nano_banana_unpriced_resolution_falls_back_to_the_default_rate(model): + cost: Final = cost_calculator( + model=f"fal_ai/{model}", image_response=_image_response(), optional_params={"resolution": "8K"} + ) + assert cost == litellm.model_cost[f"fal_ai/{model}"]["output_cost_per_image"] > 0 + + +@pytest.mark.parametrize("model", NANO_BANANA_RESOLUTION_MODELS) +def test_passthrough_num_images_multiplies_the_per_image_rate(model): + entry: Final = litellm.model_cost[f"fal_ai/{model}"] + assert fal_ai_passthrough_cost(model, {"num_images": 3}) == 3 * entry["output_cost_per_image"] > 0 + assert ( + fal_ai_passthrough_cost(model, {"resolution": "4K", "num_images": 2}) == 2 * entry["output_cost_per_image_4K"] > 0 + ) + + +@pytest.mark.parametrize("num_images", (None, 0, -2, True, 2.0, "2")) +def test_passthrough_without_a_positive_integer_num_images_charges_one_image(num_images): + body: Final = {} if num_images is None else {"num_images": num_images} + assert ( + fal_ai_passthrough_cost("fal-ai/nano-banana-2", body) + == litellm.model_cost["fal_ai/fal-ai/nano-banana-2"]["output_cost_per_image"] + > 0 + ) diff --git a/tests/test_litellm/llms/nadir/test_nadir.py b/tests/test_litellm/llms/nadir/test_nadir.py new file mode 100644 index 00000000000..2b8b8387a42 --- /dev/null +++ b/tests/test_litellm/llms/nadir/test_nadir.py @@ -0,0 +1,260 @@ +import json +import math +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +import litellm +from litellm import get_llm_provider +from litellm.types.utils import ModelResponse, Usage + +NADIR_BASE = "https://api.getnadir.com/v1" +COST_HEADER = "llm_provider-x-litellm-response-cost" + + +def _transform(payload): + raw = httpx.Response( + 200, + content=json.dumps(payload).encode(), + headers={"content-type": "application/json"}, + request=httpx.Request("POST", f"{NADIR_BASE}/chat/completions"), + ) + return litellm.NadirConfig().transform_response( + model="auto", + raw_response=raw, + model_response=ModelResponse(), + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def _payload(**extra): + return { + "id": "req-1", + "object": "chat.completion", + "created": 0, + "model": "claude-haiku-4-5", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + **extra, + } + + +def _cost(response, provider): + return litellm.completion_cost(completion_response=response, custom_llm_provider=provider) + + +def _logged_cost(response): + return litellm.response_cost_calculator( + response_object=response, + model="auto", + custom_llm_provider="nadir", + call_type="completion", + optional_params={}, + ) + + +class TestNadirProviderResolution: + def test_model_prefix_resolves_to_nadir(self): + model, provider, _, _ = get_llm_provider(model="nadir/auto", api_key="sk-test") + assert (model, provider) == ("auto", "nadir") + + def test_default_api_base(self): + _, _, _, api_base = get_llm_provider(model="nadir/auto", api_key="sk-test") + assert api_base == NADIR_BASE + + def test_api_base_override(self): + _, _, _, api_base = get_llm_provider( + model="nadir/auto", + api_key="sk-test", + api_base="https://gateway.internal/v1", + ) + assert api_base == "https://gateway.internal/v1" + + def test_endpoint_reverse_maps_to_nadir_with_the_env_key(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, provider, dynamic_api_key, _ = get_llm_provider(model="auto", api_base=NADIR_BASE) + assert (provider, dynamic_api_key) == ("nadir", "sk-server-secret") + + def test_plaintext_endpoint_never_loads_the_env_key(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, provider, dynamic_api_key, _ = get_llm_provider(model="auto", api_base="http://api.getnadir.com/v1") + assert provider == "nadir" + assert dynamic_api_key is None + + +class TestNadirCredentialScoping: + def test_env_key_used_for_default_endpoint(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, _, dynamic_api_key, _ = get_llm_provider(model="nadir/auto") + assert dynamic_api_key == "sk-server-secret" + + def test_env_key_used_when_base_matches_default(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, _, dynamic_api_key, _ = get_llm_provider(model="nadir/auto", api_base=f"{NADIR_BASE}/") + assert dynamic_api_key == "sk-server-secret" + + def test_env_key_not_leaked_to_custom_base(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, _, dynamic_api_key, _ = get_llm_provider(model="nadir/auto", api_base="https://attacker.example/v1") + assert dynamic_api_key is None + + def test_caller_key_used_for_custom_base(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + _, _, dynamic_api_key, _ = get_llm_provider( + model="nadir/auto", + api_base="https://self-hosted.internal/v1", + api_key="sk-caller-own", + ) + assert dynamic_api_key == "sk-caller-own" + + def test_env_key_used_for_operator_configured_base(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-server-secret") + monkeypatch.setenv("NADIR_API_BASE", "https://nadir.mycorp.internal/v1") + _, _, dynamic_api_key, _ = get_llm_provider(model="nadir/auto", api_base="https://nadir.mycorp.internal/v1") + assert dynamic_api_key == "sk-server-secret" + + +class TestNadirParamMapping: + def test_supported_params_are_mapped(self): + params = litellm.get_optional_params( + model="auto", + custom_llm_provider="nadir", + temperature=0.5, + max_tokens=64, + ) + assert params["temperature"] == 0.5 + assert params["max_tokens"] == 64 + + def test_streaming_is_advertised_and_tools_are_not(self): + params = litellm.get_supported_openai_params(model="auto", custom_llm_provider="nadir") + assert "stream" in params + assert "tools" not in params + + @pytest.mark.parametrize( + "unsupported", + [ + {"tools": [{"type": "function", "function": {"name": "f", "parameters": {}}}]}, + {"stop": ["\n"]}, + {"seed": 7}, + {"n": 2}, + ], + ) + def test_params_nadir_would_silently_drop_are_rejected(self, unsupported): + with pytest.raises(litellm.UnsupportedParamsError): + litellm.get_optional_params(model="auto", custom_llm_provider="nadir", **unsupported) + + def test_unsupported_params_are_dropped_when_asked(self): + params = litellm.get_optional_params( + model="auto", + custom_llm_provider="nadir", + drop_params=True, + seed=7, + temperature=0.2, + ) + assert "seed" not in params + assert params["temperature"] == 0.2 + + +class TestNadirEnvValidation: + def test_validate_environment_detects_key(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-live-xyz") + result = litellm.validate_environment(model="nadir/auto") + assert result["keys_in_environment"] is True + + def test_validate_environment_flags_missing_key(self, monkeypatch): + monkeypatch.delenv("NADIR_API_KEY", raising=False) + result = litellm.validate_environment(model="nadir/auto") + assert "NADIR_API_KEY" in result["missing_keys"] + + +class TestNadirCostAttribution: + def test_reported_cost_wins_over_model_pricing(self): + res = _transform(_payload(nadir_metadata={"cost": {"total_cost_usd": 0.00123}})) + assert _logged_cost(res) == pytest.approx(0.00123) + assert _logged_cost(res) != _cost(res, "anthropic") + + def test_routed_model_is_preserved(self): + res = _transform(_payload(nadir_metadata={"cost": {"total_cost_usd": 0.001}})) + assert res.model == "claude-haiku-4-5" + + def test_missing_cost_prices_the_routed_model_from_its_own_entry(self): + res = _transform(_payload()) + assert res.choices[0].message.content == "hi" + assert COST_HEADER not in res._hidden_params.get("additional_headers", {}) + assert _cost(res, "nadir") == _logged_cost(res) == _cost(res, "anthropic") > 0 + + def test_streamed_routed_model_prices_from_its_own_entry(self): + res = ModelResponse( + model="gemini-3.5-flash-lite", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + assert _cost(res, "nadir") == _cost(res, "gemini") > 0 + + @pytest.mark.parametrize("bad", [-0.001, math.nan, math.inf, -math.inf, True, "0.001", None]) + def test_invalid_reported_cost_falls_back_to_model_pricing(self, bad): + res = _transform(_payload(nadir_metadata={"cost": {"total_cost_usd": bad}})) + assert COST_HEADER not in res._hidden_params.get("additional_headers", {}) + assert _cost(res, "nadir") == _cost(res, "anthropic") > 0 + + @pytest.mark.parametrize("metadata", ["oops", {"cost": "free"}, {"cost": None}, {}]) + def test_malformed_metadata_falls_back_to_model_pricing(self, metadata): + res = _transform(_payload(nadir_metadata=metadata)) + assert COST_HEADER not in res._hidden_params.get("additional_headers", {}) + assert _cost(res, "nadir") == _cost(res, "anthropic") > 0 + + +class TestNadirCompletionDispatch: + def _call(self, **kwargs): + captured = {} + + def fake_completion(**call_kwargs): + captured.update(call_kwargs) + return ModelResponse() + + with patch( # test-quality-ok: these tests assert the dispatch wiring itself (nadir must reach base_llm_http_handler, and which credentials it is handed); faking HTTP would not observe that + "litellm.main.base_llm_http_handler.completion", side_effect=fake_completion + ): + litellm.completion( + model="nadir/auto", + messages=[{"role": "user", "content": "hi"}], + **kwargs, + ) + return captured + + def test_routes_through_the_http_handler_as_nadir(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-env") + captured = self._call() + assert captured["custom_llm_provider"] == "nadir" + assert captured["api_base"] == NADIR_BASE + + def test_env_key_is_used_for_the_default_endpoint(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-env") + assert self._call()["api_key"] == "sk-env" + + def test_caller_key_wins(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-env") + assert self._call(api_key="sk-caller")["api_key"] == "sk-caller" + + def test_env_key_is_not_forwarded_to_a_caller_supplied_host(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-env") + captured = self._call(api_base="https://attacker.example/v1") + assert captured["api_key"] != "sk-env" + assert captured["api_base"] == "https://attacker.example/v1" + + def test_global_key_is_not_forwarded_to_a_caller_supplied_host(self, monkeypatch): + monkeypatch.delenv("NADIR_API_KEY", raising=False) + monkeypatch.setattr(litellm, "api_key", "sk-global") + captured = self._call(api_base="https://attacker.example/v1") + assert captured["api_key"] != "sk-global" + + def test_custom_api_base_is_honoured(self, monkeypatch): + monkeypatch.setenv("NADIR_API_KEY", "sk-env") + captured = self._call(api_base="https://nadir.internal/v1", api_key="sk-own") + assert captured["api_base"] == "https://nadir.internal/v1" + assert captured["api_key"] == "sk-own" diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py index 034f85f5a0b..f3a276f7e46 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from types import SimpleNamespace import pytest @@ -6,6 +7,7 @@ from litellm.exceptions import BadRequestError from litellm.llms.vertex_ai.vector_stores.search_api.transformation import ( VertexSearchAPIVectorStoreConfig, ) +from litellm.types.vector_stores import VectorStoreSearchResponse def test_should_encode_vertex_search_vector_store_id_in_complete_url(): @@ -297,3 +299,188 @@ def test_search_request_logs_effective_query_when_extra_body_overrides_query(): assert body["query"] == "from-extra-body" assert log.model_call_details["query"] == "from-extra-body" + + +_CHUNK_NAME = ( + "projects/p/locations/global/collections/default_collection/dataStores/ds-1/" + "branches/0/documents/policy/chunks/c3" +) + + +def _search_response(payload: Mapping[str, object]) -> VectorStoreSearchResponse: + return VertexSearchAPIVectorStoreConfig().transform_search_vector_store_response( + response=SimpleNamespace(json=lambda: payload, status_code=200, headers={}), + litellm_logging_obj=SimpleNamespace(model_call_details={"query": "hello"}), + ) + + +def test_chunk_hit_uses_chunk_content_and_document_metadata(): + payload = { + "results": [ + { + "chunk": { + "id": "c3", + "name": _CHUNK_NAME, + "content": "Refunds are available within 14 days.", + "documentMetadata": { + "uri": "gs://bucket/policy.pdf", + "title": "Refund policy", + "structData": {"department": "billing"}, + }, + "pageSpan": {"pageStart": 2, "pageEnd": 2}, + "relevanceScore": 0.91, + } + } + ] + } + + result = _search_response(payload)["data"][0] + + assert result["content"] == [ + {"text": "Refunds are available within 14 days.", "type": "text"} + ] + assert result["score"] == 0.91 + assert result["file_id"] == "gs://bucket/policy.pdf" + assert result["filename"] == "Refund policy" + assert result["attributes"] == { + "document_id": "policy", + "chunk_id": "c3", + "link": "gs://bucket/policy.pdf", + "title": "Refund policy", + "structData": {"department": "billing"}, + "pageSpan": {"pageStart": 2, "pageEnd": 2}, + } + + +def test_chunk_hit_without_uri_or_title_falls_back_to_document_id(): + payload = { + "results": [ + { + "chunk": { + "id": "c1", + "name": _CHUNK_NAME.replace("policy/chunks/c3", "handbook/chunks/c1"), + "content": "Guest Services Handbook", + "documentMetadata": {"structData": {"title": "Handbook"}}, + } + } + ] + } + + result = _search_response(payload)["data"][0] + + assert result["content"] == [{"text": "Guest Services Handbook", "type": "text"}] + assert result["score"] == 1.0 + assert result["file_id"] == "handbook" + assert result["filename"] == "Unknown Document" + assert result["attributes"] == { + "document_id": "handbook", + "chunk_id": "c1", + "structData": {"title": "Handbook"}, + } + + +@pytest.mark.parametrize( + ("derived", "expected_text"), + [ + ( + { + "extractive_segments": [{"content": "seg one"}, {"content": "seg two"}], + "extractive_answers": [{"content": "ans"}], + "snippets": [{"snippet": "snip"}], + "title": "policy.pdf", + }, + "seg one\n\nseg two", + ), + ( + { + "extractive_answers": [{"content": "ans one"}, {"content": "ans two"}], + "snippets": [{"snippet": "snip"}], + "title": "policy.pdf", + }, + "ans one\n\nans two", + ), + ( + { + "snippets": [{"snippet": "snip a"}, {"htmlSnippet": "snip b"}], + "title": "policy.pdf", + }, + "snip a snip b", + ), + ( + {"extractive_segments": [{"pageNumber": "1"}, {"content": "seg", "pageNumber": "2"}]}, + "seg", + ), + ({"title": "policy.pdf"}, "policy.pdf"), + ], + ids=["segments", "answers", "snippets", "content_less_segment", "title"], +) +def test_document_hit_text_prefers_extractive_content(derived: Mapping[str, object], expected_text: str) -> None: + payload = {"results": [{"id": "doc-1", "document": {"derivedStructData": derived}}]} + + result = _search_response(payload)["data"][0] + + assert result["content"] == [{"text": expected_text, "type": "text"}] + + +def test_document_hit_surfaces_struct_data_in_attributes(): + payload = { + "results": [ + { + "id": "attr-1", + "document": { + "structData": {"title": "Thunder Loop", "waitMinutes": 45}, + "derivedStructData": {"clearbox_escorer_score": 0.5}, + }, + }, + {"id": "attr-2", "document": {"structData": {}, "derivedStructData": {}}}, + ] + } + + first, second = _search_response(payload)["data"] + + assert first["content"] == [{"text": "", "type": "text"}] + assert first["file_id"] == "attr-1" + assert first["filename"] == "Unknown Document" + assert first["attributes"] == { + "document_id": "attr-1", + "structData": {"title": "Thunder Loop", "waitMinutes": 45}, + } + assert second["attributes"] == {"document_id": "attr-2"} + + +def test_search_response_keeps_link_metadata_and_positional_scores(): + payload = { + "results": [ + { + "id": "doc-1", + "document": { + "derivedStructData": { + "title": "Terms", + "link": "gs://bucket/terms.pdf", + "displayLink": "bucket", + "formattedUrl": "https://bucket/terms.pdf", + "snippets": [{"snippet": "snip"}], + } + }, + }, + {"chunk": {"name": _CHUNK_NAME, "content": "chunk text"}}, + ] + } + + response = _search_response(payload) + + assert response["object"] == "vector_store.search_results.page" + assert response["search_query"] == "hello" + assert [result["score"] for result in response["data"]] == [1.0, 0.5] + assert response["data"][0]["file_id"] == "gs://bucket/terms.pdf" + assert response["data"][0]["attributes"] == { + "document_id": "doc-1", + "link": "gs://bucket/terms.pdf", + "title": "Terms", + "displayLink": "bucket", + "formattedUrl": "https://bucket/terms.pdf", + } + + +def test_search_response_without_results_key_is_empty(): + assert _search_response({"totalSize": 0})["data"] == [] diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py index 0dcaa4c72c2..b80d4714253 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -8,10 +8,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm -from litellm.llms.vertex_ai.vertex_model_garden.main import ( - _vertex_model_garden_model_id_in_json_body, - create_vertex_url, +from litellm.llms.vertex_ai.common_utils import ( + vertex_model_garden_model_id_in_json_body, ) +from litellm.llms.vertex_ai.vertex_model_garden.main import create_vertex_url @pytest.mark.parametrize( @@ -43,11 +43,8 @@ def test_create_vertex_url_openapi_vs_deployed_endpoint( def test_model_id_in_json_body_heuristic() -> None: - assert ( - _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") - is True - ) - assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False + assert vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert vertex_model_garden_model_id_in_json_body("5464397967697903616") is False @pytest.fixture diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py index 37a619d6400..ca4a3dedb4a 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py @@ -1,4 +1,7 @@ +import copy +import json + import pytest from litellm.anthropic_beta_headers_manager import ( @@ -771,3 +774,104 @@ def test_vertex_ai_anthropic_tool_based_response_format_still_upgrades_legacy_th assert "tools" in result_params assert result_params["thinking"] == {"type": "adaptive"} assert result_params["output_config"] == {"effort": "high"} + + + + +def _mid_conversation_system_conversation() -> list[dict]: + return [ + {"role": "system", "content": [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}]}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + {"role": "system", "content": "Answer with exactly one word."}, + {"role": "assistant", "content": "Second answer"}, + {"role": "user", "content": "Third question"}, + ] + + +def test_chat_unflagged_model_converts_mid_conversation_system_instead_of_hoisting(local_model_cost_map): + """A hoisted reminder rewrites the top-level system block and invalidates the + prompt cache for the whole conversation (#36559).""" + result = VertexAIAnthropicConfig().transform_request( + model="claude-opus-4-7", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "assistant", "user"] + texts = [b["text"] for b in result["messages"][2]["content"] if b.get("type") == "text"] + assert texts[0] == "Second question" + assert texts[-1] == "Answer with exactly one word." + + +def test_chat_flagged_model_keeps_mid_conversation_system_role_in_place(local_model_cost_map): + result = VertexAIAnthropicConfig().transform_request( + model="claude-opus-4-8", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] + assert result["messages"][3] == { + "role": "system", + "content": [{"type": "text", "text": "Answer with exactly one word."}], + } + + +def _thinking_reply(text: str) -> dict: + return { + "role": "assistant", + "content": text, + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": f"sig-{text}"}], + } + + +def _preserved_thinking_turns(reminder_after_user: bool) -> tuple[list[dict], list[dict], list[dict]]: + turn_n = [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "First question"}] + reminder = {"role": "system", "content": "Answer with exactly one word."} + second_question = {"role": "user", "content": "Second question"} + second_turn = [second_question, reminder] if reminder_after_user else [reminder, second_question] + turn_n_plus_one = [*turn_n, _thinking_reply("First answer"), *second_turn] + turn_n_plus_two = [*turn_n_plus_one, _thinking_reply("Second answer"), {"role": "user", "content": "Third question"}] + return turn_n, turn_n_plus_one, turn_n_plus_two + + +def _replayed_prefix(request: dict, message_count: int) -> str: + replayed = { + "system": request.get("system"), + "tools": request.get("tools"), + "messages": request["messages"][:message_count], + } + return json.dumps(replayed, sort_keys=True) + + +def _assert_prefix_stable(requests: list[dict]) -> None: + for earlier, later in zip(requests, requests[1:]): + count = len(earlier["messages"]) + assert _replayed_prefix(later, count) == _replayed_prefix(earlier, count) + + +@pytest.mark.parametrize("reminder_after_user", [True, False]) +def test_chat_flagged_model_replays_a_byte_identical_prefix_around_a_mid_conversation_reminder( + local_model_cost_map, reminder_after_user +): + """Preserved thinking binds each signed block to the request prefix it was created + under (``system``, ``tools`` and the earlier messages), so turn N's transformed + request must be a byte-identical prefix of turn N+1's or the block is dropped.""" + requests = [ + VertexAIAnthropicConfig().transform_request( + model="claude-fable-5-1", messages=copy.deepcopy(turn), optional_params={}, litellm_params={}, headers={} + ) + for turn in _preserved_thinking_turns(reminder_after_user) + ] + + _assert_prefix_stable(requests) + assert [m["role"] for m in requests[1]["messages"]] == ["user", "assistant", "user", "system"] + assert [m["role"] for m in requests[2]["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] 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/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 24f09e79dbb..fc4d7b45785 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -342,6 +342,15 @@ class TestMCPRequestHandler: mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) return mock_manager + def _real_manager_with_toolsets(self, toolset_perms): + """A real MCPServerManager so the real expand_tool_permissions runs; + only the DB-backed toolset lookup is stubbed""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager = MCPServerManager() + manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) + return manager + async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self): """A key granted only mcp_toolsets must reach the toolset's servers on every path (list, call, REST); regression for the list-ok/call-403 bug""" @@ -508,6 +517,141 @@ class TestMCPRequestHandler: assert result is None + @pytest.mark.parametrize( + "direct,via_toolsets,expected", + [ + (["*"], None, None), + (["*"], ["read_file"], None), + (None, None, None), + ([], None, ()), + (None, ["read_file"], ("read_file",)), + ], + ) + def test_union_tool_grants_wildcard_and_union_cases(self, direct, via_toolsets, expected): + """A direct ["*"] makes the level unrestricted even beside a toolset + list (regression: mapping ["*"] to None in expand_tool_permissions let + a same-level toolset list deny every other tool)""" + result = MCPRequestHandler._union_tool_grants(direct, via_toolsets) + + if expected is None: + assert result is None + else: + assert result is not None + assert set(result) == set(expected) + + def test_union_tool_grants_unions_two_concrete_lists(self): + result = MCPRequestHandler._union_tool_grants(["read_file"], ["write_file"]) + + assert result is not None + assert set(result) == {"read_file", "write_file"} + + async def test_key_wildcard_allows_a_tool_never_enumerated(self): + """End to end at the key level: object_permission sits on the auth + object already, no team named, so no patching is needed; the real + global manager expands ["*"] and the level reads unrestricted""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed is None + assert brand_new_tool_allowed is True + + async def test_key_wildcard_stays_capped_by_team_allowlist(self): + """A wildcard on the key must never widen a team's enumerated ceiling: + the intersection keeps only the team's named tools""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["read_file"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + assert brand_new_tool_allowed is False + + async def test_team_wildcard_stays_capped_by_key_allowlist(self): + """A wildcard on the team leaves the key's enumerated list as the + effective ceiling""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["read_file"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["*"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + + async def test_key_empty_tool_list_stays_deny_all(self): + """[] on the key is deny-all, distinct from the wildcard: it must not + be widened into allow-all""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": []}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + read_file_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="read_file", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == [] + assert read_file_allowed is False + # ------------------------------------------------------------------ # LIT-5749: toolsets attached to a TEAM, ORG, or internal USER must be # enforced exactly like inline tool allowlists, on both axes diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py index 8ec5b8642bc..0ec4b431276 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py @@ -16,7 +16,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache class _FakeRedisCache: namespace = None - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return object() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 9659eb1cbc2..d39ec063538 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -17,6 +17,7 @@ from typing import Any, Dict, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from mcp.types import CallToolResult from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.proxy._types import UserAPIKeyAuth @@ -409,7 +410,7 @@ class TestCallToolFlowsHookHeaders: manager, "_call_openapi_tool_handler", new_callable=AsyncMock, - return_value=MagicMock(), + return_value=CallToolResult(content=[], isError=False), ): import litellm.proxy._experimental.mcp_server.mcp_server_manager as mgr_mod @@ -456,7 +457,7 @@ class TestCallToolFlowsHookHeaders: manager, "_call_openapi_tool_handler", new_callable=AsyncMock, - return_value=MagicMock(), + return_value=CallToolResult(content=[], isError=False), ): proxy_logging = MagicMock(spec=ProxyLogging) @@ -1076,9 +1077,9 @@ class TestOpenApiByokCallTool: user_auth = UserAPIKeyAuth(user_id="default_user_id", api_key="sk-dashboard") captured_auth: dict[str, Optional[str]] = {} - async def fake_openapi_handler(_server, _name, _arguments): + async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): captured_auth["value"] = _request_auth_header.get() - return MagicMock() + return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): with patch( @@ -1316,9 +1317,9 @@ class TestOpenApiResolvedUpstreamAuth: user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user") captured: Dict[str, Any] = {} - async def fake_openapi_handler(_server, _name, _arguments): + async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): captured["resolved"] = _request_resolved_auth_headers.get() - return MagicMock() + return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): with patch.object( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py index e11897b65c2..0dc7ac5ecd9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py @@ -2,6 +2,7 @@ import asyncio from typing import Dict, Optional import pytest +from mcp.types import CallToolResult, TextContent from unittest.mock import patch from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager @@ -46,7 +47,7 @@ def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPS def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker): async def fake_create_mcp_client(server, **kwargs): class _ProbeClient: - async def call_tool(self, params, host_progress_callback=None): + async def call_tool(self, params, host_progress_callback=None, allow_input_required=False): tracker.enter(server.server_id) try: await asyncio.sleep(HOLD_SECONDS) @@ -145,11 +146,11 @@ async def test_openapi_backed_server_also_respects_the_cap(): server = _make_server("srv-openapi", max_concurrent_requests=2) server.spec_path = "/fake/openapi.json" - async def fake_openapi_handler(mcp_server, name, arguments): + async def fake_openapi_handler(mcp_server, name, arguments, wire_compat): tracker.enter(mcp_server.server_id) try: await asyncio.sleep(HOLD_SECONDS) - return "ok" + return CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) finally: tracker.exit(mcp_server.server_id) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 6887adf8283..97b242831a2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -12,8 +12,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException from mcp import ReadResourceResult, Resource +from mcp.server.models import InitializationOptions from mcp.types import ( INVALID_REQUEST, + METHOD_NOT_FOUND, BlobResourceContents, CallToolResult, Prompt, @@ -21,7 +23,7 @@ from mcp.types import ( TextContent, TextResourceContents, ) -from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION +from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS from pydantic import TypeAdapter from starlette.types import Message, Receive, Scope, Send @@ -6565,8 +6567,11 @@ class TestGatewayCreateInitializationOptions: async def connect_sse(scope, receive, send): yield (None, None) - async def record_request(read_stream, write_stream, options): - captured["server_name"] = server.create_initialization_options().server_name + async def record_request( + serving_server: object, read_stream: object, write_stream: object, + *, lifespan_state: object, init_options: InitializationOptions, + ) -> None: + captured["server_name"] = init_options.server_name scope = { "type": "http", @@ -6612,8 +6617,8 @@ class TestGatewayCreateInitializationOptions: True, ), patch.object( - mcp_server.server, - "run", + mcp_server, + "serve_loop", side_effect=record_request, ), ): @@ -8021,7 +8026,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): ), patch( "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", - new=AsyncMock(return_value=[]), + new=AsyncMock(return_value=CallToolResult(content=[], is_error=False)), ), patch( "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", @@ -9255,6 +9260,152 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error(): proxy_logging_mock.post_call_failure_hook.assert_not_awaited() +def _interim_input_required_result(): + from mcp.types import InputRequiredResult + + return InputRequiredResult.model_validate( + { + "resultType": "input_required", + "inputRequests": { + "req-1": { + "method": "elicitation/create", + "params": {"message": "Pick one", "requestedSchema": {"type": "object", "properties": {}}}, + } + }, + "requestState": "state-1", + } + ) + + +@contextlib.contextmanager +def _managed_tool_returning(server, upstream_result, proxy_logging_mock): + from litellm.proxy._experimental.mcp_server.server import global_mcp_server_manager + + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[server.server_id], + ), + patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server), + patch.object(global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(global_mcp_server_manager, "server_owning_tool_name_prefix", return_value=server), + patch( + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", + new_callable=AsyncMock, + return_value=[server], + ), + patch( + "litellm.proxy._experimental.mcp_server.operations._list_tools_before_first_call", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.operations._handle_managed_mcp_tool", + new_callable=AsyncMock, + return_value=upstream_result, + ) as managed_call, + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock), + ): + yield managed_call + + +@pytest.mark.asyncio +async def test_call_mcp_tool_legacy_interim_result_is_rejected_into_failure_accounting(): + """An upstream input_required interim on a legacy connection cannot be carried on the wire, so it + must come back as isError and go through the same failure accounting as any other errored call.""" + from mcp.types import CallToolResult + + from litellm.proxy._experimental.mcp_server.result_conversion import ( + INPUT_REQUIRED_UNSUPPORTED_MESSAGE, + WireCompat, + ) + from litellm.proxy._experimental.mcp_server.server import call_mcp_tool + from litellm.proxy._types import MCPTransport, UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="server-interim", + name="test_server", + alias="test_server", + server_name="test_server", + url="https://test-server.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "test_server"}, + ) + proxy_logging_mock = _mock_mcp_proxy_logging() + logging_obj = _mock_mcp_logging_obj() + + with _managed_tool_returning(server, _interim_input_required_result(), proxy_logging_mock) as managed_call: + result = await call_mcp_tool( + name="test_server-any_tool", + arguments={"x": 1}, + user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + litellm_logging_obj=logging_obj, + wire_compat=WireCompat.LEGACY, + ) + + assert managed_call.await_args.kwargs["wire_compat"] is WireCompat.LEGACY + assert isinstance(result, CallToolResult) and result.is_error is True + assert result.content[0].text == INPUT_REQUIRED_UNSUPPORTED_MESSAGE + logging_obj.async_success_handler.assert_not_awaited() + logging_obj.async_failure_handler.assert_awaited_once() + assert str(logging_obj.async_failure_handler.await_args.args[0]) == INPUT_REQUIRED_UNSUPPORTED_MESSAGE + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_call_mcp_tool_modern_interim_result_passes_through_without_completed_accounting(): + """On a modern connection the interim result is returned with its fields intact and is neither + logged as a completed success nor run through the post-call guardrail and success hooks.""" + from mcp.types import InputRequiredResult + + from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat + from litellm.proxy._experimental.mcp_server.server import call_mcp_tool + from litellm.proxy._types import MCPTransport, UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="server-interim", + name="test_server", + alias="test_server", + server_name="test_server", + url="https://test-server.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "test_server"}, + ) + proxy_logging_mock = _mock_mcp_proxy_logging() + logging_obj = _mock_mcp_logging_obj() + interim = _interim_input_required_result() + + with _managed_tool_returning(server, interim, proxy_logging_mock) as managed_call: + result = await call_mcp_tool( + name="test_server-any_tool", + arguments={"x": 1}, + user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + litellm_logging_obj=logging_obj, + wire_compat=WireCompat.MODERN, + ) + + assert managed_call.await_args.kwargs["wire_compat"] is WireCompat.MODERN + assert isinstance(result, InputRequiredResult) + assert result.request_state == "state-1" + assert result.input_requests is not None and set(result.input_requests) == {"req-1"} + logging_obj.async_success_handler.assert_not_awaited() + logging_obj.async_failure_handler.assert_not_awaited() + logging_obj.async_post_mcp_tool_call_hook.assert_not_awaited() + proxy_logging_mock.post_mcp_call_hook.assert_not_awaited() + proxy_logging_mock.post_call_failure_hook.assert_not_awaited() + assert sorted(c.kwargs["event_type"] for c in logging_obj.has_run_logging.call_args_list) == [ + "async_success", + "sync_success", + ], "the @client wrapper would otherwise log the interim result as a completed success on return" + + @pytest.mark.asyncio async def test_aggregate_listing_reports_per_server_outcomes(): """A failed server must contribute a classified outcome, not just silently shrink the list: @@ -10371,7 +10522,10 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai @pytest.mark.asyncio @pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/"))) -async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str) -> None: +@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS)) +async def test_legacy_sse_mount_emits_message_endpoint( + prefix: str, suffix: str, opening_protocol: str | None, +) -> None: from starlette.applications import Starlette from starlette.routing import Mount from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing @@ -10434,6 +10588,23 @@ async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str) await asyncio.wait_for(app(post_scope, requests.get, messages.put), 2) return (await messages.get())["status"] + if opening_protocol is not None: + discover: Final = json.dumps({ + "jsonrpc": "2.0", + "id": 0, + "method": "server/discover", + "params": {"_meta": { + "io.modelcontextprotocol/protocolVersion": opening_protocol, + "io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"}, + "io.modelcontextprotocol/clientCapabilities": {}, + }}, + }).encode() + assert await post(discover) == 202 + discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode() + discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0]) + assert discovered["id"] == 0 + assert discovered["error"]["code"] == METHOD_NOT_FOUND + initialization: Final = json.dumps( { "jsonrpc": "2.0", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cd5dae1269a..64d94065674 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -39,6 +39,7 @@ from mcp.types import Tool as MCPTool from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT +from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _deserialize_json_dict, @@ -6730,7 +6731,7 @@ class TestMCPServerManager: # Create mock client that tracks call_tool usage mock_client = AsyncMock() - async def mock_call_tool(params, host_progress_callback=None): + async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False): # Return a mock CallToolResult result = MagicMock(spec=CallToolResult) result.content = [{"type": "text", "text": "Tool executed successfully"}] @@ -8709,6 +8710,35 @@ class TestMCPServerManagerExpandToolPermissions: result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]}) assert sorted(result["uuid-a"]) == ["read_file", "write_file"] + def test_wildcard_survives_expansion_as_list_entry(self): + """["*"] stays in the expanded list so the caller's wildcard check + (``_union_tool_grants``) can read it; this function only normalizes + keys and never maps grants to None.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": ["*"]}) + assert result == {"uuid-a": ["*"]} + + def test_wildcard_unions_with_concrete_names_across_keys_for_same_server(self): + """An alias key carrying ["*"] unioned with an id key naming one tool + keeps both entries; interpretation of the wildcard belongs to the + caller, not the expansion.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a", alias="alias-a") + + result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["*"]}) + assert sorted(result["uuid-a"]) == ["*", "read_file"] + + def test_empty_list_stays_deny_all(self): + """[] is deny-all, a distinct meaning from no entry (unrestricted); + the key must survive expansion rather than disappear.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": []}) + assert result == {"uuid-a": []} + class TestOAuthDiscoverySSRFGuard: """SSRF guard for the OAuth metadata discovery follow-up fetches. @@ -10032,7 +10062,7 @@ class _RetryFakeClient: self._MCPClient = MCPClient self.attempts = 0 - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False): + async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): self.attempts += 1 if self._raises is not None: if raise_on_error: @@ -10244,7 +10274,7 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False): + async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -12221,6 +12251,38 @@ class TestOpenApiHandlerRelaysUpstreamAuth: assert result.is_error is True assert "upstream returned HTTP 503" in result.content[0].text + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("body", "compat", "expected_structured"), + [ + ('{"total": 1.10, "items": [ ]}', "legacy", {"total": 1.1, "items": []}), + ('{"total": 1.10, "items": [ ]}', "modern", {"total": 1.1, "items": []}), + ("[1, 2]", "legacy", None), + ("[1, 2]", "modern", [1, 2]), + ("plain text", "legacy", None), + ("plain text", "modern", None), + ], + ) + async def test_json_bodies_keep_verbatim_text_and_gain_structured_content(self, body, compat, expected_structured): + """The OpenAPI arm used to stringify the response; now the text block is the upstream body + byte for byte, exactly once, and JSON bodies carry structuredContent when the caller's revision admits it.""" + from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat, parse_http_body + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager = MCPServerManager() + + async def handler(**_kwargs): + return parse_http_body(body) + + tool = MagicMock() + tool.handler = handler + with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool): + result = await manager._call_openapi_tool_handler(self._server(), "list_reports", {}, WireCompat(compat)) + + assert result.is_error is False + assert [block.text for block in result.content] == [body] + assert result.structured_content == expected_structured + class TestConfigServerIdPinning: """config.yaml servers may pin ``server_id`` so permission grants survive connection edits.""" @@ -14225,7 +14287,7 @@ class TestProtectedCredentialPreparation: caller_token: Final = _request_auth_header.set(caller) extra_token: Final = _request_extra_headers.set(forwarded) try: - assert await tool() == "authenticated" + assert await tool() == TextResult("authenticated") sent: Final = destination.calls.last.request.headers assert sent.get("x-api-key") == static.get("X-API-Key", (forwarded or {}).get("X-API-Key")) if caller: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index bd351f9106e..6b0211c3866 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, _request_extra_headers, _request_resolved_auth_headers, + _request_upstream_url, _resolve_param_list, _resolve_ref, build_input_schema, @@ -31,6 +32,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( get_base_url, resolve_operation_params, ) +from litellm.proxy._experimental.mcp_server.tool_outcome import JsonResult, TextResult from litellm.proxy._experimental.mcp_server.exceptions import ( MCPOpenApiUpstreamError, @@ -40,6 +42,43 @@ from litellm.proxy._experimental.mcp_server.exceptions import ( GET_ASYNC_CLIENT_TARGET = "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client" +@pytest.mark.asyncio +async def test_unsupported_http_method_returns_text_without_sending_request( + respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + tool: Final = create_tool_function("/echo", "HEAD", {}, "https://upstream.example") + token: Final = _request_upstream_url.set("https://outer.example/request") + try: + assert await tool() == TextResult("Unsupported HTTP method: head") + assert len(respx_mock.calls) == 0 + assert _request_upstream_url.get() == "https://outer.example/request" + finally: + _request_upstream_url.reset(token) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body,expected", [ + (' { "ok": true }\n', JsonResult({"ok": True}, ' { "ok": true }\n')), + (' [1, 2]\n', JsonResult([1, 2], ' [1, 2]\n')), + ('false', JsonResult(False, 'false')), + ('0', JsonResult(0, '0')), + ('""', JsonResult("", '""')), + ('null', TextResult('null')), + ('{"unfinished":', TextResult('{"unfinished":')), + ('', TextResult('')), +]) +async def test_http_response_preserves_body_and_classifies_json( + respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, + body: str, expected: TextResult | JsonResult, +) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example") + destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text=body) + assert await tool() == expected + assert destination.call_count == 1 + + @pytest.mark.asyncio @pytest.mark.parametrize("auth_type,value,accepted", [ (MCPAuth.api_key, "Bearer Bearer", False), (MCPAuth.api_key, "ApiKey ApiKey", False), @@ -63,7 +102,7 @@ async def test_authorization_validates_credentials_before_http( caller_token: Final = _request_auth_header.set(value) try: if accepted: - assert await tool() == "authenticated" + assert await tool() == TextResult("authenticated") assert destination.call_count == 1 assert destination.calls.last.request.headers["authorization"] == value else: @@ -103,7 +142,7 @@ async def test_static_auth_validates_headers_after_existing_precedence( assert exc.value.status_code == 500 assert destination.call_count == 0 else: - assert await tool() == "authenticated" + assert await tool() == TextResult("authenticated") assert destination.call_count == 1 assert destination.calls.last.request.headers["authorization"] == expected finally: @@ -124,7 +163,7 @@ async def test_static_auth_uses_configured_custom_header( ) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") if credential: - assert await tool() == "authenticated" + assert await tool() == TextResult("authenticated") assert destination.call_count == 1 assert destination.calls.last.request.headers["x-custom"] == credential else: @@ -144,7 +183,7 @@ async def test_static_auth_accepts_api_key_carried_by_static_header( ) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") if credential: - assert await tool() == "authenticated" + assert await tool() == TextResult("authenticated") assert destination.calls.last.request.headers["apikey"] == credential assert "x-api-key" not in destination.calls.last.request.headers else: @@ -167,7 +206,7 @@ async def test_static_validation_preserves_no_auth_and_resolved_oauth( destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo") token: Final = _request_resolved_auth_headers.set(resolved) try: - assert await tool() == "echo" + assert await tool() == TextResult("echo") assert destination.call_count == 1 assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization") finally: @@ -220,7 +259,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"repository-id": "test-repo"}) - assert result == '{"id": "123"}' + assert result == JsonResult({"id": "123"}, '{"id": "123"}') # Verify URL was constructed correctly call_args = async_client.get.call_args @@ -256,7 +295,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"2fa-code": "123456"}) - assert result == "verified" + assert result == TextResult("verified") # Verify query parameter was included call_args = async_client.post.call_args @@ -290,7 +329,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"user.name": "john.doe"}) - assert result == "found" + assert result == TextResult("found") call_args = async_client.get.call_args assert call_args[1]["params"]["user.name"] == "john.doe" @@ -323,7 +362,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"$filter": "name eq 'test'"}) - assert result == "[]" + assert result == JsonResult([], "[]") call_args = async_client.get.call_args assert call_args[1]["params"]["$filter"] == "name eq 'test'" @@ -356,7 +395,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"class": "premium"}) - assert result == "items" + assert result == TextResult("items") call_args = async_client.get.call_args assert call_args[1]["params"]["class"] == "premium" @@ -407,7 +446,7 @@ class TestCreateToolFunction: "$filter": "active", } ) - assert result == "success" + assert result == TextResult("success") @pytest.mark.asyncio async def test_request_body_parameter(self): @@ -440,7 +479,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"body": {"name": "test"}}) - assert result == "created" + assert result == TextResult("created") call_args = async_client.post.call_args assert call_args[1]["json"] == {"name": "test"} @@ -464,7 +503,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func() - assert result == "ok" + assert result == TextResult("ok") @pytest.mark.asyncio async def test_all_http_methods(self): @@ -497,7 +536,7 @@ class TestCreateToolFunction: mock_client.return_value = async_client result = await func(**{"repository-id": "test"}) - assert result == "success" + assert result == TextResult("success") def test_no_exec_usage(self): """Verify that create_tool_function does not use exec().""" @@ -614,7 +653,7 @@ class TestPathSecurity: response = await tool_function(**{"filename": "../admin"}) - assert "Invalid path parameter" in response + assert isinstance(response, TextResult) and "Invalid path parameter" in response.text @pytest.mark.asyncio async def test_should_encode_and_request_safe_path_parameters(self): @@ -643,7 +682,7 @@ class TestPathSecurity: response = await tool_function(**{"filename": "report 2024.json"}) - assert response == "dummy-response" + assert response == TextResult("dummy-response") # Verify URL was properly encoded call_args = async_client.get.call_args @@ -1181,7 +1220,7 @@ class TestRequestExtraHeaders: finally: _request_extra_headers.reset(token) - assert result == "ok" + assert result == TextResult("ok") call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert headers_sent.get("X-TOKEN") == "secret-value" @@ -1204,7 +1243,7 @@ class TestRequestExtraHeaders: result = await func() - assert result == "ok" + assert result == TextResult("ok") call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert headers_sent == {"X-Static": "static-value"} @@ -1232,7 +1271,7 @@ class TestRequestExtraHeaders: finally: _request_extra_headers.reset(token) - assert result == "created" + assert result == TextResult("created") call_args = async_client.post.call_args headers_sent = call_args[1]["headers"] assert headers_sent.get("X-Static") == "static-value" @@ -1260,7 +1299,7 @@ class TestRequestExtraHeaders: finally: _request_extra_headers.reset(token) - assert result == "ok" + assert result == TextResult("ok") call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert headers_sent.get("X-Tenant") == "operator-tenant" @@ -1288,7 +1327,7 @@ class TestRequestExtraHeaders: finally: _request_extra_headers.reset(token) - assert result == "ok" + assert result == TextResult("ok") call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert headers_sent.get("X-Tenant") == "operator-tenant" @@ -1320,7 +1359,7 @@ class TestRequestExtraHeaders: _request_auth_header.reset(auth_token) _request_extra_headers.reset(extra_token) - assert result == "secure-data" + assert result == TextResult("secure-data") call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert headers_sent.get("Authorization") == "Bearer byok-credential" @@ -1380,7 +1419,7 @@ class TestRequestExtraHeaders: _request_extra_headers.reset(extra_token) _request_resolved_auth_headers.reset(resolved_token) - assert result == "secure-data" + assert result == TextResult("secure-data") headers_sent = async_client.get.call_args[1]["headers"] authorization_values = [v for k, v in headers_sent.items() if k.lower() == "authorization"] assert authorization_values == ["Bearer resolved-oauth"] @@ -1436,7 +1475,7 @@ class TestUpstreamStatusIsClassified: async def test_success_still_returns_the_body(self): tool, client = self._tool(200, text='{"reports": []}') with patch(GET_ASYNC_CLIENT_TARGET, return_value=client): - assert await tool() == '{"reports": []}' + assert await tool() == JsonResult({"reports": []}, '{"reports": []}') @pytest.mark.asyncio async def test_401_raises_the_reauth_signal_carrying_the_challenge(self): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index bb70f38285c..15d3b67e641 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -9,6 +9,7 @@ from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock, patch import pytest +from mcp.types import CallToolResult from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -46,7 +47,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): fake_tool.name = "list_pets" pre_call = AsyncMock(return_value={}) - handle_local = AsyncMock(return_value=[]) + handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False)) with ( patch.object( @@ -131,7 +132,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): pre_call = AsyncMock( side_effect=HTTPException(status_code=403, detail="not allowed") ) - handle_local = AsyncMock(return_value=[]) + handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False)) with ( patch.object( @@ -191,7 +192,7 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): fake_tool.name = "list_pets" pre_call = AsyncMock(return_value={}) - handle_local = AsyncMock(return_value=[]) + handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False)) resolve_auth = MagicMock() # `_get_mcp_server_from_tool_name` returns None — no server context. @@ -275,9 +276,9 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): fake_tool.name = "get_values" captured: dict = {} - async def handle_local(_name, _arguments): + async def handle_local(_name, _arguments, _wire_compat): captured["resolved"] = _request_resolved_auth_headers.get() - return [] + return CallToolResult(content=[], is_error=False) with ( patch.object( @@ -603,13 +604,13 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc captured["resolver_credential"] = kwargs["mcp_auth_header"] return None, kwargs["forwarded_headers"] - async def capture_local(_name, _arguments): + async def capture_local(_name, _arguments, _wire_compat): captured["injected"] = _request_auth_header.get() - return [] + return CallToolResult(content=[], is_error=False) - async def capture_openapi_handler(_server, _name, _arguments): + async def capture_openapi_handler(_server, _name, _arguments, _wire_compat): captured["injected"] = _request_auth_header.get() - return [] + return CallToolResult(content=[], is_error=False) manager = mcp_operations.global_mcp_server_manager with ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py index 81877c38389..81f81045740 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -35,6 +35,7 @@ async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(capl @pytest.mark.asyncio async def test_dispatch_uses_explicit_context_when_ambient_caller_differs(): from mcp.server.auth.middleware.auth_context import auth_context_var + from litellm.proxy._experimental.mcp_server.server import set_auth_context context = prepare_context( @@ -64,12 +65,13 @@ async def test_dispatch_uses_explicit_context_when_ambient_caller_differs(): @pytest.mark.asyncio async def test_legacy_adapter_cleans_context_after_cancelled_operation(): from types import SimpleNamespace + from litellm.proxy._experimental.mcp_server import server from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var previous_session = server.active_mcp_session_var.get() previous_request = active_mcp_request_ctx_var.get() - request = SimpleNamespace(session=object()) + request = SimpleNamespace(session=object(), protocol_version="2025-06-18") auth = (None, None, None, None, None, None, None) async def cancelled_operation(): @@ -90,6 +92,7 @@ async def test_legacy_adapter_cleans_context_after_cancelled_operation(): @pytest.mark.asyncio async def test_legacy_adapter_cleans_context_when_trace_setup_fails(): from types import SimpleNamespace + from litellm.proxy._experimental.mcp_server import server from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var @@ -111,6 +114,7 @@ async def test_legacy_adapter_cleans_context_when_trace_setup_fails(): @pytest.mark.asyncio async def test_prompt_sampling_receives_explicit_operation_caller_headers_and_ip(): from unittest.mock import MagicMock + from litellm.proxy._experimental.mcp_server import operations from litellm.types.mcp import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -201,9 +205,11 @@ def _catalog_case(method): @pytest.mark.parametrize("state", ["success", "denied", "upstream_failure", "scope_failure"]) async def test_native_catalog_operations_preserve_context_results_and_failure_policy(method, state): from types import SimpleNamespace + from fastapi import HTTPException from mcp.server.context import ServerRequestContext from mcp.types import PaginatedRequestParams + from litellm.proxy._experimental.mcp_server import operations, server from litellm.types.mcp import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -260,6 +266,7 @@ async def test_native_catalog_operations_preserve_context_results_and_failure_po async def test_explicit_proxy_context_rejects_catalog_operations_before_upstream_access(method): from mcp.shared.exceptions import MCPError from mcp.types import METHOD_NOT_FOUND + from litellm.proxy._experimental.mcp_server import operations operation, _, manager_method, _, _ = _catalog_case(method) @@ -276,6 +283,7 @@ async def test_explicit_proxy_context_rejects_catalog_operations_before_upstream @pytest.mark.parametrize("failure", ["missing_env", "pii", "guardrail", "unexpected"]) async def test_tool_operation_preserves_failure_messages_and_request_trace(failure): from mcp.types import CallToolRequest, CallToolRequestParams + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException from litellm.proxy._experimental.mcp_server import operations from litellm.proxy._experimental.mcp_server.utils import MCPMissingUserEnvVarsError @@ -333,6 +341,7 @@ async def test_catalog_operation_preserves_empty_result_for_malformed_upstream_i @pytest.mark.parametrize("catalog_unavailable", [False, True]) async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailable_catalog(catalog_unavailable): from mcp.types import ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations allowed = AsyncMock( @@ -352,6 +361,7 @@ async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailabl @pytest.mark.asyncio async def test_explicit_proxy_context_lists_builtin_tools_and_blocks_direct_tool_dispatch(): from mcp.types import CallToolRequest, CallToolRequestParams, ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations context = prepare_context(mcp_proxy_mode=True) @@ -486,8 +496,10 @@ class TestChallengeMissingTokenExchangeSubject: @pytest.mark.asyncio async def test_execute_mcp_tool_challenges_missing_subject_before_cold_listing(): """On a cold catalog the challenge fires before any listing or tool resolution is attempted.""" - from fastapi import HTTPException from datetime import datetime, timezone + + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server import operations server = _server("te-exec", MCPAuth.oauth2_token_exchange) @@ -509,3 +521,24 @@ async def test_execute_mcp_tool_challenges_missing_subject_before_cold_listing() ) assert exc_info.value.status_code == 401 listing.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("compat", ["legacy", "modern"]) +async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(compat: str) -> None: + """The local-registry arm used to convert at MODERN and let the legacy downgrade append a second + text block; converting at the caller's revision keeps the upstream body exactly once.""" + from unittest.mock import MagicMock + + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat, parse_http_body + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + body = '["a","b"]' + tool = MagicMock() + tool.handler = AsyncMock(return_value=parse_http_body(body)) + with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool): + result = await operations._handle_local_mcp_tool("reports-list_tags", {}, WireCompat(compat)) + + assert [block.text for block in result.content] == [body] + assert result.structured_content == (["a", "b"] if compat == "modern" else None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index e20d74ab60d..4da120cb26f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1,4 +1,3 @@ -from litellm.proxy._experimental.mcp_server import operations as mcp_operations import asyncio import inspect import json @@ -7,12 +6,15 @@ from datetime import datetime from typing import Any, Dict, Final, Optional from unittest.mock import AsyncMock, MagicMock +from litellm.proxy._experimental.mcp_server import operations as mcp_operations + if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 from exceptiongroup import BaseExceptionGroup import httpx import pytest from fastapi import HTTPException +from mcp.types import CallToolResult, TextContent from starlette.requests import Request from litellm.constants import MCP_TOOL_LISTING_TIMEOUT @@ -29,6 +31,8 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.mcp import MCPAuth, MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer +_OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False) + def _rendered_log_message(call): message = str(call.args[0]) @@ -1472,10 +1476,10 @@ class TestListToolsRestAPI: monkeypatch, ): """The REST tools/list path should include tools beyond the upstream first page.""" - import litellm.experimental_mcp_client.client as mcp_client_module from mcp.types import ListToolsResult, PaginatedRequestParams from mcp.types import Tool as MCPTool + import litellm.experimental_mcp_client.client as mcp_client_module from litellm.proxy._experimental.mcp_server.server import MCPServer from litellm.types.mcp import MCPTransport @@ -2470,7 +2474,7 @@ class TestCallToolRestAPI: async def fake_execute_mcp_tool(**kwargs): captured.update(kwargs) - return {"result": "ok"} + return _OK_TOOL_RESULT monkeypatch.setattr( rest_endpoints, @@ -2530,13 +2534,89 @@ class TestCallToolRestAPI: user_api_key_dict=UserAPIKeyAuth(), ) - assert result == {"result": "ok"} + assert result == _OK_TOOL_RESULT assert captured["name"] == "demo-tool" assert captured["arguments"] == {"foo": "bar"} assert captured["allowed_mcp_servers"] == [stub_server] assert captured["oauth2_headers"] is None fire_logging.assert_awaited_once() + @pytest.mark.parametrize( + ("structured", "expected_structured", "expected_texts"), + [ + ({"a": 1}, {"a": 1}, ['{"a": 1}']), + ([1, 2], None, ['{"a": 1}', "[1, 2]"]), + ], + ) + async def test_rest_keeps_its_serialization_shape_with_legacy_structured_admission( + self, monkeypatch, structured, expected_structured, expected_texts + ): + """REST has no negotiated revision, so it admits object structuredContent only and downgrades + anything else losslessly, while the response keeps the SDK model shape (resultType included) + rather than being run through the MCP legacy wire serializer.""" + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + auth_type = None + + stub_server = StubServer() + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + async def fake_execute_mcp_tool(**kwargs): + return CallToolResult( + content=[TextContent(type="text", text='{"a": 1}')], + structuredContent=structured, + isError=False, + ) + + async def fake_fire_logging(logging_obj, result, start_time, end_time, **kwargs): + return result + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", fake_get_allowed_mcp_servers + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False + ) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False) + monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False) + monkeypatch.setattr(rest_endpoints, "_fire_mcp_tool_call_logging", fake_fire_logging, raising=False) + + request = _build_request( + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}}, + ) + + result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth()) + + assert isinstance(result, CallToolResult) + dumped = result.model_dump(by_alias=True, mode="json", exclude_none=True) + assert dumped.get("structuredContent") == expected_structured + assert [block["text"] for block in dumped["content"]] == expected_texts + assert dumped["resultType"] == "complete" + assert dumped["isError"] is False + @pytest.mark.asyncio @pytest.mark.parametrize( ("auth_type", "per_user_oauth", "expected"), @@ -2580,7 +2660,7 @@ class TestCallToolRestAPI: async def fake_execute_mcp_tool(**kwargs): captured.update(kwargs) - return {"result": "ok"} + return _OK_TOOL_RESULT monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", fake_get_allowed_mcp_servers @@ -2607,7 +2687,7 @@ class TestCallToolRestAPI: result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth()) - assert result == {"result": "ok"} + assert result == _OK_TOOL_RESULT assert captured["oauth2_headers"] == expected assert captured["raw_headers"]["authorization"] == "Bearer user-subject-token" @@ -2637,7 +2717,7 @@ class TestCallToolRestAPI: return kwargs.get("data", {}) async def fake_execute_mcp_tool(**kwargs): - return {"content": [{"type": "text", "text": "jane@example.com"}]} + return CallToolResult(content=[TextContent(type="text", text="jane@example.com")], is_error=False) monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) monkeypatch.setattr( @@ -2659,7 +2739,7 @@ class TestCallToolRestAPI: ) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False) monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False) - masked_result = {"content": [{"type": "text", "text": ""}]} + masked_result = CallToolResult(content=[TextContent(type="text", text="")], is_error=False) monkeypatch.setattr( rest_endpoints, "_fire_mcp_tool_call_logging", @@ -2714,9 +2794,9 @@ class TestCallToolRestAPI: async def fake_execute_mcp_tool(**kwargs): captured.update(kwargs) - return {"result": "ok"} + return _OK_TOOL_RESULT - fire_logging = AsyncMock(return_value={"result": "ok"}) + fire_logging = AsyncMock(return_value=_OK_TOOL_RESULT) monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_result_conversion.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_result_conversion.py new file mode 100644 index 00000000000..d9b5063a811 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_result_conversion.py @@ -0,0 +1,243 @@ +import json +from typing import Final + +import pytest +from mcp.types import CallToolResult, ImageContent, InputRequiredResult, TextContent, Tool +from mcp_types.methods import serialize_server_result +from mcp_types.version import KNOWN_PROTOCOL_VERSIONS, MODERN_PROTOCOL_VERSIONS +from pydantic import JsonValue, ValidationError + +from litellm.proxy._experimental.mcp_server.result_conversion import ( + INPUT_REQUIRED_UNSUPPORTED_MESSAGE, + JsonResult, + TextResult, + WireCompat, + complete_call_tool_result, + error_text_result, + handler_outcome, + parse_http_body, + to_call_tool_result, + to_gateway_tool, + wire_compat_for, +) + +BOTH: Final = (WireCompat.LEGACY, WireCompat.MODERN) + + +def _interim() -> InputRequiredResult: + return InputRequiredResult.model_validate( + { + "resultType": "input_required", + "inputRequests": { + "req-1": { + "method": "elicitation/create", + "params": {"message": "Pick one", "requestedSchema": {"type": "object", "properties": {}}}, + } + }, + "requestState": "abc", + } + ) + + +def _wire(result: CallToolResult | InputRequiredResult, version: str) -> dict[str, object]: + return serialize_server_result( + "tools/call", version, result.model_dump(by_alias=True, mode="json", exclude_none=True) + ) + + +class TestWireCompatFor: + def test_only_modern_revisions_map_to_modern(self): + for version in KNOWN_PROTOCOL_VERSIONS: + expected: Final = WireCompat.MODERN if version in MODERN_PROTOCOL_VERSIONS else WireCompat.LEGACY + assert wire_compat_for(version) is expected, version + assert wire_compat_for("1999-01-01") is WireCompat.LEGACY + + +class TestParseHttpBody: + @pytest.mark.parametrize("body", ["", " ", "{not json", "null"]) + def test_non_structured_bodies_stay_text(self, body: str): + assert parse_http_body(body) == TextResult(body) + + @pytest.mark.parametrize( + "body, value", + [ + ('{"a": 1}', {"a": 1}), + ("[1, 2]", [1, 2]), + ("1.10", 1.1), + ("true", True), + ('"hi"', "hi"), + ], + ) + def test_json_bodies_keep_original_text(self, body: str, value: object): + assert parse_http_body(body) == JsonResult(value=value, original_text=body) + + def test_handler_outcome_stringifies_unknown_values(self): + assert handler_outcome(42) == TextResult("42") + assert handler_outcome(TextResult("x")) == TextResult("x") + + +class TestTextAndJsonArms: + @pytest.mark.parametrize("compat", BOTH) + def test_text_result(self, compat: WireCompat): + result = to_call_tool_result(TextResult("plain"), compat) + assert isinstance(result, CallToolResult) + assert result.is_error is False + assert [c.text for c in result.content if isinstance(c, TextContent)] == ["plain"] + assert result.structured_content is None + + @pytest.mark.parametrize("compat", BOTH) + def test_json_object_is_structured_everywhere_and_text_is_verbatim(self, compat: WireCompat): + body: Final = '{"n": 1.10,\n"k": "v"}' + result = to_call_tool_result(parse_http_body(body), compat) + assert isinstance(result, CallToolResult) + assert result.structured_content == {"n": 1.1, "k": "v"} + assert [c.text for c in result.content if isinstance(c, TextContent)] == [body] + + @pytest.mark.parametrize("body", ["[1, 2]", "3", "true", '"s"']) + def test_non_object_json_is_structured_only_on_modern(self, body: str): + legacy = to_call_tool_result(parse_http_body(body), WireCompat.LEGACY) + modern = to_call_tool_result(parse_http_body(body), WireCompat.MODERN) + assert isinstance(legacy, CallToolResult) and isinstance(modern, CallToolResult) + assert legacy.structured_content is None + assert modern.structured_content == json.loads(body) + for result in (legacy, modern): + assert [c.text for c in result.content if isinstance(c, TextContent)] == [body] + + @pytest.mark.parametrize("compat", BOTH) + def test_json_null_keeps_text_and_claims_no_structured_field(self, compat: WireCompat): + result = to_call_tool_result(parse_http_body("null"), compat) + assert isinstance(result, CallToolResult) + assert [c.text for c in result.content if isinstance(c, TextContent)] == ["null"] + assert "structuredContent" not in _wire(result, "2026-07-28") + + +class TestSdkResultArm: + def _incoming(self, content: list[TextContent]) -> CallToolResult: + return CallToolResult(content=content, structured_content=[1, 2], meta={"trace": "t1"}, is_error=False) + + def test_modern_passes_through_the_same_object(self): + incoming = self._incoming([]) + assert to_call_tool_result(incoming, WireCompat.MODERN) is incoming + + def test_legacy_downgrade_with_empty_content_appends_json_text(self): + incoming = self._incoming([]) + result = to_call_tool_result(incoming, WireCompat.LEGACY) + assert isinstance(result, CallToolResult) + assert result.structured_content is None + assert [c.text for c in result.content if isinstance(c, TextContent)] == ["[1, 2]"] + assert result.meta == {"trace": "t1"} + assert incoming.structured_content == [1, 2] and incoming.content == [] + + def test_legacy_downgrade_keeps_unrelated_content_and_appends_json_text(self): + incoming = self._incoming([TextContent(type="text", text="Done")]) + result = to_call_tool_result(incoming, WireCompat.LEGACY) + assert isinstance(result, CallToolResult) + assert [c.text for c in result.content if isinstance(c, TextContent)] == ["Done", "[1, 2]"] + assert incoming.content == [TextContent(type="text", text="Done")] + assert incoming.structured_content == [1, 2] + + def test_legacy_keeps_object_structured_content(self): + incoming = CallToolResult(content=[], structured_content={"a": 1}, is_error=False) + assert to_call_tool_result(incoming, WireCompat.LEGACY) is incoming + + @pytest.mark.parametrize("value", [False, 0, "", []]) + def test_legacy_downgrade_preserves_falsy_values_and_non_text_blocks(self, value: JsonValue) -> None: + incoming: Final = CallToolResult( + content=[ + ImageContent(type="image", data="AA==", mime_type="image/png"), + TextContent(type="text", text="Done"), + ], + structured_content=value, + meta={"trace": "t1"}, + is_error=True, + ) + before: Final = incoming.model_dump(by_alias=True) + result: Final = to_call_tool_result(incoming, WireCompat.LEGACY) + assert isinstance(result, CallToolResult) + assert result.content == [*incoming.content, TextContent(type="text", text=json.dumps(value))] + assert result.structured_content is None + assert result.meta == incoming.meta + assert result.is_error is True + assert incoming.model_dump(by_alias=True) == before + + def test_is_error_survives_downgrade(self): + incoming = CallToolResult(content=[], structured_content=7, is_error=True) + result = to_call_tool_result(incoming, WireCompat.LEGACY) + assert isinstance(result, CallToolResult) and result.is_error is True + + +class TestInterimAndExceptionArms: + def test_modern_interim_passes_through(self): + interim = _interim() + assert to_call_tool_result(interim, WireCompat.MODERN) is interim + + def test_legacy_interim_becomes_error_result(self): + result = to_call_tool_result(_interim(), WireCompat.LEGACY) + assert isinstance(result, CallToolResult) + assert result.is_error is True + assert [c.text for c in result.content if isinstance(c, TextContent)] == [INPUT_REQUIRED_UNSUPPORTED_MESSAGE] + + def test_complete_call_tool_result_never_returns_interim(self): + result = complete_call_tool_result(_interim(), WireCompat.MODERN) + assert isinstance(result, CallToolResult) and result.is_error is True + + @pytest.mark.parametrize("compat", BOTH) + def test_exception_arm_matches_error_text_result(self, compat: WireCompat): + exc = ValueError("boom") + result = to_call_tool_result(exc, compat) + assert result == error_text_result(exc) + assert isinstance(result, CallToolResult) and result.is_error is True + assert [c.text for c in result.content if isinstance(c, TextContent)] == ["ValueError: boom"] + + +class TestSdkWireSerialization: + @pytest.mark.parametrize("version", KNOWN_PROTOCOL_VERSIONS) + def test_converted_results_serialize_on_their_negotiated_revision(self, version: str): + compat = wire_compat_for(version) + for body in ('{"a": 1}', "[1, 2]", "3", "null", "text"): + result = to_call_tool_result(parse_http_body(body), compat) + frame = _wire(result, version) + assert frame["content"] == [{"type": "text", "text": body}] + structured = json.loads(body) if body != "text" else None + expects_structured = structured is not None and ( + compat is WireCompat.MODERN or isinstance(structured, dict) + ) + assert ("structuredContent" in frame) is expects_structured, (version, body) + if expects_structured: + assert frame["structuredContent"] == structured + assert ("resultType" in frame) is (compat is WireCompat.MODERN), (version, body) + + @pytest.mark.parametrize("version", KNOWN_PROTOCOL_VERSIONS) + def test_downgraded_sdk_result_serializes_where_the_raw_one_would_not(self, version: str): + incoming = CallToolResult(content=[TextContent(type="text", text="Done")], structured_content=[1, 2]) + converted = to_call_tool_result(incoming, wire_compat_for(version)) + frame = _wire(converted, version) + if version in MODERN_PROTOCOL_VERSIONS: + assert frame["structuredContent"] == [1, 2] + return + with pytest.raises(ValidationError): + _wire(incoming, version) + assert "structuredContent" not in frame + assert frame["content"] == [{"type": "text", "text": "Done"}, {"type": "text", "text": "[1, 2]"}] + + def test_modern_interim_serializes_with_its_fields_intact(self): + frame = _wire(_interim(), "2026-07-28") + assert frame["resultType"] == "input_required" + assert frame["requestState"] == "abc" + assert frame["inputRequests"]["req-1"]["params"]["message"] == "Pick one" + + +class TestToGatewayTool: + def test_rename_is_a_deep_copy_that_keeps_every_other_field(self): + tool = Tool( + name="orig", + description="d", + inputSchema={"type": "object", "properties": {"q": {"type": "string"}}}, + _meta={"owner": "x"}, + ) + renamed = to_gateway_tool(tool, "srv-orig") + assert renamed.name == "srv-orig" + assert tool.name == "orig" + assert renamed.input_schema == tool.input_schema and renamed.input_schema is not tool.input_schema + assert renamed.meta == {"owner": "x"} + assert renamed.description == "d" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 30f5abdbb98..e42a47a1091 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only(): assert result is True +def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): + """A key alias whose target is on the key allowlist resolves like a team alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + result = _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + object_type="key", + fallback_depth=0, + ) + + assert result is True + + +def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): + """A key alias whose target is outside the key allowlist stays denied.""" + from litellm.proxy._types import ProxyErrorTypes, ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + assert exc_info.value.code == "403" + + +@pytest.mark.asyncio +async def test_can_team_access_model_honors_key_alias(): + """A key on a team can call a model through its own alias when the target is on the team allowlist.""" + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=["gpt-4o-mini"], + ) + + assert ( + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + +@pytest.mark.asyncio +async def test_can_key_call_model_honors_key_alias(): + """The real key entry point resolves a key alias to its target before the allowlist check.""" + from litellm.proxy.auth.auth_checks import can_key_call_model + + allowed_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + assert ( + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=allowed_token, + llm_router=None, + ) + is True + ) + + denied_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=denied_token, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): + """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): + """A key alias on the globally rewritten name resolves the same way the request chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): + """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_name_alone_is_not_enough(): + """A key that may call the alias name but not its target cannot call the alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="bar", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="bar", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_team_alias_applies_before_key_alias(): + """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_on_team_alias_target(): + """A key alias on the team-rewritten name resolves like the dispatch chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +@pytest.mark.asyncio +async def test_can_user_call_model_honors_key_alias(): + """A personal-scope key alias resolves to its target before the user allowlist check.""" + from litellm.proxy.auth.auth_checks import can_user_call_model + + user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"]) + + assert ( + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + ) + + assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_honors_key_alias(): + """A key alias resolves against the member allowlist, not just the raw alias name.""" + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), + ) + + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + def test_can_object_call_model_access_via_underlying_model_only(): """ Test that a key can access a model via underlying model even when using an alias. @@ -8668,7 +8998,7 @@ async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_t def __init__(self) -> None: self.namespace = None - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return _RecordingRedisClient() local_spend_counter_cache = DualCache() @@ -8742,7 +9072,7 @@ async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_ def __init__(self) -> None: self.namespace = None - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return _RecordingRedisClient() local_spend_counter_cache = DualCache() @@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them(): assert asked == ["agent-1", "agent-1"] +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_admits_the_key_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"]) + agent_key.aliases = {"fast": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_checks_the_team_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"} + resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + with pytest.raises(ModelAccessDeniedProxyException) as exc: + await _check_agent_access_group_model_access("foo", agent_key, None, resolve) + assert exc.value.type == ProxyErrorTypes.agent_model_access_denied + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment(): + from litellm.router import Router + + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"]) + agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"} + router: Final = Router(model_list=[]) + resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True + assert asked == ["agent-1"] + + @pytest.mark.asyncio async def test_agent_access_groups_naming_no_model_deny_every_model(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[]) @@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode assert asked == ["team:None", "user:alice", "team:None", "user:alice"] +@pytest.mark.asyncio +async def test_agent_key_alias_resolves_against_the_echoed_teams_models(): + agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a") + agent_key.aliases = {"foo": "bar"} + load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None) + cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None) + + await _check_caller_models(agent_key, "foo", load_team, load_user, cache) + + assert asked == ["team:team-a"] + + @pytest.mark.asyncio async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"]) diff --git a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py index 96770ee01c4..34d9741bcc0 100644 --- a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py @@ -87,7 +87,7 @@ class _FakeRedisCache: self._client = client self.namespace = namespace - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return self._client diff --git a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py index 0f64ef2b4ca..83ed3afa293 100644 --- a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py @@ -81,14 +81,25 @@ class _FailingPublishRedisClient(Redis): raise ConnectionError("redis down") -class _NotRedisClient: - def __init__(self) -> None: +class _ScriptedPubSubClient: + """Pub/sub-capable client that is not a redis.asyncio.Redis. + + Mirrors what RedisCache.init_pubsub_client returns for a cluster backend: + a node-level client exposing publish/pubsub without being an instance of + the standalone Redis class. + """ + + def __init__(self, pubsubs: Iterable["_QueuePubSub"]) -> None: + self._scripted_pubsubs = iter(pubsubs) self.published: List[Tuple[str, str]] = [] async def publish(self, channel: str, message: str) -> int: self.published.append((channel, message)) return 1 + def pubsub(self) -> "_QueuePubSub": + return next(self._scripted_pubsubs) + class _QueuePubSub: def __init__(self, initial_messages: Iterable[str] = ()) -> None: @@ -159,14 +170,14 @@ class _FakeRedisCache: self._client = client self.namespace = namespace - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return self._client class _ExplodingRedisCache: namespace: Optional[str] = None - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: raise ConnectionError("cannot connect") @@ -215,13 +226,17 @@ async def test_publish_swallows_client_init_errors() -> None: await publish_config_change(redis_cache=_ExplodingRedisCache(), object_type="litellm_proxymodeltable") -async def test_publish_skips_clients_without_pubsub_support() -> None: - client = _NotRedisClient() +async def test_publish_reaches_cluster_derived_pubsub_clients() -> None: + """LIT-8543: a cluster-backed cache returns a node-level client from + init_pubsub_client; publishes must go out on it instead of being skipped.""" + client = _ScriptedPubSubClient(pubsubs=[]) cache = _FakeRedisCache(client) await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable") - assert client.published == [] + assert client.published == [ + (CONFIG_SYNC_CHANNEL, json.dumps({"object_type": "litellm_proxymodeltable"})) + ] async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None: @@ -558,20 +573,26 @@ async def test_stop_before_start_is_a_noop() -> None: await subscriber.stop() -async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> None: - cache = _FakeRedisCache(_NotRedisClient()) +async def test_subscriber_subscribes_on_cluster_derived_pubsub_client() -> None: + """LIT-8543: the subscriber used to disable itself on cluster caches; now it + subscribes on the node-level client init_pubsub_client returns.""" + pubsub = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_proxymodeltable"})]) + cache = _FakeRedisCache(_ScriptedPubSubClient(pubsubs=[pubsub])) resyncs: List[str] = [] + fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, - resync_callbacks=(_recording_callback(resyncs, "resync", asyncio.Event()),), + resync_callbacks=(_recording_callback(resyncs, "resync", fired),), + debounce_seconds=0.01, + jitter_max_seconds=0.0, ) subscriber.start() - task = subscriber._task - assert task is not None - await asyncio.wait_for(task, timeout=5) + await asyncio.wait_for(fired.wait(), timeout=5) + await subscriber.stop() - assert resyncs == [] + assert pubsub.subscribed_channels == [CONFIG_SYNC_CHANNEL] + assert resyncs == ["resync"] class _FakeTableActions: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 1b696669724..0a4ffbaef26 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -4,10 +4,15 @@ Tests PII detection and masking for different message formats """ import asyncio +import copy import json +import re from contextlib import asynccontextmanager +from typing import Final from unittest.mock import MagicMock, patch +from aiohttp import web +from aiohttp.test_utils import TestServer import pytest @@ -2599,6 +2604,247 @@ async def test_apply_to_output_streaming_anthropic_first_frame_split_across_tran assert joined.count("event: message_start") == 1 +def _anthropic_stream_tail() -> list[bytes]: + return [ + _anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _anthropic_sse("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {}}), + _anthropic_sse("message_stop", {"type": "message_stop"}), + ] + + +def _anthropic_stream_head() -> list[bytes]: + return [ + _anthropic_sse( + "message_start", + {"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}}, + ), + _anthropic_sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + ] + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_first_frame_split_inside_a_utf8_character_is_still_masked(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + delta = ( + "event: content_block_delta\n" + + "data: " + + json.dumps( + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "John Smith designed the café."}, + }, + ensure_ascii=False, + ) + + "\n\n" + ).encode() + cut = delta.index("é".encode()) + 1 + assert delta[cut - 1 : cut] == b"\xc3", delta + byte_chunks = [*_anthropic_stream_head(), delta[:cut], delta[cut:], *_anthropic_stream_tail()] + + async def mock_stream(): + yield b"".join(byte_chunks[:2]) + byte_chunks[2] + for b in byte_chunks[3:]: + yield b + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + joined = b"".join(collected).decode() + assert "John Smith" not in joined, joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_stream_led_by_sse_comment_keepalive_is_still_masked(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + byte_chunks = [ + b": keepalive\n\n", + *_anthropic_stream_head(), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "John Smith"}}, + ), + *_anthropic_stream_tail(), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + raw = b"".join(collected) + assert raw.startswith(b": keepalive\n\n"), raw[:200] + joined = raw.decode() + assert "John Smith" not in joined, joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_anthropic_stream_led_by_data_less_ping_event_is_still_masked(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + byte_chunks = [ + b"event: ping\n\n", + *_anthropic_stream_head(), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "John Smith"}}, + ), + *_anthropic_stream_tail(), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + raw = b"".join(collected) + assert raw.startswith(b"event: ping\n\n"), raw[:200] + joined = raw.decode() + assert "John Smith" not in joined, joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_upstream_data_arrives(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + gate = asyncio.Event() + byte_chunks = [ + *_anthropic_stream_head(), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "John Smith"}}, + ), + *_anthropic_stream_tail(), + ] + + async def mock_stream(): + yield b": keepalive\n\n" + await gate.wait() + for b in byte_chunks: + yield b + + out = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ) + assert await asyncio.wait_for(anext(out), 1) == b": keepalive\n\n" + assert not gate.is_set() + + gate.set() + collected = [chunk async for chunk in out] + + joined = b"".join(collected).decode() + assert "John Smith" not in joined, joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + keepalives = [b": keepalive\n\n" * 512] * 12 # ~72 KiB of complete comment frames, over the 64 KiB cap + byte_chunks = [ + *keepalives[:-1], + keepalives[-1] + + b"".join( + [ + *_anthropic_stream_head(), + _anthropic_sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "John Smith"}}, + ), + ] + ), + *_anthropic_stream_tail(), + ] + + async def mock_stream(): + for b in byte_chunks: + yield b + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + joined = b"".join(collected).decode() + assert "John Smith" not in joined, joined + assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "" + assert joined.count("event: message_start") == 1 + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchanged(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + mock_redacted_text={"text": ""}, + ) + + async def mock_stream(): + yield b": keepalive\n\n" + + collected = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=mock_stream(), + request_data={}, + ): + collected.append(chunk) + + assert collected == [b": keepalive\n\n"] + + @pytest.mark.asyncio async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally(): guardrail = _OPTIONAL_PresidioPIIMasking( @@ -3849,3 +4095,79 @@ async def test_chunk_fanout_bound_is_shared_across_concurrent_calls(): ) assert state["peak"] >= 2 assert state["peak"] <= PRESIDIO_ANALYZE_CHUNK_CONCURRENCY + + +_PERSON_NAME: Final = re.compile(r"\b[A-Z][a-z]+ [A-Z][a-z]+\b") + + +def _person_spans(text: str) -> list[dict]: + return [ + {"entity_type": "PERSON", "start": match.start(), "end": match.end(), "score": 0.85, "analysis_explanation": None} + for match in _PERSON_NAME.finditer(text) + ] + + +def _redacted(text: str, spans: list[dict]) -> str: + starts = [0, *(span["end"] for span in spans)] + ends = [*(span["start"] for span in spans), len(text)] + return "".join(text[start:end] for start, end in zip(starts, ends)) + + +async def _fake_analyze(request: web.Request) -> web.Response: + payload = await request.json() + return web.json_response(_person_spans(payload["text"])) + + +async def _fake_anonymize(request: web.Request) -> web.Response: + payload = await request.json() + spans = payload["analyzer_results"] + items = [{"entity_type": span["entity_type"], "operator": "replace"} for span in spans] + return web.json_response({"text": _redacted(payload["text"], spans), "items": items}) + + +def _fake_presidio_app() -> web.Application: + app = web.Application() + app.router.add_post("/analyze", _fake_analyze) + app.router.add_post("/anonymize", _fake_anonymize) + return app + + +def _pii_turns() -> tuple[list[dict], list[dict]]: + turn_n = [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "My name is John Smith and my colleague is Alice Brown."}, + ] + reply = { + "role": "assistant", + "content": "Noted.", + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": "sig-1"}], + } + return turn_n, [*turn_n, reply, {"role": "user", "content": "Now compare against Bob Jones too."}] + + +async def test_pii_masking_replays_a_byte_identical_prefix_across_turns(mock_user_api_key, mock_cache): + """Masking rewrites the history on every turn, so the rewrite of an earlier message + must not depend on the turns that came after it or the signed thinking blocks in + the history lose their binding. The analyzer and anonymizer are an in-process fake + handed to the guardrail through its api_base settings.""" + async with TestServer(_fake_presidio_app()) as server: + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base=str(server.make_url("/")), + presidio_anonymizer_api_base=str(server.make_url("/")), + pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK}, + ) + masked = [ + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key, + cache=mock_cache, + data={"model": "claude-fable-5-1", "messages": copy.deepcopy(turn)}, + call_type="completion", + ) + for turn in _pii_turns() + ] + await guardrail._close_http_session() + earlier, later = (result["messages"] for result in masked) + + assert json.dumps(later[: len(earlier)], sort_keys=True) == json.dumps(earlier, sort_keys=True) + assert earlier[1]["content"] == "My name is and my colleague is ." + assert later[3]["content"] == "Now compare against too." diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index ee4c468a460..3ec5176159d 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -4307,3 +4307,22 @@ async def test_health_services_endpoint_pointfive_blocks_non_admin(monkeypatch, assert str(raised.value.code) == "403" logger_class.assert_not_called() + + +@pytest.mark.asyncio +async def test_health_services_endpoint_langfuse_missing_keys_errors(monkeypatch): + """v2 raised out of ``auth_check`` and the endpoint printed the server's answer; the v4 check + returns the failure as a value, and the endpoint has to error with that reason rather than a + generic credentials message that reads the same for an outage and a bad key.""" + import litellm.integrations.langfuse.langfuse as langfuse_module + from litellm.integrations.langfuse.langfuse_sdk import AuthCheckFailure + + logger_class = MagicMock() + logger_class.return_value.api_client.auth_check.return_value = AuthCheckFailure( + "connection refused by lf.internal.example" + ) + monkeypatch.setattr(langfuse_module, "LangFuseLogger", logger_class) + + with pytest.raises(ProxyException, match="auth_check failed") as raised: + await health_services_endpoint(service="langfuse") + assert "connection refused by lf.internal.example" in str(raised.value.message) diff --git a/tests/test_litellm/proxy/management_endpoints/test_callback_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_callback_management_endpoints.py index 5c5cfd0814d..8f19cf6329c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_callback_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_callback_management_endpoints.py @@ -81,35 +81,29 @@ class TestCallbackManagementEndpoints: # Setup test client client = TestClient(app) - # Initialize Langfuse logger and add to callbacks - with patch("litellm.integrations.langfuse.langfuse.Langfuse") as mock_langfuse: - # Mock the Langfuse client initialization - mock_langfuse_client = MagicMock() - mock_langfuse.return_value = mock_langfuse_client + # Add string representation to callback lists (this is how the system typically works) + litellm.success_callback.append("langfuse") + litellm._async_success_callback.append("langfuse") - # Add string representation to callback lists (this is how the system typically works) - litellm.success_callback.append("langfuse") - litellm._async_success_callback.append("langfuse") + # Make request to list callbacks endpoint + response = client.get( + "/callbacks/list", headers={"Authorization": "Bearer sk-1234"} + ) - # Make request to list callbacks endpoint - response = client.get( - "/callbacks/list", headers={"Authorization": "Bearer sk-1234"} - ) + # Verify response + assert response.status_code == 200 - # Verify response - assert response.status_code == 200 + response_data = response.json() - response_data = response.json() + # Verify langfuse appears in success callbacks + assert "langfuse" in response_data["success"] + assert response_data["failure"] == [] + assert response_data["success_and_failure"] == [] - # Verify langfuse appears in success callbacks - assert "langfuse" in response_data["success"] - assert response_data["failure"] == [] - assert response_data["success_and_failure"] == [] - - # Verify the response structure is correct - assert isinstance(response_data["success"], list) - assert isinstance(response_data["failure"], list) - assert isinstance(response_data["success_and_failure"], list) + # Verify the response structure is correct + assert isinstance(response_data["success"], list) + assert isinstance(response_data["failure"], list) + assert isinstance(response_data["success_and_failure"], list) def test_alist_callbacks_with_datadog_logger(self): """Test /callbacks/list endpoint with DataDog logger configuration""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 3b86f1f6d20..aa6be328f4a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -20931,3 +20931,165 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch assert deleted.index(object_permission_cache_key(permission_id)) < deleted.index( _hash_token_if_needed("sk-lit5479") ), deleted + + +class TestTeamAdminMemberKeyBudgetUpdate: + """LIT-5647: a team admin may update budget fields on another member's team key + only when the proxy enables the 'member_key_budgets' permission.""" + + def _member_key_row(self): + return LiteLLM_VerificationToken( + token="hashed_member_key", + user_id="member-1", + team_id="team-1", + key_alias="member", + models=["m"], + max_budget=10.0, + metadata={}, + ) + + def _caller(self, user_id="team-admin-1"): + return UserAPIKeyAuth( + user_id=user_id, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + def _team(self, members): + return LiteLLM_TeamTableCachedObj(team_id="team-1", members_with_roles=members) + + def _setup(self, monkeypatch, team_obj, editable_fields): + mock_get_team = AsyncMock(return_value=team_obj) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {"team_admin_editable_team_fields": editable_fields}, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks", + SimpleNamespace( + can_team_member_execute_key_management_endpoint=AsyncMock(return_value=None), + enforce_member_can_assign_access_groups=MagicMock(return_value=None), + ), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._check_team_key_limits", + AsyncMock(return_value=None), + ) + + @pytest.mark.asyncio + async def test_team_admin_updates_member_key_budget_when_enabled(self, monkeypatch): + self._setup( + monkeypatch, + self._team([Member(user_id="team-admin-1", role="admin"), Member(user_id="member-1", role="user")]), + ["member_key_budgets"], + ) + admin_check = AsyncMock(return_value=None) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._check_key_admin_access", + admin_check, + ) + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-member", max_budget=0, budget_duration="30d"), + existing_key_row=self._member_key_row(), + user_api_key_dict=self._caller(), + llm_router=None, + premium_user=True, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + admin_check.assert_called_once() + + @pytest.mark.asyncio + async def test_team_admin_denied_when_permission_disabled(self, monkeypatch): + self._setup( + monkeypatch, + self._team([Member(user_id="team-admin-1", role="admin"), Member(user_id="member-1", role="user")]), + ["tpm_limit"], + ) + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-member", max_budget=0), + existing_key_row=self._member_key_row(), + user_api_key_dict=self._caller(), + llm_router=None, + premium_user=True, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "member_key_budgets" in str(exc.value.detail) + assert "only create keys for themselves" not in str(exc.value.detail) + + @pytest.mark.asyncio + async def test_enabled_but_non_budget_field_is_denied(self, monkeypatch): + self._setup( + monkeypatch, + self._team([Member(user_id="team-admin-1", role="admin"), Member(user_id="member-1", role="user")]), + ["member_key_budgets"], + ) + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-member", key_alias="renamed"), + existing_key_row=self._member_key_row(), + user_api_key_dict=self._caller(), + llm_router=None, + premium_user=True, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "'key_alias'" in str(exc.value.detail) + + @pytest.mark.asyncio + async def test_ordinary_member_still_denied_on_another_members_key(self, monkeypatch): + self._setup( + monkeypatch, + self._team([Member(user_id="member-2", role="user"), Member(user_id="member-1", role="user")]), + ["member_key_budgets"], + ) + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-member", max_budget=0), + existing_key_row=self._member_key_row(), + user_api_key_dict=self._caller(user_id="member-2"), + llm_router=None, + premium_user=True, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "member_key_budgets" not in str(exc.value.detail) + + @pytest.mark.asyncio + async def test_personal_key_owned_by_someone_else_still_denied(self, monkeypatch): + self._setup( + monkeypatch, + self._team([Member(user_id="team-admin-1", role="admin")]), + ["member_key_budgets"], + ) + personal_row = LiteLLM_VerificationToken( + token="hashed_personal", + user_id="member-1", + team_id=None, + max_budget=10.0, + metadata={}, + ) + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=UpdateKeyRequest(key="sk-personal", max_budget=0), + existing_key_row=personal_row, + user_api_key_dict=self._caller(), + llm_router=None, + premium_user=True, + prisma_client=AsyncMock(), + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "member_key_budgets" not in str(exc.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_admin_field_permissions.py b/tests/test_litellm/proxy/management_endpoints/test_team_admin_field_permissions.py index 1a72d1de393..02cda355621 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_admin_field_permissions.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_admin_field_permissions.py @@ -1,14 +1,22 @@ import pytest from fastapi import HTTPException -from litellm.proxy._types import LiteLLM_ModelTable, LiteLLM_TeamTable, UpdateTeamRequest +from litellm.models.team import BudgetLimitEntry +from litellm.models.verification_token import LiteLLM_VerificationToken +from litellm.proxy._types import LiteLLM_ModelTable, LiteLLM_TeamTable, UpdateKeyRequest, UpdateTeamRequest from litellm.proxy.management_endpoints.team_admin_field_permissions import ( TeamAdminEditAllowed, TeamAdminEditingDisabled, TeamAdminFieldNotPermitted, + TeamAdminKeyEditAllowed, + TeamAdminMemberKeyEditingDisabled, + changed_key_fields, changed_team_fields, resolve_team_admin_editable_fields, team_admin_edit_verdict, + team_admin_key_edit_verdict, + team_admin_key_request_or_raise, + team_admin_may_edit_member_key_budgets, team_admin_may_manage_projects, team_admin_request_or_raise, ) @@ -156,3 +164,125 @@ class TestTeamAdminRequestOrRaise: team_admin_request_or_raise(TeamAdminFieldNotPermitted(field="blocked")) assert exc.value.status_code == 403 assert "'blocked'" in exc.value.detail + + +def _key(**overrides): + return LiteLLM_VerificationToken(token="hashed", **overrides) + + +class TestTeamAdminMayEditMemberKeyBudgets: + def test_missing_setting_denies(self): + assert team_admin_may_edit_member_key_budgets({}) is False + + def test_team_fields_alone_do_not_grant(self): + configured = {"team_admin_editable_team_fields": ["tpm_limit", "max_budget", "projects"]} + assert team_admin_may_edit_member_key_budgets(configured) is False + + def test_member_key_budgets_entry_grants(self): + configured = {"team_admin_editable_team_fields": ["member_key_budgets"]} + assert team_admin_may_edit_member_key_budgets(configured) is True + + @pytest.mark.parametrize("raw", ["member_key_budgets", 7, [1, 2]]) + def test_malformed_setting_denies(self, raw): + assert team_admin_may_edit_member_key_budgets({"team_admin_editable_team_fields": raw}) is False + + +class TestChangedKeyFields: + def test_key_alone_changes_nothing(self): + assert changed_key_fields(UpdateKeyRequest(key="sk-1"), _key()) == frozenset() + + def test_columns_echoing_stored_values_are_not_a_change(self): + data = UpdateKeyRequest(key="sk-1", max_budget=10.0, models=["m"], tpm_limit=5) + existing = _key(max_budget=10.0, models=["m"], tpm_limit=5) + assert changed_key_fields(data, existing) == frozenset() + + def test_column_with_different_value_is_a_change(self): + data = UpdateKeyRequest(key="sk-1", max_budget=0) + assert changed_key_fields(data, _key(max_budget=10.0)) == frozenset({"max_budget"}) + + def test_metadata_folded_field_echo_is_not_a_change(self): + data = UpdateKeyRequest(key="sk-1", tag_rpm_limit={"fast": 3}) + existing = _key(metadata={"tag_rpm_limit": {"fast": 3}}) + assert changed_key_fields(data, existing) == frozenset() + + def test_metadata_folded_field_difference_is_named_not_metadata(self): + data = UpdateKeyRequest(key="sk-1", tag_rpm_limit={"fast": 4}) + existing = _key(metadata={"tag_rpm_limit": {"fast": 3}}) + assert changed_key_fields(data, existing) == frozenset({"tag_rpm_limit"}) + + def test_budget_limits_echo_ignores_order_and_reset_at(self): + windows = [ + {"budget_duration": "1d", "max_budget": 5.0, "reset_at": "2030-01-01T00:00:00"}, + {"budget_duration": "7d", "max_budget": 50.0, "reset_at": "2030-01-07T00:00:00"}, + ] + data = UpdateKeyRequest( + key="sk-1", + budget_limits=[ + BudgetLimitEntry(budget_duration="7d", max_budget=50.0), + BudgetLimitEntry(budget_duration="1d", max_budget=5.0), + ], + ) + assert changed_key_fields(data, _key(budget_limits=windows)) == frozenset() + + def test_budget_limits_difference_is_a_change(self): + data = UpdateKeyRequest(key="sk-1", budget_limits=[BudgetLimitEntry(budget_duration="1d", max_budget=9.0)]) + existing = _key(budget_limits=[{"budget_duration": "1d", "max_budget": 5.0, "reset_at": "2030-01-01"}]) + assert changed_key_fields(data, existing) == frozenset({"budget_limits"}) + + def test_explicit_null_clearing_a_stored_column_is_a_change(self): + data = UpdateKeyRequest(key="sk-1", budget_duration=None) + assert changed_key_fields(data, _key(budget_duration="30d")) == frozenset({"budget_duration"}) + + def test_field_without_a_stored_counterpart_counts_as_changed_when_sent(self): + data = UpdateKeyRequest(key="sk-1", duration="1h") + assert changed_key_fields(data, _key()) == frozenset({"duration"}) + + +class TestTeamAdminKeyEditVerdict: + def test_disabled_even_for_a_no_op(self): + verdict = team_admin_key_edit_verdict(UpdateKeyRequest(key="sk-1"), _key(), enabled=False) + assert verdict == TeamAdminMemberKeyEditingDisabled() + + def test_budget_only_change_is_allowed(self): + data = UpdateKeyRequest(key="sk-1", max_budget=0, budget_duration="30d") + verdict = team_admin_key_edit_verdict(data, _key(max_budget=10.0), enabled=True) + assert verdict == TeamAdminKeyEditAllowed(changed=frozenset({"max_budget", "budget_duration"})) + + def test_key_alias_change_is_blocked_and_named(self): + data = UpdateKeyRequest(key="sk-1", key_alias="renamed") + verdict = team_admin_key_edit_verdict(data, _key(key_alias="member"), enabled=True) + assert verdict == TeamAdminFieldNotPermitted(field="key_alias") + + def test_spend_is_blocked(self): + data = UpdateKeyRequest(key="sk-1", spend=0) + verdict = team_admin_key_edit_verdict(data, _key(spend=3.5), enabled=True) + assert verdict == TeamAdminFieldNotPermitted(field="spend") + + def test_spend_echo_is_blocked_even_when_unchanged(self): + data = UpdateKeyRequest(key="sk-1", spend=4.5, max_budget=0) + verdict = team_admin_key_edit_verdict(data, _key(spend=4.5, max_budget=10.0), enabled=True) + assert verdict == TeamAdminFieldNotPermitted(field="spend") + + def test_budget_plus_non_budget_names_the_non_budget_field(self): + data = UpdateKeyRequest(key="sk-1", max_budget=0, key_alias="renamed") + verdict = team_admin_key_edit_verdict(data, _key(max_budget=10.0, key_alias="member"), enabled=True) + assert verdict == TeamAdminFieldNotPermitted(field="key_alias") + + +class TestTeamAdminKeyRequestOrRaise: + def test_allowed_returns_none(self): + verdict = TeamAdminKeyEditAllowed(changed=frozenset({"max_budget"})) + assert team_admin_key_request_or_raise(verdict) is None + + def test_disabled_is_a_403_pointing_at_member_key_budgets(self): + with pytest.raises(HTTPException) as exc: + team_admin_key_request_or_raise(TeamAdminMemberKeyEditingDisabled()) + assert exc.value.status_code == 403 + assert "member_key_budgets" in exc.value.detail + assert "Settings > UI > Team admin editable fields" in exc.value.detail + + def test_field_not_permitted_is_a_403_naming_the_field(self): + with pytest.raises(HTTPException) as exc: + team_admin_key_request_or_raise(TeamAdminFieldNotPermitted(field="key_alias")) + assert exc.value.status_code == 403 + assert "'key_alias'" in exc.value.detail diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 078315c2bf8..2fba39b30f6 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -1,12 +1,16 @@ import json +from collections.abc import Iterator +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException from litellm.proxy._types import ( + LiteLLM_AccessGroupTable, LiteLLM_ObjectPermissionBase, LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTableCachedObj, ObjectPermissionDict, SpecialMCPServerName, ) @@ -217,10 +221,12 @@ def _make_team_obj( mcp_servers=None, mcp_access_groups=None, mcp_tool_permissions=None, + access_group_ids=None, ): """Create a mock team object with the given MCP permissions.""" mock_team = MagicMock() mock_team.team_id = team_id + mock_team.access_group_ids = access_group_ids or [] if ( mcp_servers is not None @@ -541,6 +547,132 @@ async def test_validate_team_no_mcp_config_blocks_all( assert exc_info.value.status_code == 403 +@pytest.fixture +def unified_mcp_prisma() -> Iterator[MagicMock]: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + prisma: Final = MagicMock() + prisma.db.litellm_accessgrouptable.find_unique = AsyncMock( + return_value=LiteLLM_AccessGroupTable( + access_group_id="ag-1", + access_group_name="group one", + access_mcp_server_ids=["server-1"], + ) + ) + prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) + manager: Final = _make_mock_mcp_manager( + "server-1", + "server-2", + servers=[_make_mock_mcp_server("server-1", alias="server-alias")], + ) + manager.config_mcp_servers = {} + manager.get_allow_all_keys_server_ids.return_value = [] + with ( + patch( # test-quality-ok: management helpers read this module singleton without a registry injection seam + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + new=manager, + ), + patch( # test-quality-ok: unified group resolver obtains its cache from the proxy singleton + "litellm.proxy.proxy_server.user_api_key_cache", + new=UserApiKeyCache(), + ), + ): + yield prisma + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("group_identifier", "requested_identifier"), + [("server-1", "server-1"), ("server-alias", "server-1"), ("server-1", "server-alias")], +) +async def test_validate_key_servers_granted_via_team_unified_access_group_pass( + unified_mcp_prisma: MagicMock, + group_identifier: str, + requested_identifier: str, +) -> None: + unified_mcp_prisma.db.litellm_accessgrouptable.find_unique.return_value = LiteLLM_AccessGroupTable( + access_group_id="ag-1", + access_group_name="group one", + access_mcp_server_ids=[group_identifier], + ) + team: Final = LiteLLM_TeamTableCachedObj(team_id="team-1", access_group_ids=["ag-1"]) + result: Final = await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": [requested_identifier]}, + team_obj=team, + prisma_client=unified_mcp_prisma, + ) + assert result == {"mcp_servers": [requested_identifier]} + + +@pytest.mark.asyncio +async def test_validate_key_servers_outside_team_unified_access_group_rejected( + unified_mcp_prisma: MagicMock, +) -> None: + team: Final = LiteLLM_TeamTableCachedObj(team_id="team-1", access_group_ids=["ag-1"]) + with pytest.raises(HTTPException) as exc_info: + await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-2"]}, + team_obj=team, + prisma_client=unified_mcp_prisma, + ) + assert exc_info.value.status_code == 403 + assert "server-2" in str(exc_info.value.detail) + assert "Team allows: ['server-1']" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_team_allowed_servers_union_object_permission_and_unified_access_group( + unified_mcp_prisma: MagicMock, +) -> None: + team: Final = LiteLLM_TeamTableCachedObj( + team_id="team-1", + access_group_ids=["ag-1"], + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["server-2"]), + ) + result: Final = await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-1", "server-2"]}, + team_obj=team, + prisma_client=unified_mcp_prisma, + ) + assert result == {"mcp_servers": ["server-1", "server-2"]} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("group_state", ["empty", "missing", "unresolved"]) +async def test_team_unified_access_group_without_servers_preserves_direct_grants( + unified_mcp_prisma: MagicMock, + group_state: Literal["empty", "missing", "unresolved"], +) -> None: + unified_mcp_prisma.db.litellm_accessgrouptable.find_unique.return_value = ( + LiteLLM_AccessGroupTable( + access_group_id="ag-1", + access_group_name="empty or stale group", + access_mcp_server_ids=["deleted-server"] if group_state == "unresolved" else [], + ) + if group_state != "missing" + else None + ) + team: Final = LiteLLM_TeamTableCachedObj( + team_id="team-1", + access_group_ids=["ag-1"], + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["server-2"]), + ) + allowed: Final = await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-2"]}, + team_obj=team, + prisma_client=unified_mcp_prisma, + ) + assert allowed == {"mcp_servers": ["server-2"]} + with pytest.raises(HTTPException) as exc_info: + await validate_key_mcp_servers_against_team( + object_permission={"mcp_servers": ["server-1"]}, + team_obj=team, + prisma_client=unified_mcp_prisma, + ) + assert exc_info.value.status_code == 403 + assert "['server-1']. Team allows:" in str(exc_info.value.detail) + + @pytest.mark.asyncio @patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 2d929a832a5..a40741c8fdb 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -5,10 +5,11 @@ import logging import os import sys import zlib -from collections.abc import Callable +from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager +from dataclasses import dataclass from io import BytesIO -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -16,7 +17,7 @@ import httpx import pytest from fastapi import HTTPException, Request, Response, UploadFile from fastapi.responses import StreamingResponse -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError +from litellm.types import utils as types_utils from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -7305,6 +7307,156 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "vertex-gemini-38-flash-dep"} +@dataclass(frozen=True, slots=True, kw_only=True) +class _PassThroughSplit: + litellm_params: Mapping[str, object] + forwarded_body: Mapping[str, object] + + +_LITELLM_PARAMS: Final = TypeAdapter(dict[str, object]) +_PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object]) + + +def _split_pass_through_body(body: str) -> _PassThroughSplit: + mock_request: Final = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.headers = Headers() + mock_request.scope = MappingProxyType({}) + + init_kwargs_for_pass_through_endpoint: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # untyped legacy helper + kwargs: Final = init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body=json.loads(body), + litellm_call_id="lit-owned-keys-call-id", + ) + validate_litellm_params: Final = _LITELLM_PARAMS.validate_python # pyright: ignore[reportUnknownArgumentType] # untyped legacy helper + litellm_params: Final = validate_litellm_params(kwargs["litellm_params"]) + return _PassThroughSplit( + litellm_params=MappingProxyType(litellm_params), + forwarded_body=MappingProxyType( + _LITELLM_PARAMS.validate_python( + _PROXY_SERVER_REQUEST.validate_python(litellm_params["proxy_server_request"])["body"] + ) + ), + ) + + +GEMINI_BODY: Final = '{"contents": [{"parts": [{"text": "hi"}]}], "generationConfig": {"temperature": 0}}' + + +def _metadata_of(split: _PassThroughSplit) -> Mapping[str, object]: + return MappingProxyType(_LITELLM_PARAMS.validate_python(split.litellm_params["metadata"])) + + +def test_passthrough_moves_every_litellm_owned_key_from_the_forwarded_body_into_litellm_params() -> None: + split: Final = _split_pass_through_body( + '{"ttl": 30, "contents": [{"parts": [{"text": "hi"}]}], "num_retries": 2,' + ' "generationConfig": {"temperature": 0}, "litellm_trace_id": "trace-a"}' + ) + + assert frozenset(split.litellm_params) == frozenset( + ("ttl", "num_retries", "litellm_trace_id", "metadata", "proxy_server_request") + ) + assert tuple(split.litellm_params[k] for k in ("ttl", "num_retries", "litellm_trace_id")) == (30, 2, "trace-a") + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +PROXY_STAMPED_NAMES: Final = frozenset( + ( + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + ) +) + + +@pytest.mark.parametrize( + "name", + sorted(frozenset(litellm.all_litellm_params) - frozenset(("metadata", "litellm_metadata")) - PROXY_STAMPED_NAMES), +) +def test_passthrough_keeps_each_registered_litellm_owned_name_out_of_the_forwarded_body(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: "owned", **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset((name, "metadata", "proxy_server_request")) + assert split.litellm_params[name] == "owned" + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +@pytest.mark.parametrize("name", sorted(PROXY_STAMPED_NAMES)) +def test_passthrough_drops_a_client_supplied_proxy_stamped_name(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: {"forged": "by-client"}, **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset(("metadata", "proxy_server_request")) + assert split.litellm_params["proxy_server_request"] != {"forged": "by-client"}, split.litellm_params + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_merges_both_metadata_carriers_from_the_body_into_one_metadata_key() -> None: + split: Final = _split_pass_through_body( + '{"metadata": {"client_tag": "a"}, "contents": [{"parts": [{"text": "hi"}]}], "ttl": 30,' + ' "litellm_metadata": {"lm": "b"}, "generationConfig": {"temperature": 0}}' + ) + + assert frozenset(split.litellm_params) == frozenset(("ttl", "metadata", "proxy_server_request")) + assert _metadata_of(split) == {**_metadata_of(_split_pass_through_body(GEMINI_BODY)), "client_tag": "a", "lm": "b"} + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_lets_metadata_win_over_litellm_metadata_on_a_shared_key() -> None: + split: Final = _split_pass_through_body( + '{"litellm_metadata": {"shared": "from-litellm-metadata", "lm": "b"},' + ' "metadata": {"shared": "from-metadata", "client_tag": "a"}, "contents": []}' + ) + + assert _metadata_of(split) == { + **_metadata_of(_split_pass_through_body('{"contents": []}')), + "shared": "from-metadata", + "lm": "b", + "client_tag": "a", + } + + +def test_passthrough_orders_extracted_litellm_params_by_the_registry() -> None: + body: Final = json.dumps({"ttl": 30, "tags": ["team-a"], "num_retries": 2, "contents": []}) + split: Final = _split_pass_through_body(body) + body_keys: Final = frozenset(json.loads(body)) + + assert tuple(k for k in split.litellm_params if k in body_keys) == tuple( + k for k in types_utils.all_litellm_params if k in body_keys + ) + + +LATE_REGISTERED_BODY: Final = '{"registered_later": 1, "contents": [{"parts": [{"text": "hi"}]}]}' + + +def test_passthrough_sees_a_name_appended_to_the_public_list_after_import() -> None: + litellm.all_litellm_params.append("registered_later") + try: + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + finally: + litellm.all_litellm_params.remove("registered_later") + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + +def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(types_utils, "all_litellm_params", (*litellm.all_litellm_params, "registered_later")) + + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + @pytest.mark.asyncio async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 36d2e16d261..6feb37e9867 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -132,6 +132,62 @@ async def test_proxy_shutdown_event_disconnects_prisma_and_resets(monkeypatch): } +@pytest.mark.asyncio +async def test_proxy_shutdown_flushes_every_langfuse_export_channel(monkeypatch): + """A generation finished just before a graceful restart is still queued in its batch + processor, so shutdown must flush every acquired export channel.""" + from litellm.integrations.langfuse import langfuse_sdk + + flushed = MagicMock(return_value=True) + monkeypatch.setattr(langfuse_sdk, "flush_langfuse_tracing", flushed) + monkeypatch.setattr(ps, "prisma_client", None, raising=False) + monkeypatch.setattr(ps, "jwt_handler", MagicMock(close=AsyncMock()), raising=False) + monkeypatch.setattr(ps, "db_writer_client", None, raising=False) + + import litellm + + monkeypatch.setattr(litellm, "cache", None, raising=False) + monkeypatch.setattr(litellm, "success_callback", [], raising=False) + + await proxy_shutdown_event() + + assert flushed.call_count == 1 + + +@pytest.mark.asyncio +async def test_proxy_shutdown_flushes_langfuse_off_the_event_loop_and_logs_a_timeout(monkeypatch, caplog): + """The flush blocks on OTLP exports for up to its deadline, so it must run on a worker thread + with the shutdown deadline, and a channel that misses it is reported instead of ignored.""" + import threading + + from litellm.constants import LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS + from litellm.integrations.langfuse import langfuse_sdk + + ran_on = MagicMock() + + def flushed(timeout_millis: int) -> bool: + ran_on(threading.current_thread(), timeout_millis) + return False + + monkeypatch.setattr(langfuse_sdk, "flush_langfuse_tracing", flushed) + monkeypatch.setattr(ps, "prisma_client", None, raising=False) + monkeypatch.setattr(ps, "jwt_handler", MagicMock(close=AsyncMock()), raising=False) + monkeypatch.setattr(ps, "db_writer_client", None, raising=False) + + import litellm + + monkeypatch.setattr(litellm, "cache", None, raising=False) + monkeypatch.setattr(litellm, "success_callback", [], raising=False) + + with caplog.at_level("WARNING", logger="LiteLLM Proxy"): + await proxy_shutdown_event() + + (flush_thread, timeout_millis), _ = ran_on.call_args + assert flush_thread is not threading.main_thread() + assert timeout_millis == LANGFUSE_SHUTDOWN_FLUSH_TIMEOUT_MILLIS + assert any("Langfuse shutdown flush incomplete" in record.getMessage() for record in caplog.records) + + @pytest.mark.asyncio async def test_proxy_shutdown_drains_gateway_requests_before_disconnecting(monkeypatch): """ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 250556c9281..884a9c81500 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -15127,7 +15127,7 @@ async def test_auth_cache_invalidation_subscriber_evicts_byok_credentials_cached def __init__(self, client: object) -> None: self._client = client - def init_async_client(self) -> object: + def init_pubsub_client(self) -> object: return self._client byok_credential_cache.flush_cache() diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py index e333da03950..72463e17c6b 100644 --- a/tests/test_litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -3,6 +3,7 @@ Test cases for spend log cleanup functionality """ import asyncio +import logging import math import time from contextlib import asynccontextmanager @@ -1421,7 +1422,10 @@ def test_the_reported_run_outcome_is_the_most_significant_reason_in_any_order(st results into one answer: a first-match-wins implementation would pass on whichever order happened to be written and fail on its mirror. """ - results = tuple(TableCleanupResult(rows_deleted=0, stop_reason=reason) for reason in stop_reasons) + results = tuple( + TableCleanupResult(table_name=f"t{i}", rows_deleted=0, stop_reason=reason) + for i, reason in enumerate(stop_reasons) + ) assert SpendLogCleanup._run_outcome(results) == expected @@ -1545,3 +1549,91 @@ async def test_progress_reported_by_an_overlapping_run_is_its_own(monkeypatch): (error_call,) = mock_logger.error.call_args_list rendered = error_call[0][0] % error_call[0][1:] assert "(rows_deleted=100, batches=1)" in rendered + + +@pytest.mark.asyncio +async def test_spend_logs_backlog_cannot_starve_tool_index_cleanup(): + """ + Both spend-log tables share one run budget. Before the fix the spend-log + loop ran against the whole deadline, so a backlog that outlasted the budget + meant LiteLLM_SpendLogToolIndex never received a single delete batch, run + after run. The index table must still get its own share of the budget. + """ + mock_prisma_client = MagicMock() + mock_db = MagicMock() + _wire_tx(mock_db) + mock_db.execute_raw = AsyncMock(return_value=1000) + mock_prisma_client.db = mock_db + + cleaner = SpendLogCleanup( + general_settings={ + "maximum_spend_logs_retention_period": "7d", + "maximum_spend_logs_cleanup_max_batches": 500, + "maximum_spend_logs_cleanup_run_budget": "1s", + } + ) + cleaner.pod_lock_manager = None + + started_at = time.monotonic() + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + elapsed = time.monotonic() - started_at + + tables = [call[0][0].split('"')[1] for call in mock_db.execute_raw.call_args_list] + assert tables.count("LiteLLM_SpendLogs") > 0 + assert tables.count("LiteLLM_SpendLogToolIndex") > 0, "tool index cleanup was starved by the spend-log backlog" + assert elapsed < 2.5, f"splitting the budget must not extend the run: {elapsed}s" + + +@pytest.mark.asyncio +async def test_run_that_leaves_backlog_logs_a_warning_summary_naming_each_table(caplog): + """ + Operators running at warning or error level saw nothing when a run stopped + with expired rows still present. A run that ends on a bound must emit one + WARNING line that names every table, its rows deleted and its stop reason. + """ + mock_prisma_client = MagicMock() + mock_db = MagicMock() + _wire_tx(mock_db) + mock_db.execute_raw = AsyncMock(return_value=1000) + mock_prisma_client.db = mock_db + + cleaner = SpendLogCleanup( + general_settings={ + "maximum_spend_logs_retention_period": "7d", + "maximum_spend_logs_cleanup_max_batches": 2, + } + ) + cleaner.pod_lock_manager = None + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + summaries = [record for record in caplog.records if "Spend log cleanup run finished" in record.getMessage()] + assert len(summaries) == 1 + summary = summaries[0] + assert summary.levelno == logging.WARNING + message = summary.getMessage() + assert "outcome=batch_cap_reached" in message + assert "LiteLLM_SpendLogs: deleted=2000 stop_reason=batch_cap_reached" in message + assert "LiteLLM_SpendLogToolIndex: deleted=2000 stop_reason=batch_cap_reached" in message + + +@pytest.mark.asyncio +async def test_run_that_drains_every_table_logs_the_summary_at_info_not_warning(caplog): + """A healthy run must not page anyone: the summary stays at INFO.""" + mock_prisma_client = MagicMock() + mock_db = MagicMock() + _wire_tx(mock_db) + mock_db.execute_raw = AsyncMock(return_value=0) + mock_prisma_client.db = mock_db + + cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"}) + cleaner.pod_lock_manager = None + + with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"): + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + summaries = [record for record in caplog.records if "Spend log cleanup run finished" in record.getMessage()] + assert len(summaries) == 1 + assert summaries[0].levelno == logging.INFO + assert "outcome=completed" in summaries[0].getMessage() diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 0140fcaba21..08d542df16c 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -3854,6 +3854,28 @@ class TestTeamAdminEditableTeamFieldsSetting: assert stored["team_admin_editable_team_fields"] == ["projects"] assert team_admin_may_manage_projects(general_settings) is True + def test_patch_accepts_the_member_key_budgets_permission(self, monkeypatch): + from litellm.proxy.management_endpoints.team_admin_field_permissions import ( + team_admin_may_edit_member_key_budgets, + ) + + mock_prisma = self._as_proxy_admin(monkeypatch) + general_settings: dict = {} + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings) + assert team_admin_may_edit_member_key_budgets(general_settings) is False + + try: + response = client.patch( + "/update/ui_settings", json={"team_admin_editable_team_fields": ["member_key_budgets"]} + ) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 200 + stored = json.loads(mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"]["create"]["ui_settings"]) + assert stored["team_admin_editable_team_fields"] == ["member_key_budgets"] + assert team_admin_may_edit_member_key_budgets(general_settings) is True + def test_patch_with_an_empty_list_turns_team_admin_editing_off_again(self, monkeypatch): mock_prisma = self._as_proxy_admin(monkeypatch) general_settings: dict = {"team_admin_editable_team_fields": ["tpm_limit"]} diff --git a/tests/test_litellm/rag/__init__.py b/tests/test_litellm/rag/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/rag/ingestion/__init__.py b/tests/test_litellm/rag/ingestion/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/rerank_api/__init__.py b/tests/test_litellm/rerank_api/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py b/tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py index 4cc8fe78811..a2c38a898e9 100644 --- a/tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py +++ b/tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py @@ -1,6 +1,8 @@ import asyncio import gc import logging +from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -9,6 +11,7 @@ import litellm from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError from litellm.router_strategy.budget_limiter import RouterBudgetLimiting +from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.router import LiteLLM_Params from litellm.types.utils import BudgetConfig @@ -30,9 +33,7 @@ async def test_get_llm_provider_for_deployment_dict_does_not_require_litellm_par ): class RaiseOnInit: def __init__(self, *args, **kwargs): - raise AssertionError( - "LiteLLM_Params should not be instantiated in hot path" - ) + raise AssertionError("LiteLLM_Params should not be instantiated in hot path") monkeypatch.setattr( "litellm.router_strategy.budget_limiter.LiteLLM_Params", @@ -99,9 +100,7 @@ async def test_get_llm_provider_for_deployment_dict_view_supports_mapping_and_at @pytest.mark.asyncio -async def test_async_filter_deployments_resolves_provider_once_per_deployment( - disable_budget_sync, monkeypatch -): +async def test_async_filter_deployments_resolves_provider_once_per_deployment(disable_budget_sync, monkeypatch): provider_budget = RouterBudgetLimiting( dual_cache=DualCache(), provider_budget_config={ @@ -207,9 +206,7 @@ def _legacy_provider_resolution(deployment): Reference implementation used before hot-path optimization. """ try: - _litellm_params = LiteLLM_Params( - **deployment.get("litellm_params", {"model": ""}) - ) + _litellm_params = LiteLLM_Params(**deployment.get("litellm_params", {"model": ""})) _, custom_llm_provider, _, _ = litellm.get_llm_provider( model=_litellm_params.model, litellm_params=_litellm_params, @@ -228,9 +225,7 @@ def _legacy_provider_resolution(deployment): ], ) @pytest.mark.asyncio -async def test_get_llm_provider_for_deployment_matches_legacy_behavior( - disable_budget_sync, deployment -): +async def test_get_llm_provider_for_deployment_matches_legacy_behavior(disable_budget_sync, deployment): provider_budget = RouterBudgetLimiting( dual_cache=DualCache(), provider_budget_config={}, @@ -242,9 +237,7 @@ async def test_get_llm_provider_for_deployment_matches_legacy_behavior( assert current_provider == legacy_provider -def test_register_deployment_budget_for_runtime_added_deployment( - disable_budget_sync, monkeypatch -): +def test_register_deployment_budget_for_runtime_added_deployment(disable_budget_sync, monkeypatch): import asyncio monkeypatch.setattr(asyncio, "create_task", lambda coro: None) @@ -274,9 +267,7 @@ def test_register_deployment_budget_for_runtime_added_deployment( assert budget_limiter._get_budget_config_for_deployment(model_id) is None -def test_router_add_deployment_registers_deployment_budget( - disable_budget_sync, monkeypatch -): +def test_router_add_deployment_registers_deployment_budget(disable_budget_sync, monkeypatch): import asyncio from litellm import Router @@ -304,9 +295,7 @@ def test_router_add_deployment_registers_deployment_budget( budget_limiter = router._get_router_deployment_budget_limiter() assert budget_limiter is not None - config = budget_limiter._get_budget_config_for_deployment( - "runtime-budget-deployment" - ) + config = budget_limiter._get_budget_config_for_deployment("runtime-budget-deployment") assert config is not None assert config.max_budget == 0.000000000001 @@ -338,7 +327,9 @@ async def test_sync_refused_by_the_open_circuit_breaker_is_quiet_and_leaks_no_ta assert caplog.records == [] unretrieved.assert_not_called() - assert limiter.redis_increment_operation_queue == [] + assert limiter.redis_increment_operation_queue == [ + {"key": "provider_spend:openai:1d", "increment_value": 0.5, "ttl": 60} + ] assert redis_cache.async_increment_pipeline.await_count == 1 @@ -353,32 +344,34 @@ async def _limiter_with_redis(redis_cache: MagicMock) -> RouterBudgetLimiting: @pytest.mark.asyncio -async def test_push_returns_before_redis_answers(disable_budget_sync): - """The push runs inside the request success callback, so it must hand the Redis round trip to a task instead of waiting on it.""" +async def test_push_waits_for_redis_before_completing(disable_budget_sync): + redis_started = asyncio.Event() redis_answered = asyncio.Event() async def wait_for_redis(**_: object) -> None: + redis_started.set() await redis_answered.wait() redis_cache = MagicMock(spec=RedisCache) redis_cache.async_increment_pipeline = AsyncMock(side_effect=wait_for_redis) limiter = await _limiter_with_redis(redis_cache) - await asyncio.wait_for(limiter._push_in_memory_increments_to_redis(), timeout=1) - await asyncio.sleep(0) - - assert not redis_answered.is_set() + push_task = asyncio.create_task(limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(redis_started.wait(), timeout=1) + assert not push_task.done() + redis_answered.set() + assert await asyncio.wait_for(push_task, timeout=1) is True assert redis_cache.async_increment_pipeline.await_count == 1 assert limiter.redis_increment_operation_queue == [] - redis_answered.set() - await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task())) @pytest.mark.asyncio async def test_push_task_failure_is_logged_once_and_not_leaked(disable_budget_sync, caplog): """A real Redis failure on the background push must surface as one error line, never as an unretrieved task exception.""" redis_cache = MagicMock(spec=RedisCache) - redis_cache.async_increment_pipeline = AsyncMock(side_effect=ConnectionError("Error 61 connecting to 127.0.0.1:6379")) + redis_cache.async_increment_pipeline = AsyncMock( + side_effect=ConnectionError("Error 61 connecting to 127.0.0.1:6379") + ) limiter = await _limiter_with_redis(redis_cache) loop = asyncio.get_running_loop() unretrieved = MagicMock() @@ -396,3 +389,398 @@ async def test_push_task_failure_is_logged_once_and_not_leaked(disable_budget_sy "Error syncing in-memory cache with Redis: Error 61 connecting to 127.0.0.1:6379" ] unretrieved.assert_not_called() + + +_SPEND_KEY = "provider_spend:openai:1d" + + +def _increment(increment_value: float) -> RedisPipelineIncrementOperation: + return RedisPipelineIncrementOperation(key=_SPEND_KEY, increment_value=increment_value, ttl=86400) + + +class _ObservedLock(asyncio.Lock): + def __init__(self) -> None: + super().__init__() + self.waiter_started = asyncio.Event() + + async def acquire(self) -> bool: + if self.locked(): + self.waiter_started.set() + return await super().acquire() + + +class _MockRedisCache: + def __init__( + self, + initial_values: dict[str, float], + pipeline_started: asyncio.Event | None = None, + allow_pipeline_to_complete: asyncio.Event | None = None, + should_fail_pipeline: bool = False, + pipeline_completed: asyncio.Event | None = None, + read_started: asyncio.Event | None = None, + allow_read_to_complete: asyncio.Event | None = None, + ) -> None: + self.values = initial_values + self.events: list[str] = [] + self.pipeline_started = pipeline_started + self.allow_pipeline_to_complete = allow_pipeline_to_complete + self.should_fail_pipeline = should_fail_pipeline + self.pipeline_completed = pipeline_completed + self.read_started = read_started + self.allow_read_to_complete = allow_read_to_complete + + async def async_increment_pipeline( + self, increment_list: list[RedisPipelineIncrementOperation], **kwargs: object + ) -> None: + self.events.append("increment_pipeline:start") + if self.pipeline_started is not None: + self.pipeline_started.set() + if self.allow_pipeline_to_complete is not None: + await self.allow_pipeline_to_complete.wait() + if self.should_fail_pipeline: + raise RuntimeError("redis down") + for op in increment_list: + key = op["key"] + current = float(self.values.get(key, 0.0) or 0.0) + self.values[key] = current + float(op["increment_value"]) + self.events.append("increment_pipeline:done") + if self.pipeline_completed is not None: + self.pipeline_completed.set() + + async def async_batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, float | None]: + self.events.append("batch_get") + snapshot = {key: self.values.get(key) for key in key_list} + if self.read_started is not None: + self.read_started.set() + if self.allow_read_to_complete is not None: + await self.allow_read_to_complete.wait() + return snapshot + + +class _MockInMemoryCache: + def __init__(self, initial_values: dict[str, float]) -> None: + self.values = initial_values + + async def async_increment(self, key: str, value: float, ttl: int, **kwargs: object) -> float: + current = float(self.values.get(key, 0.0) or 0.0) + self.values[key] = current + float(value) + return self.values[key] + + async def async_set_cache(self, key: str, value: float, **kwargs: object) -> None: + self.values[key] = float(value) + + +def _new_router_budget_limiter( + *, + redis_cache: object, + queue_lock: asyncio.Lock | None = None, + in_memory_cache: object | None = None, + redis_increment_operation_queue: list[RedisPipelineIncrementOperation] | None = None, + provider_budget_config: dict[str, BudgetConfig] | None = None, +) -> RouterBudgetLimiting: + budget_limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting) + budget_limiter.dual_cache = SimpleNamespace( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache if in_memory_cache is not None else SimpleNamespace(), + ) + budget_limiter.provider_budget_config = provider_budget_config + budget_limiter.deployment_budget_config = None + budget_limiter.tag_budget_config = None + budget_limiter.redis_increment_operation_queue = ( + list(redis_increment_operation_queue) if redis_increment_operation_queue is not None else [] + ) + budget_limiter._redis_increment_queue_lock = queue_lock if queue_lock is not None else asyncio.Lock() + budget_limiter._redis_increment_flush_lock = asyncio.Lock() + budget_limiter._detached_increment_operations = None + return budget_limiter + + +@pytest.mark.asyncio +async def test_should_await_redis_pipeline_before_sync_reads() -> None: + pipeline_started = asyncio.Event() + allow_pipeline_to_complete = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={_SPEND_KEY: 100.0}, + pipeline_started=pipeline_started, + allow_pipeline_to_complete=allow_pipeline_to_complete, + ) + in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 160.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + redis_increment_operation_queue=[_increment(60.0)], + provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=500.0)}, + ) + + sync_task = asyncio.create_task(budget_limiter._sync_in_memory_spend_with_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + assert "batch_get" not in redis_cache.events + allow_pipeline_to_complete.set() + await sync_task + + assert redis_cache.values[_SPEND_KEY] == 160.0 + assert in_memory_cache.values[_SPEND_KEY] == 160.0 + assert budget_limiter.redis_increment_operation_queue == [] + assert redis_cache.events == [ + "increment_pipeline:start", + "increment_pipeline:done", + "batch_get", + ] + + +@pytest.mark.asyncio +async def test_should_requeue_increments_when_redis_pipeline_fails() -> None: + redis_cache = _MockRedisCache(initial_values={}, should_fail_pipeline=True) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + redis_increment_operation_queue=[_increment(10.0)], + ) + + flush_succeeded = await budget_limiter._push_in_memory_increments_to_redis() + + assert flush_succeeded is False + assert budget_limiter.redis_increment_operation_queue == [_increment(10.0)] + + +@pytest.mark.asyncio +async def test_should_keep_new_increments_when_pipeline_flush_fails() -> None: + pipeline_started = asyncio.Event() + allow_pipeline_to_complete = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={}, + pipeline_started=pipeline_started, + allow_pipeline_to_complete=allow_pipeline_to_complete, + should_fail_pipeline=True, + ) + in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 0.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + redis_increment_operation_queue=[_increment(10.0)], + ) + + push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + await budget_limiter._increment_spend_in_current_window(spend_key=_SPEND_KEY, response_cost=20.0, ttl=86400) + allow_pipeline_to_complete.set() + await push_task + + assert budget_limiter.redis_increment_operation_queue == [_increment(30.0)] + + +@pytest.mark.asyncio +async def test_failed_redis_flushes_coalesce_spend_by_key() -> None: + other_spend_key: Final = "provider_spend:other:1d" + redis_cache: Final = _MockRedisCache( + initial_values={_SPEND_KEY: 0.0, other_spend_key: 0.0}, should_fail_pipeline=True + ) + in_memory_cache: Final = _MockInMemoryCache(initial_values={_SPEND_KEY: 0.0, other_spend_key: 0.0}) + budget_limiter: Final = _new_router_budget_limiter(redis_cache=redis_cache, in_memory_cache=in_memory_cache) + + for spend_key, response_cost, ttl in ( + (_SPEND_KEY, 10.0, 90), + (other_spend_key, 4.0, 50), + (_SPEND_KEY, 20.0, 80), + (_SPEND_KEY, 30.0, 70), + ): + await budget_limiter._increment_spend_in_current_window(spend_key, response_cost, ttl) + assert await budget_limiter._push_in_memory_increments_to_redis() is False + + queued: Final = {operation["key"]: operation for operation in budget_limiter.redis_increment_operation_queue} + assert len(budget_limiter.redis_increment_operation_queue) == 2 + assert queued[_SPEND_KEY] == RedisPipelineIncrementOperation(key=_SPEND_KEY, increment_value=60.0, ttl=70) + assert queued[other_spend_key] == RedisPipelineIncrementOperation(key=other_spend_key, increment_value=4.0, ttl=50) + + redis_cache.should_fail_pipeline = False + assert await budget_limiter._push_in_memory_increments_to_redis() is True + assert redis_cache.values == {_SPEND_KEY: 60.0, other_spend_key: 4.0} + assert budget_limiter.redis_increment_operation_queue == [] + + +@pytest.mark.asyncio +async def test_should_keep_in_memory_spend_when_redis_pipeline_fails() -> None: + redis_cache = _MockRedisCache(initial_values={_SPEND_KEY: 100.0}, should_fail_pipeline=True) + in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 160.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + redis_increment_operation_queue=[_increment(60.0)], + provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=500.0)}, + ) + + await budget_limiter._sync_in_memory_spend_with_redis() + + assert in_memory_cache.values[_SPEND_KEY] == 160.0 + assert redis_cache.values[_SPEND_KEY] == 100.0 + assert budget_limiter.redis_increment_operation_queue == [_increment(60.0)] + assert "batch_get" not in redis_cache.events + + +@pytest.mark.asyncio +async def test_should_keep_increments_when_flush_is_cancelled_after_success() -> None: + pipeline_started = asyncio.Event() + allow_pipeline_to_complete = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={_SPEND_KEY: 0.0}, + pipeline_started=pipeline_started, + allow_pipeline_to_complete=allow_pipeline_to_complete, + ) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + redis_increment_operation_queue=[_increment(10.0)], + ) + + push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + push_task.cancel() + allow_pipeline_to_complete.set() + with pytest.raises(asyncio.CancelledError): + await push_task + + assert redis_cache.values[_SPEND_KEY] == 10.0 + assert budget_limiter.redis_increment_operation_queue == [] + + +@pytest.mark.asyncio +async def test_cancelled_push_waiting_for_flush_lock_still_writes_spend() -> None: + redis_cache = _MockRedisCache(initial_values={_SPEND_KEY: 0.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + redis_increment_operation_queue=[_increment(10.0)], + ) + flush_lock = _ObservedLock() + budget_limiter._redis_increment_flush_lock = flush_lock + + async with flush_lock: + push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(flush_lock.waiter_started.wait(), timeout=1) + push_task.cancel() + await asyncio.sleep(0) + assert not push_task.done() + + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(push_task, timeout=1) + + assert redis_cache.values[_SPEND_KEY] == 10.0 + assert budget_limiter.redis_increment_operation_queue == [] + + +@pytest.mark.asyncio +async def test_empty_flush_does_not_block_later_increment_sync() -> None: + redis_cache = _MockRedisCache(initial_values={_SPEND_KEY: 100.0}) + in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 100.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=500.0)}, + ) + + empty_flush_succeeded = await budget_limiter._push_in_memory_increments_to_redis() + await budget_limiter._increment_spend_in_current_window(spend_key=_SPEND_KEY, response_cost=20.0, ttl=86400) + await budget_limiter._sync_in_memory_spend_with_redis() + + assert empty_flush_succeeded is True + assert budget_limiter.redis_increment_operation_queue == [] + assert redis_cache.values[_SPEND_KEY] == 120.0 + assert in_memory_cache.values[_SPEND_KEY] == 120.0 + assert redis_cache.events == [ + "increment_pipeline:start", + "increment_pipeline:done", + "batch_get", + ] + + +@pytest.mark.asyncio +async def test_should_requeue_increments_when_flush_is_cancelled_and_redis_fails() -> None: + pipeline_started = asyncio.Event() + allow_pipeline_to_complete = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={_SPEND_KEY: 0.0}, + pipeline_started=pipeline_started, + allow_pipeline_to_complete=allow_pipeline_to_complete, + should_fail_pipeline=True, + ) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + redis_increment_operation_queue=[_increment(10.0)], + ) + + push_task = asyncio.create_task(budget_limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + push_task.cancel() + allow_pipeline_to_complete.set() + with pytest.raises(asyncio.CancelledError): + await push_task + + assert redis_cache.values[_SPEND_KEY] == 0.0 + assert budget_limiter.redis_increment_operation_queue == [_increment(10.0)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("pause_during", ["write", "read"]) +async def test_sync_preserves_spend_recorded_during_redis_io(pause_during: str) -> None: + io_started = asyncio.Event() + allow_io_to_complete = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={_SPEND_KEY: 100.0}, + pipeline_started=io_started if pause_during == "write" else None, + allow_pipeline_to_complete=allow_io_to_complete if pause_during == "write" else None, + read_started=io_started if pause_during == "read" else None, + allow_read_to_complete=allow_io_to_complete if pause_during == "read" else None, + ) + in_memory_cache = _MockInMemoryCache(initial_values={_SPEND_KEY: 160.0}) + budget_limiter = _new_router_budget_limiter( + redis_cache=redis_cache, + in_memory_cache=in_memory_cache, + redis_increment_operation_queue=[_increment(60.0)], + provider_budget_config={"openai": BudgetConfig(time_period="1d", budget_limit=175.0)}, + ) + + sync_task = asyncio.create_task(budget_limiter._sync_in_memory_spend_with_redis()) + await asyncio.wait_for(io_started.wait(), timeout=1) + await budget_limiter._increment_spend_in_current_window(_SPEND_KEY, 20.0, 86400) + allow_io_to_complete.set() + await sync_task + + assert in_memory_cache.values[_SPEND_KEY] == 180.0 + assert redis_cache.values[_SPEND_KEY] == 160.0 + assert budget_limiter.redis_increment_operation_queue == [_increment(20.0)] + + await budget_limiter._sync_in_memory_spend_with_redis() + + assert in_memory_cache.values[_SPEND_KEY] == 180.0 + assert redis_cache.values[_SPEND_KEY] == 180.0 + assert budget_limiter.redis_increment_operation_queue == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancellations", [1, 2]) +async def test_cancelled_flush_does_not_requeue_an_applied_batch(cancellations: int) -> None: + pipeline_started = asyncio.Event() + pipeline_completed = asyncio.Event() + allow_pipeline = asyncio.Event() + redis_cache = _MockRedisCache( + initial_values={_SPEND_KEY: 0.0}, + pipeline_started=pipeline_started, + pipeline_completed=pipeline_completed, + allow_pipeline_to_complete=allow_pipeline, + ) + queue_lock = _ObservedLock() + limiter = _new_router_budget_limiter( + redis_cache=redis_cache, queue_lock=queue_lock, redis_increment_operation_queue=[_increment(10.0)] + ) + push_task = asyncio.create_task(limiter._push_in_memory_increments_to_redis()) + await asyncio.wait_for(pipeline_started.wait(), timeout=1) + async with limiter._redis_increment_queue_lock: + allow_pipeline.set() + await asyncio.wait_for(pipeline_completed.wait(), timeout=1) + await asyncio.wait_for(queue_lock.waiter_started.wait(), timeout=1) + for _ in range(cancellations): + push_task.cancel() + await asyncio.sleep(0) + assert not push_task.done() + with pytest.raises(asyncio.CancelledError): + await push_task + await limiter._push_in_memory_increments_to_redis() + assert redis_cache.values[_SPEND_KEY] == 10.0 + assert limiter.redis_increment_operation_queue == [] 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/rust_bridge/messages/test_route_host.py b/tests/test_litellm/rust_bridge/messages/test_route_host.py index f47333a45d9..c5a442e0709 100644 --- a/tests/test_litellm/rust_bridge/messages/test_route_host.py +++ b/tests/test_litellm/rust_bridge/messages/test_route_host.py @@ -110,3 +110,15 @@ def test_native_request_rejections_map_to_the_public_400() -> None: assert "does not support top_k=5" in mapped.message assert mapped.model == "claude-sonnet-5" assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError) + + +def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None: + hidden: Final = route_host.stream_hidden_params( + (("request-id", "req_upstream_123"), ("x-ratelimit-remaining-requests", "41")) + ) + + additional: Final = hidden["additional_headers"] + assert isinstance(additional, dict) + assert additional["llm_provider-request-id"] == "req_upstream_123" + assert additional["x-ratelimit-remaining-requests"] == "41" + assert "request-id" not in additional 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/test_litellm/types/__init__.py b/tests/test_litellm/types/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/types/proxy/__init__.py b/tests/test_litellm/types/proxy/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/types/proxy/policy_engine/__init__.py b/tests/test_litellm/types/proxy/policy_engine/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/vector_stores/__init__.py b/tests/test_litellm/vector_stores/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/videos/__init__.py b/tests/test_litellm/videos/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm_rust/AGENTS.md b/tests/test_litellm_rust/AGENTS.md new file mode 100644 index 00000000000..d65ffd613aa --- /dev/null +++ b/tests/test_litellm_rust/AGENTS.md @@ -0,0 +1 @@ +This directory holds only the tests that cannot be written in the Rust code diff --git a/tests/test_litellm_rust/cache/__init__.py b/tests/test_litellm_rust/cache/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/tests/test_litellm_rust/cache/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/test_litellm_rust/cache/conftest.py b/tests/test_litellm_rust/cache/conftest.py new file mode 100644 index 00000000000..07ccd0fde2b --- /dev/null +++ b/tests/test_litellm_rust/cache/conftest.py @@ -0,0 +1,30 @@ +import threading +from collections.abc import Generator +from typing import Final + +import fakeredis +import pytest + +from tests.test_litellm_rust.support.s3_stub import S3Stub + + +@pytest.fixture +def redis_url() -> Generator[str]: + server: Final = fakeredis.TcpFakeServer(("127.0.0.1", 0), server_type="redis") + worker: Final = threading.Thread(target=server.serve_forever, daemon=True) + worker.start() + try: + yield f"redis://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + server.server_close() + worker.join(timeout=5) + + +@pytest.fixture +def s3_stub() -> Generator[S3Stub]: + stub: Final = S3Stub() + try: + yield stub + finally: + stub.close() diff --git a/tests/test_litellm_rust/cache/test_azure_blob.py b/tests/test_litellm_rust/cache/test_azure_blob.py new file mode 100644 index 00000000000..bbbab22baca --- /dev/null +++ b/tests/test_litellm_rust/cache/test_azure_blob.py @@ -0,0 +1,173 @@ +import asyncio +import json +import os +import time +import uuid +from collections.abc import Generator +from types import SimpleNamespace +from typing import Final, cast + +import pytest +from azure.storage.blob import ContainerClient + +from litellm.caching.azure_blob_cache import AzureBlobCache +from litellm.caching.caching import Cache +from litellm.rust_bridge import _native +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import ( + CacheLookup, + CacheTestHandle, + CacheTestResolver, + assert_native_runtime, + completion_kwargs, + request, + require_rust, +) +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +@pytest.fixture +def azure_blob_facade() -> Generator[Cache]: + account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") + if account_url is None: + pytest.skip( + "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" + ) + facade: Final = Cache( + type=LiteLLMCacheType.AZURE_BLOB, + azure_account_url=account_url, + azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", + ) + backend: Final = facade.cache + assert isinstance(backend, AzureBlobCache) + try: + yield facade + finally: + backend.container_client.delete_container() + asyncio.run(backend.disconnect()) + + +def azure_blob_handle(facade: Cache) -> _native._CacheTestHandle: + backend: Final = facade.cache + assert isinstance(backend, AzureBlobCache) + return CacheTestHandle.azure_blob( + backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}"), + backend.container_client.container_name, + ) + + +def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None: + backend: Final = azure_blob_facade.cache + assert isinstance(backend, AzureBlobCache) + handle: Final = azure_blob_handle(azure_blob_facade) + assert handle.backend == "azure-blob" + account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}") + with pytest.raises(TypeError, match="containers must match"): + CacheTestHandle.azure_blob(account_url, f"{backend.container_client.container_name}-other")._bind_facade( + azure_blob_facade + ) + handle._bind_facade(azure_blob_facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)) + native: Final = resolver.resolve() + assert native.kind == "native" + + response: Final = { + "choices": [{"text": "caf\u00e9 \u2603"}], + "usage": {"total_tokens": 3}, + "flag": True, + "empty": None, + } + native.store({**request("sync"), "ttl_seconds": 0.001}, response) + native.store(request("sync"), {"choices": [{"text": "second"}]}) + time.sleep(0.01) + stored: Final = json.loads(backend.container_client.download_blob("sync").readall()) + assert stored["response"] == response + assert isinstance(stored["timestamp"], float) + assert native.lookup(request("sync")) == response + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + + backend.set_cache("python", {"timestamp": time.time(), "response": response}) + backend.set_cache("legacy", "bare legacy value") + backend.container_client.upload_blob("invalid", b"{not json", overwrite=True) + assert native.lookup(request("python")) == response + assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy") + assert native.lookup_batch([request("python"), request("missing"), request("invalid"), request("sync")]) == { + "values": [response, None, None, response], + "missing_indices": [1, 2], + } + + with rebound(azure_blob_facade, "ttl", 12): + assert resolver.resolve().kind == "python_callback" + with rebound(backend, "container_client", ContainerClient.from_container_url(backend.container_client.url)): + assert resolver.resolve().kind == "python_callback" + + def custom_get(*_args: object, **_kwargs: object) -> None: + return None + + with rebound(backend, "get_cache", custom_get): + assert resolver.resolve().kind == "python_callback" + assert resolver.resolve().kind == "python_callback" + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + + class CustomBlobCache(AzureBlobCache): + pass + + with rebound(azure_blob_facade, "cache", CustomBlobCache(account_url, backend.container_client.container_name)): + assert resolver.resolve().kind == "python_callback" + with pytest.raises(TypeError): + azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + + +async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_python(azure_blob_facade: Cache) -> None: + backend: Final = azure_blob_facade.cache + assert isinstance(backend, AzureBlobCache) + azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)).resolve() + assert binding.kind == "native" + ping: Final = cast(dict[str, object], await binding.ping()) + assert ping["status"] == "success", ping + + await binding.async_store(request("async"), {"value": 1}) + await binding.async_store({**request("async"), "ttl_seconds": 0.001}, {"value": 2}) + time.sleep(0.01) + assert await binding.async_lookup(request("async")) == {"value": 2} + assert await backend.async_get_cache("async") == json.loads( + backend.container_client.download_blob("async").readall() + ) + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="async") == {"value": 2} + + await binding.async_store_batch([request("first"), request("second")], [{"value": 3}, {"value": 4}]) + assert await binding.async_lookup_batch([request("second"), request("missing"), request("first")]) == { + "values": [{"value": 4}, None, {"value": 3}], + "missing_indices": [1], + } + await binding.async_flush() + assert [blob.name for blob in backend.container_client.list_blobs()] == [] + assert await binding.async_lookup(request("async")) is None + + +async def test_azure_blob_rust_required_rule_activates_natively(monkeypatch: pytest.MonkeyPatch) -> None: + account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") + if account_url is None: + pytest.skip( + "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" + ) + require_rust(monkeypatch, LiteLLMCacheType.AZURE_BLOB) + facade: Final = Cache( + type=LiteLLMCacheType.AZURE_BLOB, + azure_account_url=account_url, + azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", + ) + backend: Final = facade.cache + assert isinstance(backend, AzureBlobCache) + try: + assert_native_runtime(facade) + kwargs: Final = completion_kwargs("azure") + await facade.async_add_cache({"answer": "azure"}, **kwargs) + assert await facade.async_get_cache(**kwargs) == {"answer": "azure"} + assert backend.get_cache(facade.get_cache_key(**kwargs))["response"] == {"answer": "azure"} + finally: + backend.container_client.delete_container() + await backend.disconnect() diff --git a/tests/test_litellm_rust/cache/test_disk.py b/tests/test_litellm_rust/cache/test_disk.py new file mode 100644 index 00000000000..4f2907e6a09 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_disk.py @@ -0,0 +1,117 @@ +import asyncio +import json +import time +from pathlib import Path +from types import SimpleNamespace +from typing import Final + +import diskcache +import pytest + +from litellm.caching.caching import Cache +from litellm.caching.disk_cache import DiskCache +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_path: Path) -> None: + disk_cache: Final = DiskCache(disk_cache_dir=str(tmp_path)) + response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} + disk_cache.disk_cache.set( + "sync", + {"timestamp": time.time(), "response": json.dumps(response)}, + ) + disk_cache.disk_cache.set("async", json.dumps({"timestamp": time.time(), "response": response})) + disk_cache.disk_cache.set("raw", json.dumps(response)) + disk_cache.disk_cache.set("invalid", "not a cache entry") + disk_cache.disk_cache.set( + "large", + {"timestamp": time.time(), "response": {"text": "x" * 70_000}}, + ) + binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + + assert binding.lookup(request("sync")) == response + assert await binding.async_lookup(request("async")) == response + assert binding.lookup(request("raw")) == response + assert await binding.async_lookup(request("invalid")) is None + assert binding.lookup(request("large")) == {"text": "x" * 70_000} + + await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) + stored_response: Final = disk_cache.get_cache("native") + assert isinstance(stored_response, dict) + assert stored_response["response"] == response + stored, expire_time = disk_cache.disk_cache.get("native", expire_time=True) + assert stored is not None + assert time.time() < expire_time <= time.time() + 12.0 + await binding.async_store(request("no-ttl"), response) + _, no_expiry = disk_cache.disk_cache.get("no-ttl", expire_time=True) + assert no_expiry is None + + +async def test_disk_entries_survive_a_fresh_handle_and_expire_on_time(tmp_path: Path) -> None: + first: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + await first.async_store(request("persistent"), {"value": "persistent"}) + await first.async_store({**request("expiring"), "ttl_seconds": 0.3}, {"value": "expiring"}) + fresh: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + assert fresh.lookup(request("persistent")) == {"value": "persistent"} + assert fresh.lookup(request("expiring")) == {"value": "expiring"} + await asyncio.sleep(0.4) + assert fresh.lookup(request("expiring")) is None + assert fresh.lookup(request("persistent")) == {"value": "persistent"} + + +def test_disk_facade_registers_and_store_changes_fall_back(tmp_path: Path) -> None: + facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) + with pytest.raises(TypeError, match="directories must match"): + CacheTestHandle.disk(str(tmp_path / "other"))._bind_facade(facade) + handle: Final = CacheTestHandle.disk(str(tmp_path)) + handle._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + binding: Final = resolver.resolve() + assert binding.kind == "native" + binding.store(request("native"), {"value": "native"}) + assert facade.get_cache(cache_key="native") == {"value": "native"} + + with rebound(facade.cache, "disk_cache", diskcache.Cache(str(tmp_path))): + assert resolver.resolve().kind == "python_callback" + assert resolver.resolve().kind == "native" + + class CustomDiskCache(DiskCache): + pass + + with rebound(facade, "cache", CustomDiskCache(disk_cache_dir=str(tmp_path))): + assert resolver.resolve().kind == "python_callback" + + class CustomStore(diskcache.Cache): + pass + + custom_facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) + custom_facade.cache.disk_cache = CustomStore(str(tmp_path)) + with pytest.raises(TypeError, match="built-in diskcache store"): + CacheTestHandle.disk(str(tmp_path))._bind_facade(custom_facade) + + +async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path: Path) -> None: + binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + requests: Final = [request("hit"), request("miss"), request("disabled")] + requests[2]["controls"] = { + "supported_call_type": True, + "configured": True, + "native_backend": True, + "default_on": True, + "caching": False, + "no_cache": False, + "no_store": False, + "use_cache": False, + } + await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}]) + + partial: Final = await binding.async_lookup_batch(requests) + + assert partial == { + "values": [{"value": 1}, {"value": 2}, None], + "missing_indices": [2], + } diff --git a/tests/test_litellm_rust/cache/test_facade.py b/tests/test_litellm_rust/cache/test_facade.py new file mode 100644 index 00000000000..d99ea4e2baa --- /dev/null +++ b/tests/test_litellm_rust/cache/test_facade.py @@ -0,0 +1,397 @@ +import asyncio +import contextvars +import gc +import weakref +from types import SimpleNamespace +from typing import Final, cast + +import pytest + +import litellm +from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.rust_bridge import _native +from litellm.rust_bridge.catalog import CacheRule, Route, RouteRule, SecretManagerRule +from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.response_cache import ResponseCacheRuntime, resolve_response_cache +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import CacheLookup, CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +def test_existing_constructor_and_global_are_unchanged() -> None: + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + assert type(facade.cache) is InMemoryCache + assert "_native_cache_handle" not in vars(facade) + assert resolve_response_cache(facade) is None + with rebound(litellm, "cache", facade): + resolver: Final = CacheTestResolver(litellm) + assert resolver.resolve().kind == "python_callback" + resolver.resolve().store(None, {"answer": 7}, callback_kwargs={"cache_key": "key"}) + assert cast(CacheLookup, facade).get_cache(cache_key="key") == {"answer": 7} + + +async def test_catalog_constructs_native_runtime_from_public_cache_configuration() -> None: + rules: Final = ( + RouteRule(Route.OCR, Rollout.PYTHON_ONLY), + SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), + CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), + ) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + runtime: Final = resolve_response_cache(facade, rules) + assert isinstance(runtime, ResponseCacheRuntime) + assert runtime.kind == "native" + + sync_request: Final = runtime.request(facade, {"cache_key": "sync"}) + assert sync_request is not None + runtime.store(sync_request, {"answer": 1}) + assert runtime.lookup(sync_request) == {"answer": 1} + assert facade.cache.get_cache("sync") is None + + async_request: Final = runtime.request(facade, {"cache_key": "async"}) + assert async_request is not None + await runtime.async_store(async_request, {"answer": 2}) + assert await runtime.async_lookup(async_request) == {"answer": 2} + assert await facade.cache.async_get_cache("async") is None + + requests: Final = (sync_request, async_request) + expected: Final = { + "values": [{"answer": 1}, {"answer": 2}], + "missing_indices": [], + } + assert runtime.lookup_batch(requests) == expected + assert await runtime.async_lookup_batch(requests) == expected + + await runtime.async_flush() + assert runtime.lookup(sync_request) is None + assert await runtime.async_lookup(async_request) is None + + +async def test_inference_resolver_uses_the_configured_native_cache_directly() -> None: + rules: Final = ( + RouteRule(Route.OCR, Rollout.PYTHON_ONLY), + SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), + CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), + ) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + runtime: Final = resolve_response_cache(facade, rules) + assert isinstance(runtime, ResponseCacheRuntime) + facade._native_cache = runtime + + selected: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + assert selected.kind == "native" + request: Final = runtime.request(facade, {"cache_key": "inference-native"}) + assert request is not None + await selected.async_store(request, {"answer": 42}) + assert await selected.async_lookup(request) == {"answer": 42} + assert await runtime.async_lookup(request) == {"answer": 42} + assert facade.cache.get_cache("inference-native") is None + + facade._native_cache = None + fallback: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + assert fallback.kind == "python_callback" + await fallback.async_store(None, {"answer": 7}, callback_kwargs={"cache_key": "inference-python"}) + assert facade.get_cache(cache_key="inference-python") == {"answer": 7} + assert facade.cache.get_cache("inference-python") is not None + + +async def test_inference_resolver_declines_a_native_runtime_whose_facade_changed() -> None: + rules: Final = ( + RouteRule(Route.OCR, Rollout.PYTHON_ONLY), + SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), + CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), + ) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + runtime: Final = resolve_response_cache(facade, rules) + assert isinstance(runtime, ResponseCacheRuntime) + facade._native_cache = runtime + stale_request: Final = runtime.request(facade, {"cache_key": "stale-only"}) + assert stale_request is not None + await runtime.async_store(stale_request, {"answer": "stale"}) + + replacement: Final = InMemoryCache() + facade.cache = replacement + with pytest.raises(_native.RustBridgeDeclined): + _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + assert await runtime.async_lookup(stale_request) == {"answer": "stale"} + assert replacement.get_cache("stale-only") is None + assert replacement.get_cache("swapped-backend") is None + + +def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> None: + resolver: Final = CacheTestResolver(litellm) + + enable_cache(type=LiteLLMCacheType.LOCAL, ttl=30) + enabled: Final = litellm.cache + assert isinstance(enabled, Cache) + assert enabled.ttl == 30 + assert resolver.resolve().kind == "python_callback" + + enable_cache(type=LiteLLMCacheType.LOCAL, ttl=60) + assert litellm.cache is enabled + + update_cache(type=LiteLLMCacheType.LOCAL, ttl=60) + updated: Final = litellm.cache + assert isinstance(updated, Cache) + assert updated is not enabled + assert updated.ttl == 60 + + disable_cache() + assert litellm.cache is None + assert resolver.resolve().kind == "disabled" + + +async def test_native_bindings_survive_replacement_and_capture_writes_before_dispatch() -> None: + namespace: Final = SimpleNamespace(cache=CacheTestHandle.memory()) + resolver: Final = CacheTestResolver(namespace) + selected: Final = resolver.resolve() + assert selected.kind == "native" + selected.store(request(), {"answer": 1}) + assert await selected.async_lookup(request()) == {"answer": 1} + with rebound(namespace, "cache", CacheTestHandle.memory()): + replacement: Final = resolver.resolve() + await selected.async_store(request(), {"answer": 2}) + assert replacement.lookup(request()) is None + assert selected.lookup(request()) == {"answer": 2} + with rebound(namespace, "cache", None): + disabled: Final = resolver.resolve() + assert disabled.kind == "disabled" + assert disabled.lookup(None) is None + await disabled.async_store(None, object()) + assert await disabled.async_lookup(None) is None + assert selected.lookup(request()) == {"answer": 2} + + +async def test_python_callback_preserves_identity_caller_task_context_and_errors() -> None: + context: Final = contextvars.ContextVar("cache_context", default="caller") + caller: Final = asyncio.current_task() + sentinel: Final = object() + failure: Final = RuntimeError("callback failed") + + class CustomCache: + async def async_get_cache(self, *, marker: object) -> object: + assert marker is sentinel + assert asyncio.current_task() is caller + context.set("callback") + return marker + + async def async_add_cache(self, response: object, *, marker: object) -> None: + assert response is sentinel + assert marker is sentinel + raise failure + + namespace: Final = SimpleNamespace(cache=CustomCache()) + binding: Final = CacheTestResolver(namespace).resolve() + assert binding.kind == "python_callback" + assert await binding.async_lookup(None, callback_kwargs={"marker": sentinel}) is sentinel + assert context.get() == "callback" + with pytest.raises(RuntimeError) as caught: + await binding.async_store(None, sentinel, callback_kwargs={"marker": sentinel}) + assert caught.value is failure + + +async def test_callback_cancellation_stays_in_the_callers_task() -> None: + entered: Final = asyncio.Event() + finished: Final = asyncio.Event() + + class CustomCache: + async def async_get_cache(self) -> None: + entered.set() + try: + await asyncio.Future() + finally: + finished.set() + + binding: Final = CacheTestResolver(SimpleNamespace(cache=CustomCache())).resolve() + + async def lookup() -> object: + return await binding.async_lookup(None, callback_kwargs={}) + + task: Final = asyncio.create_task(lookup()) + await entered.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert finished.is_set() + + +def test_registered_facade_uses_native_and_instance_overrides_fall_back() -> None: + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + handle: Final = CacheTestHandle.memory() + handle._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + native: Final = resolver.resolve() + assert native.kind == "native" + native.store(request(), {"source": "native"}) + assert native.lookup(request()) == {"source": "native"} + assert cast(CacheLookup, facade).get_cache(cache_key="key") is None + sentinel: Final = object() + + def outer_override(**_kwargs: object) -> object: + return sentinel + + def backend_override(*_args: object, **_kwargs: object) -> dict[str, str]: + return {"source": "override"} + + with rebound(facade, "get_cache", outer_override): + fallback: Final = resolver.resolve() + assert fallback.kind == "python_callback" + assert fallback.lookup(None, callback_kwargs={"cache_key": "key"}) is sentinel + assert resolver.resolve().kind == "python_callback" + delattr(facade, "get_cache") + assert resolver.resolve().kind == "native" + with rebound(facade.cache, "get_cache", backend_override): + backend_fallback: Final = resolver.resolve() + assert backend_fallback.kind == "python_callback" + assert backend_fallback.lookup(None, callback_kwargs={"cache_key": "key"}) == {"source": "override"} + + +def test_facade_subclasses_backend_replacement_and_configuration_changes_are_not_bypassed() -> None: + class CustomCache(Cache): + pass + + handle: Final = CacheTestHandle.memory() + with pytest.raises(TypeError): + handle._bind_facade(CustomCache(type=LiteLLMCacheType.LOCAL)) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + handle._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + with rebound(facade, "cache", InMemoryCache()): + assert resolver.resolve().kind == "python_callback" + with rebound(facade, "ttl", 12): + assert resolver.resolve().kind == "python_callback" + with rebound(facade, "semantic_cache_scope", "end_user"): + assert resolver.resolve().kind == "python_callback" + + def custom_key(**_kwargs: object) -> str: + return "custom" + + with rebound(facade, "get_cache_key", custom_key): + assert resolver.resolve().kind == "python_callback" + assert resolver.resolve().kind == "python_callback" + delattr(facade, "get_cache_key") + assert resolver.resolve().kind == "native" + + +def test_resolver_and_callback_cycles_can_be_collected() -> None: + class CustomCache: + pass + + def cyclic_reference() -> weakref.ReferenceType[CustomCache]: + callback: Final = CustomCache() + namespace: Final = SimpleNamespace(cache=callback) + binding: Final = CacheTestResolver(namespace).resolve() + setattr(callback, "binding", binding) + return weakref.ref(callback) + + reference: Final = cyclic_reference() + gc.collect() + assert reference() is None + + +def test_invalid_duration_and_request_shape_fail_before_storage() -> None: + binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory())).resolve() + for seconds in (-1.0, float("nan"), float("inf")): + with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): + binding.store({**request(), "ttl_seconds": seconds}, {"answer": 1}) + assert binding.lookup(request()) is None + with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): + CacheTestHandle.memory(ttl_seconds=-1) + + +async def test_memory_size_policy_is_applied_by_the_native_host() -> None: + handle: Final = CacheTestHandle.memory(capacity=2, max_entry_bytes=128) + binding: Final = CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + small: Final = {"answer": "ok"} + binding.store(request("small"), small) + assert await binding.async_lookup(request("small")) == small + await binding.async_store(request("large"), {"answer": "x" * 256}) + assert binding.lookup(request("large")) is None + assert binding.lookup(request("small")) == small + disabled: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory(capacity=0))).resolve() + await disabled.async_store(request(), small) + assert await disabled.async_lookup(request()) is None + + +async def test_native_batch_lookup_and_store_report_partial_hits() -> None: + binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory())).resolve() + requests: Final = [request("hit"), request("miss"), request("disabled")] + requests[2]["controls"] = { + "supported_call_type": True, + "configured": True, + "native_backend": True, + "default_on": True, + "caching": False, + "no_cache": False, + "no_store": False, + "use_cache": False, + } + await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}]) + + partial: Final = await binding.async_lookup_batch(requests) + + assert partial == { + "values": [{"value": 1}, {"value": 2}, None], + "missing_indices": [2], + } + + +async def test_python_batch_callbacks_use_the_builtin_cache_api() -> None: + result: Final = object() + marker: Final = object() + + class CustomCache(Cache): + def get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object: + return ("sync", kwargs) + + async def async_get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object: + return ("async", kwargs) + + async def async_add_cache_pipeline( + self, result: object, dynamic_cache_object: object = None, **kwargs: object + ) -> object: + return result, kwargs + + binding: Final = CacheTestResolver(SimpleNamespace(cache=CustomCache(type=LiteLLMCacheType.LOCAL))).resolve() + assert binding.kind == "python_callback" + requests: Final = [request("first"), request("second")] + kwargs: Final = [{"cache_key": "first"}, {"cache_key": "second"}] + + assert binding.lookup_batch(requests, callback_kwargs=kwargs) == [("sync", kwargs[0]), ("sync", kwargs[1])] + assert await binding.async_lookup_batch(requests, callback_kwargs=kwargs) == [ + ("async", kwargs[0]), + ("async", kwargs[1]), + ] + with pytest.raises(ValueError, match="equal lengths"): + binding.lookup_batch(requests, callback_kwargs=kwargs[:1]) + with pytest.raises(TypeError, match="callback_result"): + await binding.async_store_batch(requests, [1, 2], callback_kwargs={"marker": marker}) + stored: Final = cast( + tuple[object, dict[str, object]], + await binding.async_store_batch(requests, [1, 2], callback_result=result, callback_kwargs={"marker": marker}), + ) + assert stored[0] is result + assert stored[1] == {"marker": marker} + + +async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None: + async def ping() -> str: + return "pong" + + cache: Final = Cache(type=LiteLLMCacheType.LOCAL) + cache.cache.set_cache("key", "value") + binding: Final = CacheTestResolver(SimpleNamespace(cache=cache)).resolve() + assert binding.kind == "python_callback" + + setattr(cache.cache, "ping", ping) + assert await binding.ping() == "pong" + await binding.async_flush() + assert cache.cache.get_cache("key") is None + + +def test_facade_registration_rejects_mismatched_capacity() -> None: + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + with pytest.raises(TypeError, match="capacities must match"): + CacheTestHandle.memory(capacity=7)._bind_facade(facade) diff --git a/tests/test_litellm_rust/cache/test_gcs.py b/tests/test_litellm_rust/cache/test_gcs.py new file mode 100644 index 00000000000..bfc9ebbb4d7 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_gcs.py @@ -0,0 +1,242 @@ +import json +import time +from collections.abc import Generator +from types import SimpleNamespace +from typing import Final, cast + +import pytest + +from litellm.caching.caching import Cache +from litellm.caching.gcs_cache import GCSCache +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import CacheLookup, CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.fake_gcs import FakeGcs +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +@pytest.fixture +def fake_gcs() -> Generator[FakeGcs]: + server: Final = FakeGcs() + try: + yield server + finally: + server.close() + + +async def test_gcs_reads_python_entries_and_writes_python_compatible_objects( + fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} + fake_gcs.put( + "bucket", + "cache/sync", + json.dumps({"timestamp": time.time(), "response": json.dumps(response)}).encode(), + ) + fake_gcs.put("bucket", "cache/async", json.dumps({"timestamp": time.time(), "response": response}).encode()) + fake_gcs.put("bucket", "cache/raw", json.dumps(response).encode()) + fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") + binding: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + ) + ).resolve() + + assert binding.lookup(request("sync")) == response + assert await binding.async_lookup(request("async")) == response + assert binding.lookup(request("raw")) == response + assert await binding.async_lookup(request("invalid")) is None + assert binding.lookup(request("missing")) is None + + await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) + stored: Final = fake_gcs.objects[("bucket", "cache/native")] + stored_value: Final = cast(dict[str, object], json.loads(stored)) + assert stored_value["response"] == response + assert isinstance(stored_value["timestamp"], float) + upload: Final = next(item for item in fake_gcs.requests if item.method == "POST") + assert upload.path == "/upload/storage/v1/b/bucket/o" + assert upload.query == "uploadType=media&name=cache%2Fnative" + assert upload.headers["Authorization"] == f"Bearer {fake_gcs.token}" + assert upload.headers["Content-Type"] == "application/json" + upload_text: Final = f"{upload.path}?{upload.query}{upload.headers}" + assert "ttl" not in upload_text.lower() + assert "expiry" not in upload_text.lower() + download: Final = next(item for item in fake_gcs.requests if item.path.endswith("/cache%2Fsync")) + assert download.path == "/storage/v1/b/bucket/o/cache%2Fsync" + assert download.query == "alt=media" + + binding.store(request("sync2"), response) + assert binding.lookup(request("sync2")) == response + assert GCSCache(bucket_name="bucket", gcs_path="cache").key_prefix == "cache/" + assert GCSCache(bucket_name="bucket", gcs_path="cache/").key_prefix == "cache/" + assert GCSCache(bucket_name="bucket").key_prefix == "" + + +async def test_gcs_batch_lookup_preserves_order_and_treats_malformed_entries_as_misses(fake_gcs: FakeGcs) -> None: + fake_gcs.put("bucket", "cache/hit", json.dumps({"timestamp": time.time(), "response": {"value": 1}}).encode()) + fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") + binding: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + ) + ).resolve() + requests: Final = [request("hit"), request("missing"), request("invalid")] + expected: Final = {"values": [{"value": 1}, None, None], "missing_indices": [1, 2]} + + assert await binding.async_lookup_batch(requests) == expected + assert binding.lookup_batch(requests) == expected + await binding.async_store_batch([request("first"), request("second")], [{"value": 1}, {"value": 2}]) + assert ("bucket", "cache/first") in fake_gcs.objects + assert ("bucket", "cache/second") in fake_gcs.objects + + +async def test_gcs_facade_binds_only_exact_matching_configuration( + fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/nonexistent") + facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") + assert type(facade.cache) is GCSCache + + mismatched_bucket: Final = CacheTestHandle.gcs( + "other", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + with pytest.raises(TypeError, match="buckets must match"): + mismatched_bucket._bind_facade(facade) + mismatched_prefix: Final = CacheTestHandle.gcs( + "bucket", + gcs_path="x", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + with pytest.raises(TypeError, match="key prefixes must match"): + mismatched_prefix._bind_facade(facade) + mismatched_credentials: Final = CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + path_service_account="sa.json", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + with pytest.raises(TypeError, match="credentials must match"): + mismatched_credentials._bind_facade(facade) + with pytest.raises(TypeError, match="types must match"): + CacheTestHandle.memory()._bind_facade(facade) + + matching: Final = CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + matching._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + binding: Final = resolver.resolve() + assert binding.kind == "native" + await binding.async_store(request("native"), {"value": "native"}) + assert await binding.async_lookup(request("native")) == {"value": "native"} + assert cast(CacheLookup, facade).get_cache(cache_key="native") is None + + with rebound(facade.cache, "bucket_name", "other"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "key_prefix", "x/"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "path_service_account", "sa.json"): + assert resolver.resolve().kind == "python_callback" + + def no_get_cache(*args: object, **kwargs: object) -> None: + return None + + with rebound(facade.cache, "get_cache", no_get_cache): + assert resolver.resolve().kind == "python_callback" + with rebound(facade, "ttl", 12): + assert resolver.resolve().kind == "python_callback" + + class CustomGcs(GCSCache): + pass + + with rebound(facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): + assert resolver.resolve().kind == "python_callback" + custom_facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") + with rebound(custom_facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): + with pytest.raises(TypeError, match="types must match"): + matching._bind_facade(custom_facade) + + missing_bucket: Final = Cache(type=LiteLLMCacheType.GCS) + with pytest.raises(TypeError, match="requires a configured bucket name"): + matching._bind_facade(missing_bucket) + + +async def test_gcs_flush_is_a_no_op_and_ping_is_not_implemented( + fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + binding: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + ) + ).resolve() + await binding.async_store(request("key"), {"value": "stored"}) + await binding.async_flush() + assert ("bucket", "cache/key") in fake_gcs.objects + assert await binding.async_lookup(request("key")) == {"value": "stored"} + with pytest.raises(NotImplementedError): + await binding.ping() + + facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") + with pytest.raises(AttributeError): + await facade.ping() + assert cast(CacheLookup, facade.cache).flush_cache() is None + + +async def test_gcs_unauthorized_and_server_errors_surface_as_runtime_errors(fake_gcs: FakeGcs) -> None: + wrong_token: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token="wrong-token", + ) + ) + ).resolve() + with pytest.raises(RuntimeError): + wrong_token.lookup(request("missing")) + assert not fake_gcs.objects + + binding: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.gcs( + "bucket", + gcs_path="cache", + endpoint=fake_gcs.url, + token=fake_gcs.token, + ) + ) + ).resolve() + with pytest.raises(RuntimeError): + binding.lookup(request("server-error")) + assert binding.lookup(request("missing")) is None diff --git a/tests/test_litellm_rust/cache/test_qdrant_semantic.py b/tests/test_litellm_rust/cache/test_qdrant_semantic.py new file mode 100644 index 00000000000..160089c9002 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_qdrant_semantic.py @@ -0,0 +1,286 @@ +import hashlib +import http.server +import json +import math +import os +import threading +import time +from collections.abc import Generator +from types import SimpleNamespace +from typing import Final +from uuid import uuid4 + +import pytest + +from litellm.caching.caching import Cache +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import ( + CacheTestHandle, + CacheTestResolver, + assert_native_runtime, + request, + require_rust, +) + +pytestmark: Final = pytest.mark.requires_rust_extension + + +def qdrant_request( + key: str, + messages: list[dict[str, object]], + **kwargs: object, +) -> dict[str, object]: + return {**request(key), "messages": messages, **kwargs} + + +def embedding_vector(text: str) -> list[float]: + raw: Final = hashlib.sha256(text.encode()).digest()[:8] + values: Final = [byte / 127.5 - 1 for byte in raw] + norm: Final = math.sqrt(sum(value * value for value in values)) + return [value / norm for value in values] + + +@pytest.fixture +def qdrant_url() -> str: + value: Final[str | None] = os.environ.get("QDRANT_URL") + if not value: + pytest.skip("QDRANT_URL is required for Qdrant semantic cache tests") + return value.rstrip("/") + + +@pytest.fixture +def fake_embedding_endpoint(monkeypatch: pytest.MonkeyPatch) -> Generator[str]: + class EmbeddingHandler(http.server.BaseHTTPRequestHandler): + def do_POST(self) -> None: + length: Final = int(self.headers["Content-Length"]) + body: Final = json.loads(self.rfile.read(length)) + text: Final = body["input"] + response: Final = { + "object": "list", + "data": [ + { + "object": "embedding", + "index": 0, + "embedding": embedding_vector(text), + } + ], + "model": body["model"], + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + encoded: Final = json.dumps(response).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(encoded))) + self.end_headers() + self.wfile.write(encoded) + + def log_message(self, *_args: object) -> None: + return + + server: Final = http.server.ThreadingHTTPServer(("127.0.0.1", 0), EmbeddingHandler) + worker: Final = threading.Thread(target=server.serve_forever, daemon=True) + worker.start() + monkeypatch.setenv("OPENAI_API_BASE", f"http://127.0.0.1:{server.server_address[1]}") + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + try: + yield f"http://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + server.server_close() + worker.join(timeout=5) + + +def qdrant_facade(qdrant_url: str, collection_name: str) -> Cache: + return Cache( + type=LiteLLMCacheType.QDRANT_SEMANTIC, + qdrant_api_base=qdrant_url, + qdrant_collection_name=collection_name, + similarity_threshold=0.99, + qdrant_semantic_cache_embedding_model="text-embedding-3-small", + qdrant_semantic_cache_vector_size=8, + ) + + +def test_qdrant_semantic_facade_binds_native_and_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None: + del fake_embedding_endpoint + messages: Final = [{"role": "user", "content": "shared prompt"}] + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + facade.cache.set_cache( + "python-key", + {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, + messages=messages, + ) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + assert binding.lookup(qdrant_request("python-key", messages)) == {"id": "py"} + binding.store(qdrant_request("native-key", messages), {"id": "native"}) + python_value: Final = facade.cache.get_cache("native-key", messages=messages) + assert isinstance(python_value, dict) + assert python_value["response"] == {"id": "native"} + unrelated: Final = [{"role": "user", "content": "unrelated prompt"}] + assert binding.lookup(qdrant_request("native-key", unrelated)) is None + assert facade.cache.get_cache("native-key", messages=unrelated) is None + assert binding.lookup(qdrant_request("different-key", messages)) is None + assert facade.cache.get_cache("different-key", messages=messages) is None + + +async def test_qdrant_semantic_async_parity(qdrant_url: str, fake_embedding_endpoint: str) -> None: + del fake_embedding_endpoint + messages: Final = [{"role": "user", "content": "async prompt"}] + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + await facade.cache.async_set_cache( + "python-key", + {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, + messages=messages, + ) + assert await binding.async_lookup(qdrant_request("python-key", messages)) == {"id": "py"} + await binding.async_store(qdrant_request("native-key", messages), {"id": "native"}) + python_value: Final = await facade.cache.async_get_cache("native-key", messages=messages) + assert isinstance(python_value, dict) + assert python_value["response"] == {"id": "native"} + + +async def test_qdrant_semantic_async_store_batch_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None: + del fake_embedding_endpoint + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + entries: Final = [ + qdrant_request("batch-one", [{"role": "user", "content": "first batch prompt"}]), + qdrant_request("batch-two", [{"role": "user", "content": "second batch prompt"}]), + ] + await binding.async_store_batch(entries, [{"id": "one"}, {"id": "two"}]) + + assert binding.lookup(entries[0]) == {"id": "one"} + assert binding.lookup(entries[1]) == {"id": "two"} + assert (await facade.cache.async_get_cache("batch-one", messages=entries[0]["messages"]))["response"] == { + "id": "one" + } + assert (await facade.cache.async_get_cache("batch-two", messages=entries[1]["messages"]))["response"] == { + "id": "two" + } + + +async def test_qdrant_semantic_malformed_entries_and_unsupported_operations( + qdrant_url: str, fake_embedding_endpoint: str +) -> None: + del fake_embedding_endpoint + messages: Final = [{"role": "user", "content": "malformed prompt"}] + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + key: Final = "malformed-key" + response: Final = { + "points": [ + { + "id": str(uuid4()), + "vector": embedding_vector("malformed prompt"), + "payload": { + "litellm_cache_key": key, + "text": "malformed prompt", + "response": "not json", + }, + } + ] + } + facade.cache.sync_client.put( + url=f"{qdrant_url}/collections/{collection}/points", + headers=facade.cache.headers, + json=response, + ) + assert binding.lookup(qdrant_request(key, messages)) is None + with pytest.raises(RuntimeError, match="operation is not supported"): + binding.lookup_batch([qdrant_request(key, messages)]) + with pytest.raises(RuntimeError, match="operation is not supported"): + await binding.async_flush() + with pytest.raises(RuntimeError, match="operation is not supported"): + await binding.ping() + + +def test_qdrant_semantic_ignores_request_expiry(qdrant_url: str, fake_embedding_endpoint: str) -> None: + del fake_embedding_endpoint + messages: Final = [{"role": "user", "content": "persistent prompt"}] + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + binding.store(qdrant_request("persistent-key", messages, ttl_seconds=1.0), {"id": "persistent"}) + time.sleep(1.2) + assert binding.lookup(qdrant_request("persistent-key", messages)) == {"id": "persistent"} + python_value: Final = facade.cache.get_cache("persistent-key", messages=messages) + assert isinstance(python_value, dict) + assert python_value["response"] == {"id": "persistent"} + + +def test_qdrant_semantic_mutation_and_projection_fallback(qdrant_url: str, fake_embedding_endpoint: str) -> None: + del fake_embedding_endpoint + collection: Final = f"cache_{uuid4().hex}" + facade: Final = qdrant_facade(qdrant_url, collection) + handle: Final = CacheTestHandle.qdrant_semantic( + qdrant_url, + collection_name=collection, + similarity_threshold=0.99, + vector_size=8, + ) + handle._bind_facade(facade) + facade.cache.qdrant_api_key = "rotated" + assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + facade.cache.similarity_threshold = 0.5 + assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + unsupported: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") + unsupported.cache.embedding_max_input_tokens = 100 + with pytest.raises(TypeError, match="requires Python"): + handle._bind_facade(unsupported) + unsupported.cache.embedding_max_input_tokens = None + unsupported.cache.qdrant_api_base = "http://127.0.0.1:7777" + with pytest.raises(TypeError, match="gRPC"): + handle._bind_facade(unsupported) + + +def test_qdrant_semantic_rust_required_rule_activates_natively( + qdrant_url: str, fake_embedding_endpoint: str, monkeypatch: pytest.MonkeyPatch +) -> None: + del fake_embedding_endpoint + require_rust(monkeypatch, LiteLLMCacheType.QDRANT_SEMANTIC) + facade: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") + assert_native_runtime(facade) + kwargs: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "qdrant activation"}]} + facade.add_cache({"answer": "qdrant"}, **kwargs) + assert facade.get_cache(**kwargs) == {"answer": "qdrant"} diff --git a/tests/test_litellm_rust/cache/test_redis.py b/tests/test_litellm_rust/cache/test_redis.py new file mode 100644 index 00000000000..dd88145ef21 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_redis.py @@ -0,0 +1,228 @@ +import json +import os +import time +from types import SimpleNamespace +from typing import Final +from urllib.parse import urlparse + +import pytest +import redis + +import litellm +from litellm.caching.caching import Cache +from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.rust_bridge import catalog +from litellm.rust_bridge.catalog import CacheRule +from litellm.rust_bridge.configuration import Rollout +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import ( + CacheTestHandle, + CacheTestResolver, + assert_native_runtime, + completion_kwargs, + request, + require_rust, +) +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +@pytest.fixture +def cluster_nodes() -> tuple[tuple[str, int], ...]: + configured: Final = os.environ.get("LITELLM_TEST_REDIS_CLUSTER_NODES") + if not configured: + pytest.skip("LITELLM_TEST_REDIS_CLUSTER_NODES is not set") + return tuple((host, int(port)) for host, _, port in (node.partition(":") for node in configured.split(","))) + + +async def test_redis_reads_python_sync_and_async_entries_and_writes_without_hidden_prefix(redis_url: str) -> None: + client: Final = redis.Redis.from_url(redis_url) + namespace: Final = SimpleNamespace(cache=CacheTestHandle.redis(redis_url, namespace="team")) + binding: Final = CacheTestResolver(namespace).resolve() + response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} + envelope: Final = {"timestamp": time.time(), "response": json.dumps(response)} + client.set("team:sync", str(envelope)) + client.set("team:async", json.dumps({"timestamp": time.time(), "response": response})) + client.set("team:raw", json.dumps(response)) + client.set("team:invalid", "not a cache entry") + assert binding.lookup(request("sync")) == response + assert await binding.async_lookup(request("team:async")) == response + assert binding.lookup(request("raw")) == response + assert await binding.async_lookup(request("invalid")) is None + await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) + stored: Final = client.get("team:native") + assert isinstance(stored, bytes) + assert json.loads(stored)["response"] == response + assert 0 < client.ttl("team:native") <= 12 + assert client.get("litellm-cache:team:native") is None + assert client.get("team:team:async") is None + client.close() + + +async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: + parsed: Final = urlparse(redis_url) + with rebound(litellm, "default_redis_ttl", 60): + facade: Final = Cache( + type=LiteLLMCacheType.REDIS, + host=parsed.hostname, + port=str(parsed.port), + redis_flush_size=2, + ) + with pytest.raises(TypeError, match="default TTLs must match"): + CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade) + with pytest.raises(TypeError, match="namespaces must match"): + CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade) + CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(redis_url) + + with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "ssl": True}): + assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + + pool: Final = facade.cache.redis_client.connection_pool + with rebound(pool, "connection_kwargs", {**pool.connection_kwargs, "db": 1}): + assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + + await binding.async_store(request("first"), {"value": 1}) + assert client.get("first") is None + await binding.async_store(request("second"), {"value": 2}) + + assert client.get("first") is not None + assert client.get("second") is not None + await facade.cache.disconnect() + client.close() + + +async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_natively( + cluster_nodes: tuple[tuple[str, int], ...], +) -> None: + startup_nodes: Final = [{"host": host, "port": port} for host, port in cluster_nodes] + url: Final = f"redis://{cluster_nodes[0][0]}:{cluster_nodes[0][1]}" + with rebound(litellm, "default_redis_ttl", 60): + facade: Final = Cache(type=LiteLLMCacheType.REDIS, redis_startup_nodes=startup_nodes, namespace="parity") + assert type(facade.cache) is RedisClusterCache + with pytest.raises(TypeError, match="types must match"): + CacheTestHandle.redis(url, namespace="parity")._bind_facade(facade) + CacheTestHandle.redis(url, namespace="parity", startup_nodes=list(cluster_nodes))._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "native" + + manager: Final = facade.cache.redis_client.nodes_manager + with rebound(manager, "connection_kwargs", {**manager.connection_kwargs, "db": 1}): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "startup_nodes": startup_nodes[:1]}): + assert resolver.resolve().kind == "python_callback" + binding: Final = resolver.resolve() + assert binding.kind == "native" + + client: Final = redis.RedisCluster(startup_nodes=[redis.cluster.ClusterNode(*node) for node in cluster_nodes]) + keys: Final = tuple(f"slot-{index}" for index in range(12)) + slots: Final = {client.keyslot(f"parity:{key}") for key in keys} + assert len(slots) > 1, slots + requests: Final = [request(key) for key in keys] + values: Final = [{"index": index} for index in range(len(keys))] + await binding.async_store_batch(requests, values) + client.set("parity:slot-3", "not a cache entry") + client.set("parity:slot-7", json.dumps({"timestamp": time.time(), "response": {"index": 7, "python": True}})) + + batch: Final = await binding.async_lookup_batch(requests) + assert batch == { + "values": [ + None if index == 3 else {"index": 7, "python": True} if index == 7 else value + for index, value in enumerate(values) + ], + "missing_indices": [3], + } + assert facade.cache.get_cache("parity:slot-0")["response"] == {"index": 0} + assert (await facade.cache.async_get_cache("parity:slot-11"))["response"] == {"index": 11} + assert facade.cache.redis_client.mget_nonatomic([f"parity:{key}" for key in keys[:2]]) == [ + client.get("parity:slot-0"), + client.get("parity:slot-1"), + ] + + await binding.async_store({**request("pinned"), "ttl_seconds": 12.0}, {"pinned": True}) + assert 0 < client.ttl("parity:pinned") <= 12 + client.set("unscoped", "stays") + + await binding.async_flush() + + remaining: Final = tuple( + sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node)) + ) + assert remaining == (), remaining + assert client.get("unscoped") == b"stays" + client.delete("unscoped") + client.close() + facade.cache.redis_client.close() + + +def redis_facade(redis_url: str, **settings: object) -> Cache: + parsed: Final = urlparse(redis_url) + return Cache(type=LiteLLMCacheType.REDIS, host=parsed.hostname, port=str(parsed.port), **settings) + + +@pytest.mark.parametrize( + ("settings", "message"), + [ + pytest.param({"max_connections": 10}, "max_connections requires Python", id="pool-size"), + pytest.param({"socket_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="socket-timeout"), + pytest.param( + {"socket_connect_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="connect-timeout" + ), + pytest.param({"socket_keepalive": True}, "does not support socket_keepalive", id="keepalive"), + pytest.param({"health_check_interval": 5}, "does not support health_check_interval", id="health-check"), + pytest.param({"client_name": "litellm"}, "does not support client_name", id="client-name"), + pytest.param({"ssl": True}, "ssl_check_hostname=false require Python", id="tls-default-hostname-check"), + pytest.param({"ssl": True, "ssl_cert_reqs": "none"}, "ssl_cert_reqs=none", id="tls-without-verification"), + pytest.param( + {"ssl": True, "ssl_check_hostname": True, "ssl_ca_certs": "/ca.pem"}, + "does not support ssl_ca_certs", + id="tls-custom-ca", + ), + pytest.param( + {"ssl": True, "ssl_check_hostname": True, "ssl_certfile": "/client.pem", "ssl_keyfile": "/client.key"}, + "does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or ssl_keyfile", + id="tls-client-certificate", + ), + ], +) +def test_redis_settings_the_native_client_cannot_honor_decline( + redis_url: str, monkeypatch: pytest.MonkeyPatch, settings: dict[str, object], message: str +) -> None: + require_rust(monkeypatch, LiteLLMCacheType.REDIS) + with pytest.raises(RuntimeError, match=f"declined the cache: native Redis.*{message}"): + redis_facade(redis_url, **settings) + + +def test_redis_verified_tls_activates_natively(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: + require_rust(monkeypatch, LiteLLMCacheType.REDIS) + assert_native_runtime(redis_facade(redis_url, ssl=True, ssl_check_hostname=True)) + + +async def test_redis_flush_size_buffers_native_facade_writes(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: + require_rust(monkeypatch, LiteLLMCacheType.REDIS) + facade: Final = redis_facade(redis_url, redis_flush_size=2, namespace="team") + assert_native_runtime(facade) + client: Final = redis.Redis.from_url(redis_url) + first: Final = completion_kwargs("first") + await facade.async_add_cache({"value": 1}, **first) + first_key: Final = facade.get_cache_key(**first) + assert first_key.startswith("team:") + assert client.get(first_key) is None + second: Final = completion_kwargs("second") + await facade.async_add_cache({"value": 2}, **second) + assert client.get(first_key) is not None + assert client.get(facade.get_cache_key(**second)) is not None + client.close() + + +def test_rust_with_fallback_keeps_python_when_the_native_client_declines( + redis_url: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + catalog, + "RULES", + (CacheRule(Rollout.RUST_OPT_OUT, backends=frozenset({LiteLLMCacheType.REDIS})),), + ) + assert redis_facade(redis_url, socket_timeout=1.0)._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor diff --git a/tests/test_litellm_rust/cache/test_redis_semantic.py b/tests/test_litellm_rust/cache/test_redis_semantic.py new file mode 100644 index 00000000000..279330d9060 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_redis_semantic.py @@ -0,0 +1,606 @@ +import asyncio +import contextvars +import hashlib +import json +import math +import os +from collections.abc import Callable, Generator +from contextlib import ExitStack +from types import SimpleNamespace +from typing import Final, cast +from uuid import uuid4 + +import pytest +import redis + +import litellm +from litellm.caching.caching import Cache +from litellm.caching.redis_semantic_cache import RedisSemanticCache +from litellm.types.caching import LiteLLMCacheType +from litellm.types.llms.custom_llm import CustomLLMItem +from litellm.types.utils import EmbeddingResponse +from tests.test_litellm_rust.support.cache import ( + CacheTestHandle, + CacheTestResolver, + assert_native_runtime, + request, + require_rust, +) +from tests.test_litellm_rust.support.isolation import rebound + +pytestmark: Final = pytest.mark.requires_rust_extension + + +PARAPHRASE_MARKER: Final = " (paraphrase)" + + +SEMANTIC_EMBEDDING_MODEL: Final = "semantic-test/deterministic" + + +SEMANTIC_INDEX_PREFIX: Final = "litellm_test_semantic_" + + +SEMANTIC_CONTEXT: Final = contextvars.ContextVar("semantic_test_context", default="unset") + + +def _normalized(vector: list[float]) -> list[float]: + norm: Final = math.sqrt(sum(component * component for component in vector)) + return [component / norm for component in vector] + + +def _base_embedding(prompt: str) -> list[float]: + digest: Final = hashlib.sha256(prompt.encode("utf-8")).digest() + return _normalized([float(digest[index] + 1) for index in range(8)]) + + +def _semantic_embedding(prompt: str) -> list[float]: + if PARAPHRASE_MARKER not in prompt: + return _base_embedding(prompt) + base: Final = _base_embedding(prompt.replace(PARAPHRASE_MARKER, "").strip()) + pivot: Final = min(range(8), key=lambda index: abs(base[index])) + direction: Final = _normalized( + [(1.0 - base[pivot] * base[pivot]) if index == pivot else -base[index] * base[pivot] for index in range(8)] + ) + # Rotating an orthogonal unit direction by 0.329 produces ~0.05 cosine distance + return _normalized([base[index] + 0.329 * direction[index] for index in range(8)]) + + +class DeterministicEmbedding(litellm.CustomLLM): + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + self.async_calls: list[dict[str, object]] = [] + self.entered = asyncio.Event() + self.gate: asyncio.Event | None = None + + def _respond( + self, + model: str, + input: object, + model_response: EmbeddingResponse, + ) -> EmbeddingResponse: + texts: Final = cast(list[object], input if isinstance(input, list) else [input]) + self.calls.append({"model": model, "input": texts}) + model_response.model = model + model_response.data = [ + {"object": "embedding", "index": index, "embedding": _semantic_embedding(str(text))} + for index, text in enumerate(texts) + ] + return model_response + + def embedding( + self, + model: str, + input: list[object], + model_response: EmbeddingResponse, + print_verbose: Callable[..., object], + logging_obj: object, + optional_params: dict[str, object], + api_key: object = None, + api_base: object = None, + timeout: object = None, + litellm_params: object = None, + ) -> EmbeddingResponse: + return self._respond(model, input, model_response) + + async def aembedding( + self, + model: str, + input: list[object], + model_response: EmbeddingResponse, + print_verbose: Callable[..., object], + logging_obj: object, + optional_params: dict[str, object], + api_key: object = None, + api_base: object = None, + timeout: object = None, + litellm_params: object = None, + ) -> EmbeddingResponse: + texts: Final = cast(list[object], input if isinstance(input, list) else [input]) + self.async_calls.append( + { + "model": model, + "input": texts, + "task": asyncio.current_task(), + "context": SEMANTIC_CONTEXT.get(), + } + ) + SEMANTIC_CONTEXT.set("written-in-aembedding") + self.entered.set() + if self.gate is not None: + await self.gate.wait() + return self._respond(model, input, model_response) + + +@pytest.fixture +def semantic_embedding() -> Generator[DeterministicEmbedding]: + handler: Final = DeterministicEmbedding() + with ExitStack() as stack: + stack.enter_context( + rebound( + litellm, + "custom_provider_map", + [ + *litellm.custom_provider_map, + cast( + CustomLLMItem, + {"provider": "semantic-test", "custom_handler": handler}, + ), + ], + ) + ) + stack.enter_context( + rebound( + litellm, + "_custom_providers", # pyright: ignore[reportPrivateUsage] # no public provider-registration hook + [*litellm._custom_providers, "semantic-test"], # pyright: ignore[reportPrivateUsage] # no public provider-registration hook + ) + ) + stack.enter_context(rebound(litellm, "provider_list", [*litellm.provider_list, "semantic-test"])) + yield handler + + +@pytest.fixture +def redis_stack() -> Generator[tuple[str, str]]: + url: Final = os.environ.get("LITELLM_REDIS_STACK_URL") + if url is None: + pytest.skip("LITELLM_REDIS_STACK_URL is not set") + index: Final = f"{SEMANTIC_INDEX_PREFIX}{uuid4().hex}" + yield url, index + client: Final = redis.Redis.from_url(url) + try: + client.execute_command("FT.DROPINDEX", index, "DD") # pyright: ignore[reportUnknownMemberType] # redis-py leaves execute_command partially unknown + except redis.RedisError: + pass + client.close() + + +def semantic_request(key: str, prompt: str, **extra: object) -> dict[str, object]: + return { + "key": {"preset": key}, + "messages": [{"role": "user", "content": prompt}], + **extra, + } + + +def semantic_messages(prompt: str) -> list[dict[str, object]]: + return [{"role": "user", "content": prompt}] + + +def semantic_entry_id(prompt: str, tag: str) -> str: + return hashlib.sha256(f"{prompt}litellm_cache_key{tag}".encode()).hexdigest() + + +def semantic_facade(url: str, index: str, *, similarity_threshold: float = 0.8) -> Cache: + facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=similarity_threshold, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + CacheTestHandle.redis_semantic(facade.cache)._bind_facade(facade) + return facade + + +def test_redis_semantic_constructor_identity_and_provenance( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + backend: Final = cast(RedisSemanticCache, facade.cache) + assert backend.__class__.__module__ == "litellm.caching.redis_semantic_cache" + assert type(backend) is RedisSemanticCache + assert backend._redis_url == url # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config + assert backend._index_name == index # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config + assert backend.similarity_threshold == 0.8 + assert backend.embedding_model == SEMANTIC_EMBEDDING_MODEL + handle: Final = cast(object, getattr(facade, "_native_cache_handle")) + assert isinstance(handle, CacheTestHandle) + assert handle.backend == "redis_semantic" + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + + +def test_redis_semantic_native_and_python_sync_entries_share_one_layout( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + response: Final = {"choices": [{"text": "paris"}], "usage": {"total_tokens": 2}} + + binding.store(semantic_request("geo", "what is the capital of france"), response) + + native_hash_key: Final = f"{index}:{semantic_entry_id('what is the capital of france', 'geo')}" + stored: Final = client.hgetall(native_hash_key) + assert set(stored) == { + b"entry_id", + b"prompt", + b"response", + b"prompt_vector", + b"inserted_at", + b"updated_at", + b"litellm_cache_key", + }, stored + assert stored[b"entry_id"].decode() == native_hash_key.split(":", 1)[1] + assert stored[b"prompt"] == b"what is the capital of france" + assert stored[b"litellm_cache_key"] == b"geo" + assert len(stored[b"prompt_vector"]) == 32 + decoded: Final = cast(dict[str, object], json.loads(stored[b"response"])) + assert decoded["response"] == response + assert ( + cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "geo", messages=semantic_messages("what is the capital of france") + ) + == decoded + ) + assert semantic_embedding.calls == [ + {"model": "deterministic", "input": ["what is the capital of france"]}, + {"model": "deterministic", "input": ["what is the capital of france"]}, + {"model": "deterministic", "input": ["dimension test"]}, + ] + + cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "math", + json.dumps({"timestamp": 1700000000.0, "response": {"answer": 42}}), + messages=semantic_messages("what is 6 times 7"), + ) + python_hash_key: Final = f"{index}:{semantic_entry_id('what is 6 times 7', 'math')}" + assert json.loads(cast(bytes, client.hget(python_hash_key, "response"))) == { + "timestamp": 1700000000.0, + "response": {"answer": 42}, + } + assert binding.lookup(semantic_request("math", "what is 6 times 7")) == {"answer": 42} + client.close() + + +async def test_redis_semantic_async_paths_and_store_batch_share_one_layout( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + await binding.async_store(semantic_request("async", "name a primary color"), {"answer": "blue"}) + hash_key: Final = f"{index}:{semantic_entry_id('name a primary color', 'async')}" + decoded: Final = cast(dict[str, object], json.loads(cast(bytes, client.hget(hash_key, "response")))) + python_read: Final = await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "async", messages=semantic_messages("name a primary color") + ) + assert python_read == decoded + + await binding.async_store_batch( + [ + semantic_request("batch-one", "first batch prompt"), + semantic_request("batch-two", "second batch prompt"), + ], + [{"answer": 1}, {"answer": 2}], + ) + expected: Final = { + key: json.loads(cast(bytes, client.hget(f"{index}:{semantic_entry_id(prompt, key)}", "response"))) + for key, prompt in ( + ("batch-one", "first batch prompt"), + ("batch-two", "second batch prompt"), + ) + } + for key, prompt in ( + ("batch-one", "first batch prompt"), + ("batch-two", "second batch prompt"), + ): + assert ( + cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + key, messages=semantic_messages(prompt) + ) + == expected[key] + ), key + + cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "async-python", + json.dumps({"timestamp": 1700000000.0, "response": {"answer": "python"}}), + messages=semantic_messages("python written prompt"), + ) + assert await binding.async_lookup(semantic_request("async-python", "python written prompt")) == {"answer": "python"} + client.close() + + +async def test_native_semantic_async_embedding_runs_inline_in_the_callers_task( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + caller: Final = asyncio.current_task() + SEMANTIC_CONTEXT.set("caller-sentinel") + response: Final = {"choices": [{"text": "paris"}]} + + await binding.async_store(semantic_request("inline", "what is the capital of france"), response) + assert ( + await binding.async_lookup(semantic_request("inline", f"what is the capital of france{PARAPHRASE_MARKER}")) + == response + ) + assert await binding.async_lookup(semantic_request("inline", "python written prompt")) is None + assert SEMANTIC_CONTEXT.get() == "written-in-aembedding" + assert semantic_embedding.async_calls == [ + { + "model": "deterministic", + "input": ["what is the capital of france"], + "task": caller, + "context": "caller-sentinel", + }, + { + "model": "deterministic", + "input": [f"what is the capital of france{PARAPHRASE_MARKER}"], + "task": caller, + "context": "written-in-aembedding", + }, + { + "model": "deterministic", + "input": ["python written prompt"], + "task": caller, + "context": "written-in-aembedding", + }, + ], semantic_embedding.async_calls + + +async def test_native_semantic_cancellation_during_embedding_skips_the_backend( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + semantic_embedding.gate = asyncio.Event() + + async def lookup() -> object: + return await binding.async_lookup(semantic_request("cancel", "cancelled prompt")) + + task: Final = asyncio.create_task(lookup()) + await semantic_embedding.entered.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + semantic_embedding.gate.set() + + assert len(semantic_embedding.async_calls) == 1 + assert ( + await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "cancel", messages=semantic_messages("cancelled prompt") + ) + is None + ) + + +def test_redis_semantic_similarity_tag_and_threshold_boundaries( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + + binding.store(semantic_request("sim", "tell me a joke"), {"answer": "haha"}) + paraphrase: Final = f"tell me a joke{PARAPHRASE_MARKER}" + assert binding.lookup(semantic_request("sim", paraphrase)) == {"answer": "haha"} + assert binding.lookup(semantic_request("sim", "an unrelated question about spreadsheets")) is None + assert binding.lookup(semantic_request("other-key", "tell me a joke")) is None + + strict: Final = semantic_facade(url, index, similarity_threshold=0.99) + strict_binding: Final = CacheTestResolver(SimpleNamespace(cache=strict)).resolve() + assert strict_binding.lookup(semantic_request("sim", paraphrase)) is None + assert strict_binding.lookup(semantic_request("sim", "tell me a joke")) == {"answer": "haha"} + + +def test_redis_semantic_ttl_is_written_only_when_requested( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store({**semantic_request("ttl", "ttl prompt"), "ttl_seconds": 12.0}, {"answer": 1}) + expiring: Final = f"{index}:{semantic_entry_id('ttl prompt', 'ttl')}" + assert 0 < client.ttl(expiring) <= 12 + + binding.store(semantic_request("ttl-none", "untimed prompt"), {"answer": 2}) + persistent: Final = f"{index}:{semantic_entry_id('untimed prompt', 'ttl-none')}" + assert client.ttl(persistent) == -1 + + binding.store( + {**semantic_request("ttl-fraction", "fractional prompt"), "ttl_seconds": 1.5}, + {"answer": 3}, + ) + fractional: Final = f"{index}:{semantic_entry_id('fractional prompt', 'ttl-fraction')}" + assert client.ttl(fractional) == 2 + client.close() + + +def test_redis_semantic_malformed_response_is_a_miss_for_both_readers( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store(semantic_request("bad", "corrupt me"), {"answer": 1}) + hash_key: Final = f"{index}:{semantic_entry_id('corrupt me', 'bad')}" + client.hset(hash_key, "response", b"{not json") + assert binding.lookup(semantic_request("bad", "corrupt me")) is None + assert ( + cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "bad", messages=semantic_messages("corrupt me") + ) + is None + ) + client.close() + + +async def test_redis_semantic_unsupported_operations_raise_not_implemented( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + + with pytest.raises(NotImplementedError): + binding.lookup_batch([semantic_request("batch", "prompt one")]) + with pytest.raises(NotImplementedError): + await binding.async_lookup_batch([semantic_request("batch", "prompt one")]) + with pytest.raises(NotImplementedError): + await binding.async_flush() + with pytest.raises(NotImplementedError): + await binding.ping() + + +def test_redis_semantic_requests_without_prompt_are_noops( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store(request("plain"), {"answer": 1}) + assert binding.lookup(request("plain")) is None + assert semantic_embedding.calls == [] + assert client.keys(f"{index}:*") == [] + client.close() + + +def test_redis_semantic_scope_overrides_the_tag_and_isolates_entries( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + scoped: Final = {**semantic_request("scoped", "scoped prompt"), "scope": "team-a"} + binding.store(scoped, {"answer": "kept"}) + hash_key: Final = f"{index}:{semantic_entry_id('scoped prompt', 'team-a')}" + assert client.hget(hash_key, "litellm_cache_key") == b"team-a" + assert binding.lookup(scoped) == {"answer": "kept"} + assert binding.lookup(semantic_request("scoped", "scoped prompt")) is None + assert binding.lookup({**scoped, "scope": "team-b"}) is None + client.close() + + +def test_redis_semantic_configuration_drift_falls_back_to_python( + redis_stack: tuple[str, str], + semantic_embedding: DeterministicEmbedding, + monkeypatch: pytest.MonkeyPatch, +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "native" + + with rebound(facade.cache, "similarity_threshold", 0.5): + assert resolver.resolve().kind == "python_callback" + with rebound(facade, "semantic_cache_scope", "end_user"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "embedding_model", "other-model"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "_index_name", "other-index"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "CACHE_KEY_FIELD_NAME", "other-field"): + assert resolver.resolve().kind == "python_callback" + + def patched_embedding(self: object, prompt: str, metadata: object = None) -> list[float]: + return _semantic_embedding(prompt) + + monkeypatch.setattr(RedisSemanticCache, "_get_embedding", patched_embedding) + assert resolver.resolve().kind == "python_callback" + + +def test_redis_semantic_handle_rejects_wrong_backends( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + + class CustomSemanticCache(RedisSemanticCache): + pass + + with pytest.raises(TypeError, match="built-in RedisSemanticCache"): + CacheTestHandle.redis_semantic(object()) + with pytest.raises(TypeError, match="built-in RedisSemanticCache"): + CacheTestHandle.redis_semantic( + CustomSemanticCache( + redis_url=url, + similarity_threshold=0.8, + embedding_model=SEMANTIC_EMBEDDING_MODEL, + index_name=f"{index}_subclass", + ) + ) + + facade: Final = semantic_facade(url, index) + with pytest.raises(TypeError, match="backend types must match"): + CacheTestHandle.redis(url)._bind_facade(facade) + + subclassed_facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + subclassed_facade.cache = CustomSemanticCache( # pyright: ignore[reportAttributeAccessIssue] # facade backend slot is not declared + redis_url=url, + similarity_threshold=0.8, + embedding_model=SEMANTIC_EMBEDDING_MODEL, + index_name=index, + ) + with pytest.raises(TypeError): + CacheTestHandle.redis_semantic(subclassed_facade.cache)._bind_facade(subclassed_facade) + + replacement_facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + with pytest.raises(TypeError, match="must be the native embedder"): + CacheTestHandle.redis_semantic(facade.cache)._bind_facade(replacement_facade) + + +async def test_redis_semantic_rust_required_rule_activates_natively( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding, monkeypatch: pytest.MonkeyPatch +) -> None: + del semantic_embedding + url, index = redis_stack + require_rust(monkeypatch, LiteLLMCacheType.REDIS_SEMANTIC) + facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + assert_native_runtime(facade) + kwargs: Final = {"model": "gpt-4o", "messages": semantic_messages("name a primary color")} + await facade.async_add_cache({"answer": "blue"}, **kwargs) + assert await facade.async_get_cache(**kwargs) == {"answer": "blue"} diff --git a/tests/test_litellm_rust/cache/test_rollout.py b/tests/test_litellm_rust/cache/test_rollout.py new file mode 100644 index 00000000000..7f33e31599f --- /dev/null +++ b/tests/test_litellm_rust/cache/test_rollout.py @@ -0,0 +1,264 @@ +import asyncio +from collections.abc import Callable +from pathlib import Path +from types import SimpleNamespace +from typing import Final, TypeAlias, cast +from urllib.parse import urlparse +from uuid import uuid4 + +import pytest + +from litellm.caching.caching import Cache +from litellm.rust_bridge.response_cache import NativeResponseCacheRuntime, ResponseCacheRuntime, resolve_response_cache +from litellm.types.caching import LiteLLMCacheType +from litellm.types.utils import EmbeddingResponse +from tests.test_litellm_rust.support.cache import assert_native_runtime, completion_kwargs, require_rust +from tests.test_litellm_rust.support.s3_stub import S3Stub + +pytestmark: Final = pytest.mark.requires_rust_extension + + +CacheFactory: TypeAlias = Callable[[], Cache] + + +@pytest.fixture +def cache_factory(request: pytest.FixtureRequest, tmp_path: Path) -> CacheFactory: + backend: Final = cast(LiteLLMCacheType, request.param) + match backend: + case LiteLLMCacheType.LOCAL: + return lambda: Cache(type=backend) + case LiteLLMCacheType.DISK: + return lambda: Cache(type=backend, disk_cache_dir=str(tmp_path)) + case LiteLLMCacheType.REDIS: + parsed: Final = urlparse(cast(str, request.getfixturevalue("redis_url"))) + return lambda: Cache(type=backend, host=parsed.hostname, port=str(parsed.port)) + case LiteLLMCacheType.S3: + stub: Final = cast(S3Stub, request.getfixturevalue("s3_stub")) + return lambda: Cache( + type=backend, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + case LiteLLMCacheType.GCS: + return lambda: Cache(type=backend, gcs_bucket_name="bucket", gcs_path="cache/") + case LiteLLMCacheType.REDIS_SEMANTIC: + return lambda: Cache( + type=backend, + redis_url="redis://127.0.0.1:6379", + similarity_threshold=0.8, + redis_semantic_cache_embedding_model="text-embedding-3-small", + ) + case LiteLLMCacheType.VALKEY_SEMANTIC: + return lambda: Cache(type=backend, redis_url="redis://127.0.0.1:6390/0", similarity_threshold=0.8) + case _: + raise AssertionError(f"no local factory for {backend}") + + +ROUND_TRIP_BACKENDS: Final = ( + LiteLLMCacheType.LOCAL, + LiteLLMCacheType.DISK, + LiteLLMCacheType.REDIS, + LiteLLMCacheType.S3, +) + + +SHARED_STORE_BACKENDS: Final = (LiteLLMCacheType.DISK, LiteLLMCacheType.REDIS, LiteLLMCacheType.S3) + + +@pytest.mark.parametrize("backend", list(LiteLLMCacheType)) +def test_shipped_rules_keep_every_backend_on_python(backend: LiteLLMCacheType) -> None: + assert resolve_response_cache(cast(Cache, SimpleNamespace(type=backend))) is None + + +@pytest.mark.parametrize( + "cache_factory", + [ + LiteLLMCacheType.LOCAL, + LiteLLMCacheType.DISK, + LiteLLMCacheType.REDIS, + LiteLLMCacheType.S3, + LiteLLMCacheType.GCS, + LiteLLMCacheType.REDIS_SEMANTIC, + LiteLLMCacheType.VALKEY_SEMANTIC, + ], + indirect=True, +) +def test_shipped_rules_construct_python_backed_facades(cache_factory: CacheFactory) -> None: + assert cache_factory()._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor + + +@pytest.mark.parametrize( + "cache_factory", + [ + LiteLLMCacheType.LOCAL, + LiteLLMCacheType.DISK, + LiteLLMCacheType.REDIS, + LiteLLMCacheType.S3, + LiteLLMCacheType.GCS, + LiteLLMCacheType.REDIS_SEMANTIC, + LiteLLMCacheType.VALKEY_SEMANTIC, + ], + indirect=True, +) +def test_rust_required_rule_activates_the_native_backend( + cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest +) -> None: + require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) + assert_native_runtime(cache_factory()) + + +@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True) +async def test_facade_storage_calls_round_trip_through_the_native_backend( + cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest +) -> None: + require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) + facade: Final = cache_factory() + assert_native_runtime(facade) + + sync_kwargs: Final = completion_kwargs("sync") + facade.add_cache({"answer": 1}, **sync_kwargs) + assert facade.get_cache(**sync_kwargs) == {"answer": 1} + + async_kwargs: Final = completion_kwargs("async") + await facade.async_add_cache({"answer": 2}, **async_kwargs) + assert await facade.async_get_cache(**async_kwargs) == {"answer": 2} + assert facade.get_cache(**completion_kwargs("absent")) is None + + +async def test_memory_facade_writes_bypass_the_python_backend(monkeypatch: pytest.MonkeyPatch) -> None: + require_rust(monkeypatch, LiteLLMCacheType.LOCAL) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + assert_native_runtime(facade) + kwargs: Final = completion_kwargs("memory") + facade.add_cache({"answer": 1}, **kwargs) + assert facade.cache.get_cache(facade.get_cache_key(**kwargs)) is None + assert facade.get_cache(**kwargs) == {"answer": 1} + + +@pytest.mark.parametrize("cache_factory", SHARED_STORE_BACKENDS, indirect=True) +async def test_native_and_python_facades_share_one_wire_format( + cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest +) -> None: + python_facade: Final = cache_factory() + assert python_facade._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor + require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) + native_facade: Final = cache_factory() + assert_native_runtime(native_facade) + + native_written: Final = completion_kwargs("native") + native_facade.add_cache({"writer": "native"}, **native_written) + assert python_facade.get_cache(**native_written) == {"writer": "native"} + + python_written: Final = completion_kwargs("python") + python_facade.add_cache({"writer": "python"}, **python_written) + assert native_facade.get_cache(**python_written) == {"writer": "python"} + + async_native: Final = completion_kwargs("async-native") + await native_facade.async_add_cache({"writer": "async-native"}, **async_native) + assert await python_facade.async_get_cache(**async_native) == {"writer": "async-native"} + + async_python: Final = completion_kwargs("async-python") + await python_facade.async_add_cache({"writer": "async-python"}, **async_python) + assert await native_facade.async_get_cache(**async_python) == {"writer": "async-python"} + + +@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True) +async def test_embedding_pipeline_stores_one_native_entry_per_input( + cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest +) -> None: + require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) + facade: Final = cache_factory() + assert_native_runtime(facade) + inputs: Final = [f"alpha {uuid4().hex}", f"beta {uuid4().hex}"] + result: Final = EmbeddingResponse( + model="text-embedding-3-small", + data=[ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}, + {"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}, + ], + ) + await facade.async_add_cache_pipeline(result, model="text-embedding-3-small", input=inputs) + + keys: Final = [facade.get_cache_key(model="text-embedding-3-small", input=text) for text in inputs] + assert len(set(keys)) == len(inputs) + for text, expected in zip(inputs, ([0.1, 0.2], [0.3, 0.4]), strict=True): + cached = await facade.async_get_cache(model="text-embedding-3-small", input=text) + assert isinstance(cached, dict) + assert cached["embedding"] == expected + assert await facade.async_get_cache(model="text-embedding-3-small", input=inputs) is None + + +@pytest.mark.parametrize( + ("backend", "settings", "message"), + [ + pytest.param( + LiteLLMCacheType.VALKEY_SEMANTIC, + {"redis_url": "rediss://127.0.0.1:6390/0", "similarity_threshold": 0.8}, + "native Valkey semantic cache does not support TLS connections", + id="valkey-tls", + ), + pytest.param( + LiteLLMCacheType.VALKEY_SEMANTIC, + {"redis_url": "redis://127.0.0.1:6390/0?socket_timeout=1", "similarity_threshold": 0.8}, + "native Redis uses fixed socket timeouts; socket_timeout and socket_connect_timeout require Python", + id="valkey-socket-timeout", + ), + pytest.param( + LiteLLMCacheType.REDIS_SEMANTIC, + {"redis_url": "rediss://127.0.0.1:6380", "similarity_threshold": 0.8}, + "native Redis semantic cache does not support TLS or query options in redis_url", + id="redis-semantic-tls", + ), + pytest.param( + LiteLLMCacheType.REDIS_SEMANTIC, + {"redis_url": "redis://127.0.0.1:6379?socket_timeout=1", "similarity_threshold": 0.8}, + "native Redis semantic cache does not support TLS or query options in redis_url", + id="redis-semantic-query", + ), + ], +) +def test_semantic_settings_the_native_client_cannot_honor_decline( + monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType, settings: dict[str, object], message: str +) -> None: + require_rust(monkeypatch, backend) + with pytest.raises(RuntimeError, match=f"declined the cache: {message}"): + Cache(type=backend, **settings) + + +class _SemanticHit: + """A native semantic runtime that answers every lookup with one cached response.""" + + kind: Final = "native" + + def lookup_semantic(self, request: object) -> tuple[object, float | None]: + return {"answer": 42}, 0.97 + + async def async_lookup_semantic(self, request: object) -> tuple[object, float | None]: + return {"answer": 42}, 0.97 + + +@pytest.mark.parametrize("semantic_type", [LiteLLMCacheType.QDRANT_SEMANTIC, LiteLLMCacheType.REDIS_SEMANTIC]) +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +def test_native_semantic_hit_stamps_similarity_on_request_metadata( + semantic_type: LiteLLMCacheType, use_async: bool +) -> None: + """Python semantic backends write `metadata["semantic-similarity"]` on every lookup, and the + facade copies it to the caller's metadata; the native path must report it the same way.""" + facade: Final = Cache() + facade.type = semantic_type + facade._native_cache = ResponseCacheRuntime(cast(NativeResponseCacheRuntime, _SemanticHit())) # pyright: ignore[reportPrivateUsage] # the native path under test has no public setter + metadata: Final[dict[str, object]] = {} + kwargs: Final = { + "cache_key": "semantic-key", + "messages": [{"role": "user", "content": "hello"}], + "metadata": metadata, + } + + result: Final = asyncio.run(facade.async_get_cache(**kwargs)) if use_async else facade.get_cache(**kwargs) + + assert result == {"answer": 42} + assert metadata["semantic-similarity"] == 0.97 diff --git a/tests/test_litellm_rust/cache/test_s3.py b/tests/test_litellm_rust/cache/test_s3.py new file mode 100644 index 00000000000..044bfc39f8d --- /dev/null +++ b/tests/test_litellm_rust/cache/test_s3.py @@ -0,0 +1,187 @@ +import json +import time +from datetime import datetime +from types import SimpleNamespace +from typing import Final, cast +from unittest.mock import Mock + +import boto3 +import botocore.config +import pytest + +from litellm.caching.caching import Cache +from litellm.caching.s3_cache import S3Cache +from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.isolation import rebound +from tests.test_litellm_rust.support.s3_stub import S3Stub + +pytestmark: Final = pytest.mark.requires_rust_extension + + +def python_s3(url: str) -> S3Cache: + return S3Cache( + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + + +async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: S3Stub) -> None: + python_cache: Final = python_s3(s3_stub.url) + response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} + python_cache.set_cache("sync:key", {"timestamp": time.time(), "response": response}, ttl=90) + python_cache.set_cache("plain", {"timestamp": time.time(), "response": response}) + s3_stub.put_object("team/malformed", b"not a cache entry") + s3_stub.put_object( + "team/expired", + json.dumps({"timestamp": time.time(), "response": response}).encode(), + {"expires": "Thu, 01 Jan 1970 00:00:00 GMT"}, + ) + binding: Final = CacheTestResolver( + SimpleNamespace( + cache=CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + ) + ).resolve() + + assert binding.lookup(request("sync:key")) == response + assert await binding.async_lookup(request("plain")) == response + assert binding.lookup(request("malformed")) is None + assert binding.lookup(request("expired")) is None + assert binding.lookup(request("absent")) is None + + binding.store({**request("native:key"), "ttl_seconds": 90.0}, response) + await binding.async_store(request("no_ttl"), response) + stored: Final = s3_stub.objects["team/native/key"] + assert stored.headers["content-type"] == "application/json" + assert stored.headers["content-language"] == "en" + assert stored.headers["content-disposition"] == 'inline; filename="team/native/key.json"' + assert stored.headers["cache-control"] == "immutable, max-age=90, s-maxage=90" + expires: Final = cast(datetime, s3_stub.expires("team/native/key")) + remaining: Final = (expires - datetime.now(expires.tzinfo)).total_seconds() + assert 60 < remaining <= 91 + no_ttl: Final = s3_stub.objects["team/no_ttl"] + assert no_ttl.headers["cache-control"] == "immutable, max-age=31536000, s-maxage=31536000" + assert "expires" not in no_ttl.headers + assert python_cache.get_cache("native:key")["response"] == response + + partial: Final = await binding.async_lookup_batch([request("native:key"), request("absent"), request("malformed")]) + assert partial == {"values": [response, None, None], "missing_indices": [1, 2]} + + +def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_stub: S3Stub) -> None: + facade: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + handle: Final = CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + with pytest.raises(TypeError, match="buckets must match"): + CacheTestHandle.s3("other", region="us-east-1", endpoint_url=s3_stub.url)._bind_facade(facade) + with pytest.raises(TypeError, match="key prefixes must match"): + CacheTestHandle.s3( + "cache-bucket", region="us-east-1", endpoint_url=s3_stub.url, key_prefix="other/" + )._bind_facade(facade) + handle._bind_facade(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + binding: Final = resolver.resolve() + assert binding.kind == "native" + + handler: Final = Mock() + facade.cache.s3_client.meta.events.register("before-call.s3.*", handler) + binding.store(request("native"), {"answer": 1}) + assert binding.lookup(request("native")) == {"answer": 1} + assert handler.call_count == 0 + assert "team/native" in s3_stub.objects + + with rebound(facade.cache, "bucket_name", "other"): + assert resolver.resolve().kind == "python_callback" + other_client: Final = boto3.client( + "s3", + region_name="us-east-1", + endpoint_url=s3_stub.url, + aws_access_key_id="key", + aws_secret_access_key="secret", + ) + with rebound(facade.cache, "s3_client", other_client): + assert resolver.resolve().kind == "python_callback" + + class CustomS3Cache(S3Cache): + pass + + subclassed: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + subclassed.cache = CustomS3Cache( + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + with pytest.raises(TypeError): + handle._bind_facade(subclassed) + assert CacheTestResolver(SimpleNamespace(cache=subclassed)).resolve().kind == "python_callback" + + +def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) -> None: + handle: Final = CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + unverified: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url="https://s3.example.test", + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + s3_verify=False, + ) + with pytest.raises(TypeError, match="requires Python"): + handle._bind_facade(unverified) + proxied: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + s3_config=botocore.config.Config(proxies={"https": "http://proxy.test"}), + ) + with pytest.raises(TypeError, match="requires Python"): + handle._bind_facade(proxied) diff --git a/tests/test_litellm_rust/test_valkey_semantic_cache_native.py b/tests/test_litellm_rust/cache/test_valkey_semantic.py similarity index 100% rename from tests/test_litellm_rust/test_valkey_semantic_cache_native.py rename to tests/test_litellm_rust/cache/test_valkey_semantic.py diff --git a/tests/test_litellm_rust/messages/test_callbacks.py b/tests/test_litellm_rust/messages/test_callbacks.py index 19043780eb6..dc66852d214 100644 --- a/tests/test_litellm_rust/messages/test_callbacks.py +++ b/tests/test_litellm_rust/messages/test_callbacks.py @@ -127,7 +127,7 @@ async def test_native_messages_stream_relays_provider_events_and_logs_success_on **arguments(messages_server, stream=True, callbacks=[recorder]) ) assert isinstance(stream, AsyncIterator) - assert get_hidden_params_dict(stream) == {"additional_headers": {"x-litellm-rust": "true"}} + assert get_hidden_params_dict(stream)["additional_headers"]["x-litellm-rust"] == "true" first: Final = await anext(stream) await drain_logging() assert "async_log_success_event" not in recorder.names @@ -171,7 +171,7 @@ def test_native_sync_messages_stream_relays_provider_events_and_logs_success_onc stream: Final = litellm.anthropic.messages.create(**arguments(messages_server, stream=True, callbacks=[recorder])) assert isinstance(stream, Iterator) - assert get_hidden_params_dict(stream) == {"additional_headers": {"x-litellm-rust": "true"}} + assert get_hidden_params_dict(stream)["additional_headers"]["x-litellm-rust"] == "true" assert b"".join(stream) == sse_payload() assert_served_natively(messages_server) @@ -186,3 +186,56 @@ def test_native_sync_messages_returns_the_provider_message(messages_server: Reco assert_served_natively(messages_server) assert response["content"] == MESSAGES_RESPONSE["content"] assert len(recorder.wait_for("log_success_event")) == 1 + + +@pytest.mark.asyncio +async def test_native_messages_pre_call_sees_the_shaped_optional_params( + messages_server: RecordingServer, +) -> None: + recorder: Final = RecordingLogger() + + await litellm.anthropic.messages.acreate( + **arguments(messages_server, callbacks=[recorder], temperature=0.2, top_k=3, drop_params=True) + ) + + sent: Final = messages_server.requests[0].body + assert not {"temperature", "top_k"} & sent.keys() + pre_call: Final = recorder.wait_for("log_pre_api_call")[0].kwargs + assert isinstance(pre_call, dict) + optional_params: Final = pre_call["optional_params"] + assert isinstance(optional_params, dict) + assert not {"model", "messages", "temperature", "top_k"} & optional_params.keys() + assert optional_params["max_tokens"] == sent["max_tokens"] + + +@pytest.mark.asyncio +async def test_native_messages_failing_pre_call_logger_does_not_fail_the_call(messages_server: RecordingServer) -> None: + class Broken(CustomLogger): + def log_pre_api_call(self, model, messages, kwargs): + raise RuntimeError("logger exploded") + + response: Final = await litellm.anthropic.messages.acreate(**arguments(messages_server, callbacks=[Broken()])) + + assert_served_natively(messages_server) + assert response["content"] == MESSAGES_RESPONSE["content"] + + +@pytest.mark.asyncio +async def test_native_messages_stream_success_log_carries_usage_rebuilt_from_the_relayed_events( + messages_server: RecordingServer, +) -> None: + messages_server.enqueue(STREAM) + recorder: Final = RecordingLogger() + + stream: Final = await litellm.anthropic.messages.acreate( + **arguments(messages_server, stream=True, callbacks=[recorder]) + ) + assert isinstance(stream, AsyncIterator) + async for _ in stream: + pass + + success: Final = await recorder.wait_for_async("async_log_success_event") + usage: Final = success[0].response.usage + assert usage.completion_tokens == MESSAGES_EVENTS[4][1]["usage"]["output_tokens"] + assert usage.prompt_tokens == MESSAGES_RESPONSE["usage"]["input_tokens"] + assert success[0].response.choices[0].message.content == "Hello from native Messages" diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 0d3b8ba472d..d09e60784fa 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -244,6 +244,21 @@ def test_native_ocr_maps_provider_400_with_public_provider_details(ocr_server: R assert "invalid OCR request" in str(caught.value) +def test_native_ocr_encodes_python_file_input_and_drops_unknown_arguments(ocr_server: RecordingServer) -> None: + response: Final = call_native_ocr( + ocr_server, + document={"type": "file", "file": BytesIO(b"abc"), "mime_type": "image/png"}, + opaque_extension=object(), + ) + + assert response.pages[0].markdown == "native OCR response" + assert_native_request(ocr_server) + assert ocr_server.requests[0].body == { + "model": "mistral-ocr-latest", + "document": {"type": "image_url", "image_url": "data:image/png;base64,YWJj"}, + } + + class TokenAbort(BaseException): pass diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py new file mode 100644 index 00000000000..41eb4d25257 --- /dev/null +++ b/tests/test_litellm_rust/support/cache.py @@ -0,0 +1,40 @@ +from typing import Final, Protocol +from uuid import uuid4 + +import pytest + +from litellm.caching.caching import Cache +from litellm.rust_bridge import _native, catalog +from litellm.rust_bridge.catalog import CacheRule +from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.response_cache import ResponseCacheRuntime +from litellm.types.caching import LiteLLMCacheType + +CacheTestHandle: Final = _native._CacheTestHandle # pyright: ignore[reportPrivateUsage] # test-only handle has no public module name + + +CacheTestResolver: Final = _native._CacheTestResolver # pyright: ignore[reportPrivateUsage] # test-only resolver has no public module name + + +class CacheLookup(Protocol): + def get_cache(self, **kwargs: object) -> object: ... + def flush_cache(self) -> object: ... + + +def request(key: str = "key") -> dict[str, object]: + return {"key": {"preset": key}} + + +def require_rust(monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType) -> None: + monkeypatch.setattr(catalog, "RULES", (CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({backend})),)) + + +def assert_native_runtime(facade: Cache) -> ResponseCacheRuntime: + runtime: Final = facade._native_cache # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor + assert isinstance(runtime, ResponseCacheRuntime) + assert runtime.kind == "native" + return runtime + + +def completion_kwargs(label: str) -> dict[str, object]: + return {"model": "gpt-4o", "messages": [{"role": "user", "content": f"{label} {uuid4().hex}"}]} diff --git a/tests/test_litellm_rust/test_cache.py b/tests/test_litellm_rust/test_cache.py deleted file mode 100644 index 96b3674fde3..00000000000 --- a/tests/test_litellm_rust/test_cache.py +++ /dev/null @@ -1,2397 +0,0 @@ -import asyncio -import contextvars -import gc -import hashlib -import http.server -import json -import math -import os -import threading -import time -import uuid -import weakref -from collections.abc import Callable, Generator -from contextlib import ExitStack -from datetime import datetime -from pathlib import Path -from types import SimpleNamespace -from typing import Final, Protocol, TypeAlias, cast -from unittest.mock import Mock -from urllib.parse import urlparse -from uuid import uuid4 - -import boto3 -import botocore.config -import diskcache -import fakeredis -import pytest -import redis -from azure.storage.blob import ContainerClient - -import litellm -from litellm.caching.azure_blob_cache import AzureBlobCache -from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache -from litellm.caching.disk_cache import DiskCache -from litellm.caching.gcs_cache import GCSCache -from litellm.caching.in_memory_cache import InMemoryCache -from litellm.caching.redis_cluster_cache import RedisClusterCache -from litellm.caching.redis_semantic_cache import RedisSemanticCache -from litellm.caching.s3_cache import S3Cache -from litellm.rust_bridge import _native, catalog -from litellm.rust_bridge.catalog import CacheRule, Route, RouteRule, SecretManagerRule -from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.response_cache import NativeResponseCacheRuntime, ResponseCacheRuntime, resolve_response_cache -from litellm.types.caching import LiteLLMCacheType -from litellm.types.llms.custom_llm import CustomLLMItem -from litellm.types.utils import EmbeddingResponse -from tests.test_litellm_rust.support.fake_gcs import FakeGcs -from tests.test_litellm_rust.support.isolation import rebound -from tests.test_litellm_rust.support.s3_stub import S3Stub - -_CacheTestHandle: Final = _native._CacheTestHandle # pyright: ignore[reportPrivateUsage] # test-only handle has no public module name -_CacheTestResolver: Final = _native._CacheTestResolver # pyright: ignore[reportPrivateUsage] # test-only resolver has no public module name - -pytestmark: Final = pytest.mark.requires_rust_extension - - -class CacheLookup(Protocol): - def get_cache(self, **kwargs: object) -> object: ... - def flush_cache(self) -> object: ... - - -def request(key: str = "key") -> dict[str, object]: - return {"key": {"preset": key}} - - -def qdrant_request( - key: str, - messages: list[dict[str, object]], - **kwargs: object, -) -> dict[str, object]: - return {**request(key), "messages": messages, **kwargs} - - -def embedding_vector(text: str) -> list[float]: - raw: Final = hashlib.sha256(text.encode()).digest()[:8] - values: Final = [byte / 127.5 - 1 for byte in raw] - norm: Final = math.sqrt(sum(value * value for value in values)) - return [value / norm for value in values] - - -@pytest.fixture -def qdrant_url() -> str: - value: Final[str | None] = os.environ.get("QDRANT_URL") - if not value: - pytest.skip("QDRANT_URL is required for Qdrant semantic cache tests") - return value.rstrip("/") - - -@pytest.fixture -def fake_embedding_endpoint(monkeypatch: pytest.MonkeyPatch) -> Generator[str]: - class EmbeddingHandler(http.server.BaseHTTPRequestHandler): - def do_POST(self) -> None: - length: Final = int(self.headers["Content-Length"]) - body: Final = json.loads(self.rfile.read(length)) - text: Final = body["input"] - response: Final = { - "object": "list", - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": embedding_vector(text), - } - ], - "model": body["model"], - "usage": {"prompt_tokens": 1, "total_tokens": 1}, - } - encoded: Final = json.dumps(response).encode() - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(encoded))) - self.end_headers() - self.wfile.write(encoded) - - def log_message(self, *_args: object) -> None: - return - - server: Final = http.server.ThreadingHTTPServer(("127.0.0.1", 0), EmbeddingHandler) - worker: Final = threading.Thread(target=server.serve_forever, daemon=True) - worker.start() - monkeypatch.setenv("OPENAI_API_BASE", f"http://127.0.0.1:{server.server_address[1]}") - monkeypatch.setenv("OPENAI_API_KEY", "sk-test") - try: - yield f"http://127.0.0.1:{server.server_address[1]}" - finally: - server.shutdown() - server.server_close() - worker.join(timeout=5) - - -@pytest.fixture -def redis_url() -> Generator[str]: - server: Final = fakeredis.TcpFakeServer(("127.0.0.1", 0), server_type="redis") - worker: Final = threading.Thread(target=server.serve_forever, daemon=True) - worker.start() - try: - yield f"redis://127.0.0.1:{server.server_address[1]}" - finally: - server.shutdown() - server.server_close() - worker.join(timeout=5) - - -@pytest.fixture -def fake_gcs() -> Generator[FakeGcs]: - server: Final = FakeGcs() - try: - yield server - finally: - server.close() - - -@pytest.fixture -def azure_blob_facade() -> Generator[Cache]: - account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") - if account_url is None: - pytest.skip( - "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" - ) - facade: Final = Cache( - type=LiteLLMCacheType.AZURE_BLOB, - azure_account_url=account_url, - azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", - ) - backend: Final = facade.cache - assert isinstance(backend, AzureBlobCache) - try: - yield facade - finally: - backend.container_client.delete_container() - asyncio.run(backend.disconnect()) - - -def azure_blob_handle(facade: Cache) -> _native._CacheTestHandle: - backend: Final = facade.cache - assert isinstance(backend, AzureBlobCache) - return _native._CacheTestHandle.azure_blob( - backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}"), - backend.container_client.container_name, - ) - - -@pytest.fixture -def cluster_nodes() -> tuple[tuple[str, int], ...]: - configured: Final = os.environ.get("LITELLM_TEST_REDIS_CLUSTER_NODES") - if not configured: - pytest.skip("LITELLM_TEST_REDIS_CLUSTER_NODES is not set") - return tuple((host, int(port)) for host, _, port in (node.partition(":") for node in configured.split(","))) - - -def test_existing_constructor_and_global_are_unchanged() -> None: - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - assert type(facade.cache) is InMemoryCache - assert "_native_cache_handle" not in vars(facade) - assert resolve_response_cache(facade) is None - with rebound(litellm, "cache", facade): - resolver: Final = _CacheTestResolver(litellm) - assert resolver.resolve().kind == "python_callback" - resolver.resolve().store(None, {"answer": 7}, callback_kwargs={"cache_key": "key"}) - assert cast(CacheLookup, facade).get_cache(cache_key="key") == {"answer": 7} - - -async def test_catalog_constructs_native_runtime_from_public_cache_configuration() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) - assert isinstance(runtime, ResponseCacheRuntime) - assert runtime.kind == "native" - - sync_request: Final = runtime.request(facade, {"cache_key": "sync"}) - assert sync_request is not None - runtime.store(sync_request, {"answer": 1}) - assert runtime.lookup(sync_request) == {"answer": 1} - assert facade.cache.get_cache("sync") is None - - async_request: Final = runtime.request(facade, {"cache_key": "async"}) - assert async_request is not None - await runtime.async_store(async_request, {"answer": 2}) - assert await runtime.async_lookup(async_request) == {"answer": 2} - assert await facade.cache.async_get_cache("async") is None - - requests: Final = (sync_request, async_request) - expected: Final = { - "values": [{"answer": 1}, {"answer": 2}], - "missing_indices": [], - } - assert runtime.lookup_batch(requests) == expected - assert await runtime.async_lookup_batch(requests) == expected - - await runtime.async_flush() - assert runtime.lookup(sync_request) is None - assert await runtime.async_lookup(async_request) is None - - -async def test_inference_resolver_uses_the_configured_native_cache_directly() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) - assert isinstance(runtime, ResponseCacheRuntime) - facade._native_cache = runtime - - selected: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() - assert selected.kind == "native" - request: Final = runtime.request(facade, {"cache_key": "inference-native"}) - assert request is not None - await selected.async_store(request, {"answer": 42}) - assert await selected.async_lookup(request) == {"answer": 42} - assert await runtime.async_lookup(request) == {"answer": 42} - assert facade.cache.get_cache("inference-native") is None - - facade._native_cache = None - fallback: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() - assert fallback.kind == "python_callback" - await fallback.async_store(None, {"answer": 7}, callback_kwargs={"cache_key": "inference-python"}) - assert facade.get_cache(cache_key="inference-python") == {"answer": 7} - assert facade.cache.get_cache("inference-python") is not None - - -async def test_inference_resolver_declines_a_native_runtime_whose_facade_changed() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) - assert isinstance(runtime, ResponseCacheRuntime) - facade._native_cache = runtime - stale_request: Final = runtime.request(facade, {"cache_key": "stale-only"}) - assert stale_request is not None - await runtime.async_store(stale_request, {"answer": "stale"}) - - replacement: Final = InMemoryCache() - facade.cache = replacement - with pytest.raises(_native.RustBridgeDeclined): - _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() - assert await runtime.async_lookup(stale_request) == {"answer": "stale"} - assert replacement.get_cache("stale-only") is None - assert replacement.get_cache("swapped-backend") is None - - -def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> None: - resolver: Final = _CacheTestResolver(litellm) - - enable_cache(type=LiteLLMCacheType.LOCAL, ttl=30) - enabled: Final = litellm.cache - assert isinstance(enabled, Cache) - assert enabled.ttl == 30 - assert resolver.resolve().kind == "python_callback" - - enable_cache(type=LiteLLMCacheType.LOCAL, ttl=60) - assert litellm.cache is enabled - - update_cache(type=LiteLLMCacheType.LOCAL, ttl=60) - updated: Final = litellm.cache - assert isinstance(updated, Cache) - assert updated is not enabled - assert updated.ttl == 60 - - disable_cache() - assert litellm.cache is None - assert resolver.resolve().kind == "disabled" - - -async def test_native_bindings_survive_replacement_and_capture_writes_before_dispatch() -> None: - namespace: Final = SimpleNamespace(cache=_CacheTestHandle.memory()) - resolver: Final = _CacheTestResolver(namespace) - selected: Final = resolver.resolve() - assert selected.kind == "native" - selected.store(request(), {"answer": 1}) - assert await selected.async_lookup(request()) == {"answer": 1} - with rebound(namespace, "cache", _CacheTestHandle.memory()): - replacement: Final = resolver.resolve() - await selected.async_store(request(), {"answer": 2}) - assert replacement.lookup(request()) is None - assert selected.lookup(request()) == {"answer": 2} - with rebound(namespace, "cache", None): - disabled: Final = resolver.resolve() - assert disabled.kind == "disabled" - assert disabled.lookup(None) is None - await disabled.async_store(None, object()) - assert await disabled.async_lookup(None) is None - assert selected.lookup(request()) == {"answer": 2} - - -async def test_python_callback_preserves_identity_caller_task_context_and_errors() -> None: - context: Final = contextvars.ContextVar("cache_context", default="caller") - caller: Final = asyncio.current_task() - sentinel: Final = object() - failure: Final = RuntimeError("callback failed") - - class CustomCache: - async def async_get_cache(self, *, marker: object) -> object: - assert marker is sentinel - assert asyncio.current_task() is caller - context.set("callback") - return marker - - async def async_add_cache(self, response: object, *, marker: object) -> None: - assert response is sentinel - assert marker is sentinel - raise failure - - namespace: Final = SimpleNamespace(cache=CustomCache()) - binding: Final = _CacheTestResolver(namespace).resolve() - assert binding.kind == "python_callback" - assert await binding.async_lookup(None, callback_kwargs={"marker": sentinel}) is sentinel - assert context.get() == "callback" - with pytest.raises(RuntimeError) as caught: - await binding.async_store(None, sentinel, callback_kwargs={"marker": sentinel}) - assert caught.value is failure - - -async def test_callback_cancellation_stays_in_the_callers_task() -> None: - entered: Final = asyncio.Event() - finished: Final = asyncio.Event() - - class CustomCache: - async def async_get_cache(self) -> None: - entered.set() - try: - await asyncio.Future() - finally: - finished.set() - - binding: Final = _CacheTestResolver(SimpleNamespace(cache=CustomCache())).resolve() - - async def lookup() -> object: - return await binding.async_lookup(None, callback_kwargs={}) - - task: Final = asyncio.create_task(lookup()) - await entered.wait() - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - assert finished.is_set() - - -def test_registered_facade_uses_native_and_instance_overrides_fall_back() -> None: - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - handle: Final = _CacheTestHandle.memory() - handle._bind_facade(facade) - resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) - native: Final = resolver.resolve() - assert native.kind == "native" - native.store(request(), {"source": "native"}) - assert native.lookup(request()) == {"source": "native"} - assert cast(CacheLookup, facade).get_cache(cache_key="key") is None - sentinel: Final = object() - - def outer_override(**_kwargs: object) -> object: - return sentinel - - def backend_override(*_args: object, **_kwargs: object) -> dict[str, str]: - return {"source": "override"} - - with rebound(facade, "get_cache", outer_override): - fallback: Final = resolver.resolve() - assert fallback.kind == "python_callback" - assert fallback.lookup(None, callback_kwargs={"cache_key": "key"}) is sentinel - assert resolver.resolve().kind == "python_callback" - delattr(facade, "get_cache") - assert resolver.resolve().kind == "native" - with rebound(facade.cache, "get_cache", backend_override): - backend_fallback: Final = resolver.resolve() - assert backend_fallback.kind == "python_callback" - assert backend_fallback.lookup(None, callback_kwargs={"cache_key": "key"}) == {"source": "override"} - - -def test_facade_subclasses_backend_replacement_and_configuration_changes_are_not_bypassed() -> None: - class CustomCache(Cache): - pass - - handle: Final = _CacheTestHandle.memory() - with pytest.raises(TypeError): - handle._bind_facade(CustomCache(type=LiteLLMCacheType.LOCAL)) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - handle._bind_facade(facade) - resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) - with rebound(facade, "cache", InMemoryCache()): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "semantic_cache_scope", "end_user"): - assert resolver.resolve().kind == "python_callback" - - def custom_key(**_kwargs: object) -> str: - return "custom" - - with rebound(facade, "get_cache_key", custom_key): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "python_callback" - delattr(facade, "get_cache_key") - assert resolver.resolve().kind == "native" - - -def test_resolver_and_callback_cycles_can_be_collected() -> None: - class CustomCache: - pass - - def cyclic_reference() -> weakref.ReferenceType[CustomCache]: - callback: Final = CustomCache() - namespace: Final = SimpleNamespace(cache=callback) - binding: Final = _CacheTestResolver(namespace).resolve() - setattr(callback, "binding", binding) - return weakref.ref(callback) - - reference: Final = cyclic_reference() - gc.collect() - assert reference() is None - - -async def test_redis_reads_python_sync_and_async_entries_and_writes_without_hidden_prefix(redis_url: str) -> None: - client: Final = redis.Redis.from_url(redis_url) - namespace: Final = SimpleNamespace(cache=_CacheTestHandle.redis(redis_url, namespace="team")) - binding: Final = _CacheTestResolver(namespace).resolve() - response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} - envelope: Final = {"timestamp": time.time(), "response": json.dumps(response)} - client.set("team:sync", str(envelope)) - client.set("team:async", json.dumps({"timestamp": time.time(), "response": response})) - client.set("team:raw", json.dumps(response)) - client.set("team:invalid", "not a cache entry") - assert binding.lookup(request("sync")) == response - assert await binding.async_lookup(request("team:async")) == response - assert binding.lookup(request("raw")) == response - assert await binding.async_lookup(request("invalid")) is None - await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) - stored: Final = client.get("team:native") - assert isinstance(stored, bytes) - assert json.loads(stored)["response"] == response - assert 0 < client.ttl("team:native") <= 12 - assert client.get("litellm-cache:team:native") is None - assert client.get("team:team:async") is None - client.close() - - -def test_invalid_duration_and_request_shape_fail_before_storage() -> None: - binding: Final = _CacheTestResolver(SimpleNamespace(cache=_CacheTestHandle.memory())).resolve() - for seconds in (-1.0, float("nan"), float("inf")): - with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): - binding.store({**request(), "ttl_seconds": seconds}, {"answer": 1}) - assert binding.lookup(request()) is None - with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): - _CacheTestHandle.memory(ttl_seconds=-1) - - -async def test_memory_size_policy_is_applied_by_the_native_host() -> None: - handle: Final = _CacheTestHandle.memory(capacity=2, max_entry_bytes=128) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=handle)).resolve() - small: Final = {"answer": "ok"} - binding.store(request("small"), small) - assert await binding.async_lookup(request("small")) == small - await binding.async_store(request("large"), {"answer": "x" * 256}) - assert binding.lookup(request("large")) is None - assert binding.lookup(request("small")) == small - disabled: Final = _CacheTestResolver(SimpleNamespace(cache=_CacheTestHandle.memory(capacity=0))).resolve() - await disabled.async_store(request(), small) - assert await disabled.async_lookup(request()) is None - - -async def test_native_batch_lookup_and_store_report_partial_hits() -> None: - binding: Final = _CacheTestResolver(SimpleNamespace(cache=_CacheTestHandle.memory())).resolve() - requests: Final = [request("hit"), request("miss"), request("disabled")] - requests[2]["controls"] = { - "supported_call_type": True, - "configured": True, - "native_backend": True, - "default_on": True, - "caching": False, - "no_cache": False, - "no_store": False, - "use_cache": False, - } - await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}]) - - partial: Final = await binding.async_lookup_batch(requests) - - assert partial == { - "values": [{"value": 1}, {"value": 2}, None], - "missing_indices": [2], - } - - -async def test_python_batch_callbacks_use_the_builtin_cache_api() -> None: - result: Final = object() - marker: Final = object() - - class CustomCache(Cache): - def get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object: - return ("sync", kwargs) - - async def async_get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object: - return ("async", kwargs) - - async def async_add_cache_pipeline( - self, result: object, dynamic_cache_object: object = None, **kwargs: object - ) -> object: - return result, kwargs - - binding: Final = _CacheTestResolver(SimpleNamespace(cache=CustomCache(type=LiteLLMCacheType.LOCAL))).resolve() - assert binding.kind == "python_callback" - requests: Final = [request("first"), request("second")] - kwargs: Final = [{"cache_key": "first"}, {"cache_key": "second"}] - - assert binding.lookup_batch(requests, callback_kwargs=kwargs) == [("sync", kwargs[0]), ("sync", kwargs[1])] - assert await binding.async_lookup_batch(requests, callback_kwargs=kwargs) == [ - ("async", kwargs[0]), - ("async", kwargs[1]), - ] - with pytest.raises(ValueError, match="equal lengths"): - binding.lookup_batch(requests, callback_kwargs=kwargs[:1]) - with pytest.raises(TypeError, match="callback_result"): - await binding.async_store_batch(requests, [1, 2], callback_kwargs={"marker": marker}) - stored: Final = cast( - tuple[object, dict[str, object]], - await binding.async_store_batch(requests, [1, 2], callback_result=result, callback_kwargs={"marker": marker}), - ) - assert stored[0] is result - assert stored[1] == {"marker": marker} - - -async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None: - async def ping() -> str: - return "pong" - - cache: Final = Cache(type=LiteLLMCacheType.LOCAL) - cache.cache.set_cache("key", "value") - binding: Final = _CacheTestResolver(SimpleNamespace(cache=cache)).resolve() - assert binding.kind == "python_callback" - - setattr(cache.cache, "ping", ping) - assert await binding.ping() == "pong" - await binding.async_flush() - assert cache.cache.get_cache("key") is None - - -def test_facade_registration_rejects_mismatched_capacity() -> None: - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - with pytest.raises(TypeError, match="capacities must match"): - _CacheTestHandle.memory(capacity=7)._bind_facade(facade) - - -def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None: - backend: Final = azure_blob_facade.cache - assert isinstance(backend, AzureBlobCache) - handle: Final = azure_blob_handle(azure_blob_facade) - assert handle.backend == "azure-blob" - account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}") - with pytest.raises(TypeError, match="containers must match"): - _native._CacheTestHandle.azure_blob( - account_url, f"{backend.container_client.container_name}-other" - )._bind_facade(azure_blob_facade) - handle._bind_facade(azure_blob_facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)) - native: Final = resolver.resolve() - assert native.kind == "native" - - response: Final = { - "choices": [{"text": "caf\u00e9 \u2603"}], - "usage": {"total_tokens": 3}, - "flag": True, - "empty": None, - } - native.store({**request("sync"), "ttl_seconds": 0.001}, response) - native.store(request("sync"), {"choices": [{"text": "second"}]}) - time.sleep(0.01) - stored: Final = json.loads(backend.container_client.download_blob("sync").readall()) - assert stored["response"] == response - assert isinstance(stored["timestamp"], float) - assert native.lookup(request("sync")) == response - assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response - - backend.set_cache("python", {"timestamp": time.time(), "response": response}) - backend.set_cache("legacy", "bare legacy value") - backend.container_client.upload_blob("invalid", b"{not json", overwrite=True) - assert native.lookup(request("python")) == response - assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy") - assert native.lookup_batch([request("python"), request("missing"), request("invalid"), request("sync")]) == { - "values": [response, None, None, response], - "missing_indices": [1, 2], - } - - with rebound(azure_blob_facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" - with rebound(backend, "container_client", ContainerClient.from_container_url(backend.container_client.url)): - assert resolver.resolve().kind == "python_callback" - - def custom_get(*_args: object, **_kwargs: object) -> None: - return None - - with rebound(backend, "get_cache", custom_get): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "python_callback" - assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response - - class CustomBlobCache(AzureBlobCache): - pass - - with rebound(azure_blob_facade, "cache", CustomBlobCache(account_url, backend.container_client.container_name)): - assert resolver.resolve().kind == "python_callback" - with pytest.raises(TypeError): - azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) - - -async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_python(azure_blob_facade: Cache) -> None: - backend: Final = azure_blob_facade.cache - assert isinstance(backend, AzureBlobCache) - azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)).resolve() - assert binding.kind == "native" - ping: Final = cast(dict[str, object], await binding.ping()) - assert ping["status"] == "success", ping - - await binding.async_store(request("async"), {"value": 1}) - await binding.async_store({**request("async"), "ttl_seconds": 0.001}, {"value": 2}) - time.sleep(0.01) - assert await binding.async_lookup(request("async")) == {"value": 2} - assert await backend.async_get_cache("async") == json.loads( - backend.container_client.download_blob("async").readall() - ) - assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="async") == {"value": 2} - - await binding.async_store_batch([request("first"), request("second")], [{"value": 3}, {"value": 4}]) - assert await binding.async_lookup_batch([request("second"), request("missing"), request("first")]) == { - "values": [{"value": 4}, None, {"value": 3}], - "missing_indices": [1], - } - await binding.async_flush() - assert [blob.name for blob in backend.container_client.list_blobs()] == [] - assert await binding.async_lookup(request("async")) is None - - -async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: - parsed: Final = urlparse(redis_url) - with rebound(litellm, "default_redis_ttl", 60): - facade: Final = Cache( - type=LiteLLMCacheType.REDIS, - host=parsed.hostname, - port=str(parsed.port), - redis_flush_size=2, - ) - with pytest.raises(TypeError, match="default TTLs must match"): - _CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade) - with pytest.raises(TypeError, match="namespaces must match"): - _CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade) - _CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(redis_url) - - with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "ssl": True}): - assert _CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" - - pool: Final = facade.cache.redis_client.connection_pool - with rebound(pool, "connection_kwargs", {**pool.connection_kwargs, "db": 1}): - assert _CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" - - await binding.async_store(request("first"), {"value": 1}) - assert client.get("first") is None - await binding.async_store(request("second"), {"value": 2}) - - assert client.get("first") is not None - assert client.get("second") is not None - await facade.cache.disconnect() - client.close() - - -async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_path: Path) -> None: - disk_cache: Final = DiskCache(disk_cache_dir=str(tmp_path)) - response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} - disk_cache.disk_cache.set( - "sync", - {"timestamp": time.time(), "response": json.dumps(response)}, - ) - disk_cache.disk_cache.set("async", json.dumps({"timestamp": time.time(), "response": response})) - disk_cache.disk_cache.set("raw", json.dumps(response)) - disk_cache.disk_cache.set("invalid", "not a cache entry") - disk_cache.disk_cache.set( - "large", - {"timestamp": time.time(), "response": {"text": "x" * 70_000}}, - ) - binding: Final = _native._CacheTestResolver( - SimpleNamespace(cache=_native._CacheTestHandle.disk(str(tmp_path))) - ).resolve() - - assert binding.lookup(request("sync")) == response - assert await binding.async_lookup(request("async")) == response - assert binding.lookup(request("raw")) == response - assert await binding.async_lookup(request("invalid")) is None - assert binding.lookup(request("large")) == {"text": "x" * 70_000} - - await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) - stored_response: Final = disk_cache.get_cache("native") - assert isinstance(stored_response, dict) - assert stored_response["response"] == response - stored, expire_time = disk_cache.disk_cache.get("native", expire_time=True) - assert stored is not None - assert time.time() < expire_time <= time.time() + 12.0 - await binding.async_store(request("no-ttl"), response) - _, no_expiry = disk_cache.disk_cache.get("no-ttl", expire_time=True) - assert no_expiry is None - - -async def test_disk_entries_survive_a_fresh_handle_and_expire_on_time(tmp_path: Path) -> None: - first: Final = _native._CacheTestResolver( - SimpleNamespace(cache=_native._CacheTestHandle.disk(str(tmp_path))) - ).resolve() - await first.async_store(request("persistent"), {"value": "persistent"}) - await first.async_store({**request("expiring"), "ttl_seconds": 0.3}, {"value": "expiring"}) - fresh: Final = _native._CacheTestResolver( - SimpleNamespace(cache=_native._CacheTestHandle.disk(str(tmp_path))) - ).resolve() - assert fresh.lookup(request("persistent")) == {"value": "persistent"} - assert fresh.lookup(request("expiring")) == {"value": "expiring"} - await asyncio.sleep(0.4) - assert fresh.lookup(request("expiring")) is None - assert fresh.lookup(request("persistent")) == {"value": "persistent"} - - -def test_disk_facade_registers_and_store_changes_fall_back(tmp_path: Path) -> None: - facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) - with pytest.raises(TypeError, match="directories must match"): - _native._CacheTestHandle.disk(str(tmp_path / "other"))._bind_facade(facade) - handle: Final = _native._CacheTestHandle.disk(str(tmp_path)) - handle._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) - binding: Final = resolver.resolve() - assert binding.kind == "native" - binding.store(request("native"), {"value": "native"}) - assert facade.get_cache(cache_key="native") == {"value": "native"} - - with rebound(facade.cache, "disk_cache", diskcache.Cache(str(tmp_path))): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "native" - - class CustomDiskCache(DiskCache): - pass - - with rebound(facade, "cache", CustomDiskCache(disk_cache_dir=str(tmp_path))): - assert resolver.resolve().kind == "python_callback" - - class CustomStore(diskcache.Cache): - pass - - custom_facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) - custom_facade.cache.disk_cache = CustomStore(str(tmp_path)) - with pytest.raises(TypeError, match="built-in diskcache store"): - _native._CacheTestHandle.disk(str(tmp_path))._bind_facade(custom_facade) - - -async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path: Path) -> None: - binding: Final = _native._CacheTestResolver( - SimpleNamespace(cache=_native._CacheTestHandle.disk(str(tmp_path))) - ).resolve() - requests: Final = [request("hit"), request("miss"), request("disabled")] - requests[2]["controls"] = { - "supported_call_type": True, - "configured": True, - "native_backend": True, - "default_on": True, - "caching": False, - "no_cache": False, - "no_store": False, - "use_cache": False, - } - await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}]) - - partial: Final = await binding.async_lookup_batch(requests) - - assert partial == { - "values": [{"value": 1}, {"value": 2}, None], - "missing_indices": [2], - } - - -@pytest.fixture -def s3_stub() -> Generator[S3Stub]: - stub: Final = S3Stub() - try: - yield stub - finally: - stub.close() - - -def python_s3(url: str) -> S3Cache: - return S3Cache( - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - ) - - -async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: S3Stub) -> None: - python_cache: Final = python_s3(s3_stub.url) - response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} - python_cache.set_cache("sync:key", {"timestamp": time.time(), "response": response}, ttl=90) - python_cache.set_cache("plain", {"timestamp": time.time(), "response": response}) - s3_stub.put_object("team/malformed", b"not a cache entry") - s3_stub.put_object( - "team/expired", - json.dumps({"timestamp": time.time(), "response": response}).encode(), - {"expires": "Thu, 01 Jan 1970 00:00:00 GMT"}, - ) - binding: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.s3( - "cache-bucket", - region="us-east-1", - endpoint_url=s3_stub.url, - key_prefix="team/", - access_key_id="key", - secret_access_key="secret", - ) - ) - ).resolve() - - assert binding.lookup(request("sync:key")) == response - assert await binding.async_lookup(request("plain")) == response - assert binding.lookup(request("malformed")) is None - assert binding.lookup(request("expired")) is None - assert binding.lookup(request("absent")) is None - - binding.store({**request("native:key"), "ttl_seconds": 90.0}, response) - await binding.async_store(request("no_ttl"), response) - stored: Final = s3_stub.objects["team/native/key"] - assert stored.headers["content-type"] == "application/json" - assert stored.headers["content-language"] == "en" - assert stored.headers["content-disposition"] == 'inline; filename="team/native/key.json"' - assert stored.headers["cache-control"] == "immutable, max-age=90, s-maxage=90" - expires: Final = cast(datetime, s3_stub.expires("team/native/key")) - remaining: Final = (expires - datetime.now(expires.tzinfo)).total_seconds() - assert 60 < remaining <= 91 - no_ttl: Final = s3_stub.objects["team/no_ttl"] - assert no_ttl.headers["cache-control"] == "immutable, max-age=31536000, s-maxage=31536000" - assert "expires" not in no_ttl.headers - assert python_cache.get_cache("native:key")["response"] == response - - partial: Final = await binding.async_lookup_batch([request("native:key"), request("absent"), request("malformed")]) - assert partial == {"values": [response, None, None], "missing_indices": [1, 2]} - - -def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_stub: S3Stub) -> None: - facade: Final = Cache( - type=LiteLLMCacheType.S3, - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=s3_stub.url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - ) - handle: Final = _native._CacheTestHandle.s3( - "cache-bucket", - region="us-east-1", - endpoint_url=s3_stub.url, - key_prefix="team/", - access_key_id="key", - secret_access_key="secret", - ) - with pytest.raises(TypeError, match="buckets must match"): - _native._CacheTestHandle.s3("other", region="us-east-1", endpoint_url=s3_stub.url)._bind_facade(facade) - with pytest.raises(TypeError, match="key prefixes must match"): - _native._CacheTestHandle.s3( - "cache-bucket", region="us-east-1", endpoint_url=s3_stub.url, key_prefix="other/" - )._bind_facade(facade) - handle._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) - binding: Final = resolver.resolve() - assert binding.kind == "native" - - handler: Final = Mock() - facade.cache.s3_client.meta.events.register("before-call.s3.*", handler) - binding.store(request("native"), {"answer": 1}) - assert binding.lookup(request("native")) == {"answer": 1} - assert handler.call_count == 0 - assert "team/native" in s3_stub.objects - - with rebound(facade.cache, "bucket_name", "other"): - assert resolver.resolve().kind == "python_callback" - other_client: Final = boto3.client( - "s3", - region_name="us-east-1", - endpoint_url=s3_stub.url, - aws_access_key_id="key", - aws_secret_access_key="secret", - ) - with rebound(facade.cache, "s3_client", other_client): - assert resolver.resolve().kind == "python_callback" - - class CustomS3Cache(S3Cache): - pass - - subclassed: Final = Cache( - type=LiteLLMCacheType.S3, - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=s3_stub.url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - ) - subclassed.cache = CustomS3Cache( - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=s3_stub.url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - ) - with pytest.raises(TypeError): - handle._bind_facade(subclassed) - assert _native._CacheTestResolver(SimpleNamespace(cache=subclassed)).resolve().kind == "python_callback" - - -def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) -> None: - handle: Final = _native._CacheTestHandle.s3( - "cache-bucket", - region="us-east-1", - endpoint_url=s3_stub.url, - key_prefix="team/", - access_key_id="key", - secret_access_key="secret", - ) - unverified: Final = Cache( - type=LiteLLMCacheType.S3, - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url="https://s3.example.test", - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - s3_verify=False, - ) - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(unverified) - proxied: Final = Cache( - type=LiteLLMCacheType.S3, - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=s3_stub.url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - s3_config=botocore.config.Config(proxies={"https": "http://proxy.test"}), - ) - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(proxied) - - -async def test_gcs_reads_python_entries_and_writes_python_compatible_objects( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) - monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} - fake_gcs.put( - "bucket", - "cache/sync", - json.dumps({"timestamp": time.time(), "response": json.dumps(response)}).encode(), - ) - fake_gcs.put("bucket", "cache/async", json.dumps({"timestamp": time.time(), "response": response}).encode()) - fake_gcs.put("bucket", "cache/raw", json.dumps(response).encode()) - fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") - binding: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - - assert binding.lookup(request("sync")) == response - assert await binding.async_lookup(request("async")) == response - assert binding.lookup(request("raw")) == response - assert await binding.async_lookup(request("invalid")) is None - assert binding.lookup(request("missing")) is None - - await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) - stored: Final = fake_gcs.objects[("bucket", "cache/native")] - stored_value: Final = cast(dict[str, object], json.loads(stored)) - assert stored_value["response"] == response - assert isinstance(stored_value["timestamp"], float) - upload: Final = next(item for item in fake_gcs.requests if item.method == "POST") - assert upload.path == "/upload/storage/v1/b/bucket/o" - assert upload.query == "uploadType=media&name=cache%2Fnative" - assert upload.headers["Authorization"] == f"Bearer {fake_gcs.token}" - assert upload.headers["Content-Type"] == "application/json" - upload_text: Final = f"{upload.path}?{upload.query}{upload.headers}" - assert "ttl" not in upload_text.lower() - assert "expiry" not in upload_text.lower() - download: Final = next(item for item in fake_gcs.requests if item.path.endswith("/cache%2Fsync")) - assert download.path == "/storage/v1/b/bucket/o/cache%2Fsync" - assert download.query == "alt=media" - - binding.store(request("sync2"), response) - assert binding.lookup(request("sync2")) == response - assert GCSCache(bucket_name="bucket", gcs_path="cache").key_prefix == "cache/" - assert GCSCache(bucket_name="bucket", gcs_path="cache/").key_prefix == "cache/" - assert GCSCache(bucket_name="bucket").key_prefix == "" - - -async def test_gcs_batch_lookup_preserves_order_and_treats_malformed_entries_as_misses(fake_gcs: FakeGcs) -> None: - fake_gcs.put("bucket", "cache/hit", json.dumps({"timestamp": time.time(), "response": {"value": 1}}).encode()) - fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") - binding: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - requests: Final = [request("hit"), request("missing"), request("invalid")] - expected: Final = {"values": [{"value": 1}, None, None], "missing_indices": [1, 2]} - - assert await binding.async_lookup_batch(requests) == expected - assert binding.lookup_batch(requests) == expected - await binding.async_store_batch([request("first"), request("second")], [{"value": 1}, {"value": 2}]) - assert ("bucket", "cache/first") in fake_gcs.objects - assert ("bucket", "cache/second") in fake_gcs.objects - - -async def test_gcs_facade_binds_only_exact_matching_configuration( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) - monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/nonexistent") - facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - assert type(facade.cache) is GCSCache - - mismatched_bucket: Final = _native._CacheTestHandle.gcs( - "other", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="buckets must match"): - mismatched_bucket._bind_facade(facade) - mismatched_prefix: Final = _native._CacheTestHandle.gcs( - "bucket", - gcs_path="x", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="key prefixes must match"): - mismatched_prefix._bind_facade(facade) - mismatched_credentials: Final = _native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - path_service_account="sa.json", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="credentials must match"): - mismatched_credentials._bind_facade(facade) - with pytest.raises(TypeError, match="types must match"): - _native._CacheTestHandle.memory()._bind_facade(facade) - - matching: Final = _native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - matching._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) - binding: Final = resolver.resolve() - assert binding.kind == "native" - await binding.async_store(request("native"), {"value": "native"}) - assert await binding.async_lookup(request("native")) == {"value": "native"} - assert cast(CacheLookup, facade).get_cache(cache_key="native") is None - - with rebound(facade.cache, "bucket_name", "other"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "key_prefix", "x/"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "path_service_account", "sa.json"): - assert resolver.resolve().kind == "python_callback" - - def no_get_cache(*args: object, **kwargs: object) -> None: - return None - - with rebound(facade.cache, "get_cache", no_get_cache): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" - - class CustomGcs(GCSCache): - pass - - with rebound(facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): - assert resolver.resolve().kind == "python_callback" - custom_facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - with rebound(custom_facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): - with pytest.raises(TypeError, match="types must match"): - matching._bind_facade(custom_facade) - - missing_bucket: Final = Cache(type=LiteLLMCacheType.GCS) - with pytest.raises(TypeError, match="requires a configured bucket name"): - matching._bind_facade(missing_bucket) - - -async def test_gcs_flush_is_a_no_op_and_ping_is_not_implemented( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) - monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - binding: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - await binding.async_store(request("key"), {"value": "stored"}) - await binding.async_flush() - assert ("bucket", "cache/key") in fake_gcs.objects - assert await binding.async_lookup(request("key")) == {"value": "stored"} - with pytest.raises(NotImplementedError): - await binding.ping() - - facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - with pytest.raises(AttributeError): - await facade.ping() - assert cast(CacheLookup, facade.cache).flush_cache() is None - - -async def test_gcs_unauthorized_and_server_errors_surface_as_runtime_errors(fake_gcs: FakeGcs) -> None: - wrong_token: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token="wrong-token", - ) - ) - ).resolve() - with pytest.raises(RuntimeError): - wrong_token.lookup(request("missing")) - assert not fake_gcs.objects - - binding: Final = _native._CacheTestResolver( - SimpleNamespace( - cache=_native._CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - with pytest.raises(RuntimeError): - binding.lookup(request("server-error")) - assert binding.lookup(request("missing")) is None - - -async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_natively( - cluster_nodes: tuple[tuple[str, int], ...], -) -> None: - startup_nodes: Final = [{"host": host, "port": port} for host, port in cluster_nodes] - url: Final = f"redis://{cluster_nodes[0][0]}:{cluster_nodes[0][1]}" - with rebound(litellm, "default_redis_ttl", 60): - facade: Final = Cache(type=LiteLLMCacheType.REDIS, redis_startup_nodes=startup_nodes, namespace="parity") - assert type(facade.cache) is RedisClusterCache - with pytest.raises(TypeError, match="types must match"): - _native._CacheTestHandle.redis(url, namespace="parity")._bind_facade(facade) - _native._CacheTestHandle.redis(url, namespace="parity", startup_nodes=list(cluster_nodes))._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) - assert resolver.resolve().kind == "native" - - manager: Final = facade.cache.redis_client.nodes_manager - with rebound(manager, "connection_kwargs", {**manager.connection_kwargs, "db": 1}): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "startup_nodes": startup_nodes[:1]}): - assert resolver.resolve().kind == "python_callback" - binding: Final = resolver.resolve() - assert binding.kind == "native" - - client: Final = redis.RedisCluster(startup_nodes=[redis.cluster.ClusterNode(*node) for node in cluster_nodes]) - keys: Final = tuple(f"slot-{index}" for index in range(12)) - slots: Final = {client.keyslot(f"parity:{key}") for key in keys} - assert len(slots) > 1, slots - requests: Final = [request(key) for key in keys] - values: Final = [{"index": index} for index in range(len(keys))] - await binding.async_store_batch(requests, values) - client.set("parity:slot-3", "not a cache entry") - client.set("parity:slot-7", json.dumps({"timestamp": time.time(), "response": {"index": 7, "python": True}})) - - batch: Final = await binding.async_lookup_batch(requests) - assert batch == { - "values": [ - None if index == 3 else {"index": 7, "python": True} if index == 7 else value - for index, value in enumerate(values) - ], - "missing_indices": [3], - } - assert facade.cache.get_cache("parity:slot-0")["response"] == {"index": 0} - assert (await facade.cache.async_get_cache("parity:slot-11"))["response"] == {"index": 11} - assert facade.cache.redis_client.mget_nonatomic([f"parity:{key}" for key in keys[:2]]) == [ - client.get("parity:slot-0"), - client.get("parity:slot-1"), - ] - - await binding.async_store({**request("pinned"), "ttl_seconds": 12.0}, {"pinned": True}) - assert 0 < client.ttl("parity:pinned") <= 12 - client.set("unscoped", "stays") - - await binding.async_flush() - - remaining: Final = tuple( - sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node)) - ) - assert remaining == (), remaining - assert client.get("unscoped") == b"stays" - client.delete("unscoped") - client.close() - facade.cache.redis_client.close() - - -PARAPHRASE_MARKER: Final = " (paraphrase)" -SEMANTIC_EMBEDDING_MODEL: Final = "semantic-test/deterministic" -SEMANTIC_INDEX_PREFIX: Final = "litellm_test_semantic_" -SEMANTIC_CONTEXT: Final = contextvars.ContextVar("semantic_test_context", default="unset") - - -def _normalized(vector: list[float]) -> list[float]: - norm: Final = math.sqrt(sum(component * component for component in vector)) - return [component / norm for component in vector] - - -def _base_embedding(prompt: str) -> list[float]: - digest: Final = hashlib.sha256(prompt.encode("utf-8")).digest() - return _normalized([float(digest[index] + 1) for index in range(8)]) - - -def _semantic_embedding(prompt: str) -> list[float]: - if PARAPHRASE_MARKER not in prompt: - return _base_embedding(prompt) - base: Final = _base_embedding(prompt.replace(PARAPHRASE_MARKER, "").strip()) - pivot: Final = min(range(8), key=lambda index: abs(base[index])) - direction: Final = _normalized( - [(1.0 - base[pivot] * base[pivot]) if index == pivot else -base[index] * base[pivot] for index in range(8)] - ) - # Rotating an orthogonal unit direction by 0.329 produces ~0.05 cosine distance - return _normalized([base[index] + 0.329 * direction[index] for index in range(8)]) - - -class DeterministicEmbedding(litellm.CustomLLM): - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - self.async_calls: list[dict[str, object]] = [] - self.entered = asyncio.Event() - self.gate: asyncio.Event | None = None - - def _respond( - self, - model: str, - input: object, - model_response: EmbeddingResponse, - ) -> EmbeddingResponse: - texts: Final = cast(list[object], input if isinstance(input, list) else [input]) - self.calls.append({"model": model, "input": texts}) - model_response.model = model - model_response.data = [ - {"object": "embedding", "index": index, "embedding": _semantic_embedding(str(text))} - for index, text in enumerate(texts) - ] - return model_response - - def embedding( - self, - model: str, - input: list[object], - model_response: EmbeddingResponse, - print_verbose: Callable[..., object], - logging_obj: object, - optional_params: dict[str, object], - api_key: object = None, - api_base: object = None, - timeout: object = None, - litellm_params: object = None, - ) -> EmbeddingResponse: - return self._respond(model, input, model_response) - - async def aembedding( - self, - model: str, - input: list[object], - model_response: EmbeddingResponse, - print_verbose: Callable[..., object], - logging_obj: object, - optional_params: dict[str, object], - api_key: object = None, - api_base: object = None, - timeout: object = None, - litellm_params: object = None, - ) -> EmbeddingResponse: - texts: Final = cast(list[object], input if isinstance(input, list) else [input]) - self.async_calls.append( - { - "model": model, - "input": texts, - "task": asyncio.current_task(), - "context": SEMANTIC_CONTEXT.get(), - } - ) - SEMANTIC_CONTEXT.set("written-in-aembedding") - self.entered.set() - if self.gate is not None: - await self.gate.wait() - return self._respond(model, input, model_response) - - -@pytest.fixture -def semantic_embedding() -> Generator[DeterministicEmbedding]: - handler: Final = DeterministicEmbedding() - with ExitStack() as stack: - stack.enter_context( - rebound( - litellm, - "custom_provider_map", - [ - *litellm.custom_provider_map, - cast( - CustomLLMItem, - {"provider": "semantic-test", "custom_handler": handler}, - ), - ], - ) - ) - stack.enter_context( - rebound( - litellm, - "_custom_providers", # pyright: ignore[reportPrivateUsage] # no public provider-registration hook - [*litellm._custom_providers, "semantic-test"], # pyright: ignore[reportPrivateUsage] # no public provider-registration hook - ) - ) - stack.enter_context(rebound(litellm, "provider_list", [*litellm.provider_list, "semantic-test"])) - yield handler - - -@pytest.fixture -def redis_stack() -> Generator[tuple[str, str]]: - url: Final = os.environ.get("LITELLM_REDIS_STACK_URL") - if url is None: - pytest.skip("LITELLM_REDIS_STACK_URL is not set") - index: Final = f"{SEMANTIC_INDEX_PREFIX}{uuid4().hex}" - yield url, index - client: Final = redis.Redis.from_url(url) - try: - client.execute_command("FT.DROPINDEX", index, "DD") # pyright: ignore[reportUnknownMemberType] # redis-py leaves execute_command partially unknown - except redis.RedisError: - pass - client.close() - - -def semantic_request(key: str, prompt: str, **extra: object) -> dict[str, object]: - return { - "key": {"preset": key}, - "messages": [{"role": "user", "content": prompt}], - **extra, - } - - -def semantic_messages(prompt: str) -> list[dict[str, object]]: - return [{"role": "user", "content": prompt}] - - -def semantic_entry_id(prompt: str, tag: str) -> str: - return hashlib.sha256(f"{prompt}litellm_cache_key{tag}".encode()).hexdigest() - - -def semantic_facade(url: str, index: str, *, similarity_threshold: float = 0.8) -> Cache: - facade: Final = Cache( - type=LiteLLMCacheType.REDIS_SEMANTIC, - redis_url=url, - similarity_threshold=similarity_threshold, - redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, - redis_semantic_cache_index_name=index, - ) - _CacheTestHandle.redis_semantic(facade.cache)._bind_facade(facade) - return facade - - -def test_redis_semantic_constructor_identity_and_provenance( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - backend: Final = cast(RedisSemanticCache, facade.cache) - assert backend.__class__.__module__ == "litellm.caching.redis_semantic_cache" - assert type(backend) is RedisSemanticCache - assert backend._redis_url == url # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config - assert backend._index_name == index # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config - assert backend.similarity_threshold == 0.8 - assert backend.embedding_model == SEMANTIC_EMBEDDING_MODEL - handle: Final = cast(object, getattr(facade, "_native_cache_handle")) - assert isinstance(handle, _CacheTestHandle) - assert handle.backend == "redis_semantic" - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - assert binding.kind == "native" - - -def test_redis_semantic_native_and_python_sync_entries_share_one_layout( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - response: Final = {"choices": [{"text": "paris"}], "usage": {"total_tokens": 2}} - - binding.store(semantic_request("geo", "what is the capital of france"), response) - - native_hash_key: Final = f"{index}:{semantic_entry_id('what is the capital of france', 'geo')}" - stored: Final = client.hgetall(native_hash_key) - assert set(stored) == { - b"entry_id", - b"prompt", - b"response", - b"prompt_vector", - b"inserted_at", - b"updated_at", - b"litellm_cache_key", - }, stored - assert stored[b"entry_id"].decode() == native_hash_key.split(":", 1)[1] - assert stored[b"prompt"] == b"what is the capital of france" - assert stored[b"litellm_cache_key"] == b"geo" - assert len(stored[b"prompt_vector"]) == 32 - decoded: Final = cast(dict[str, object], json.loads(stored[b"response"])) - assert decoded["response"] == response - assert ( - cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "geo", messages=semantic_messages("what is the capital of france") - ) - == decoded - ) - assert semantic_embedding.calls == [ - {"model": "deterministic", "input": ["what is the capital of france"]}, - {"model": "deterministic", "input": ["what is the capital of france"]}, - {"model": "deterministic", "input": ["dimension test"]}, - ] - - cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "math", - json.dumps({"timestamp": 1700000000.0, "response": {"answer": 42}}), - messages=semantic_messages("what is 6 times 7"), - ) - python_hash_key: Final = f"{index}:{semantic_entry_id('what is 6 times 7', 'math')}" - assert json.loads(cast(bytes, client.hget(python_hash_key, "response"))) == { - "timestamp": 1700000000.0, - "response": {"answer": 42}, - } - assert binding.lookup(semantic_request("math", "what is 6 times 7")) == {"answer": 42} - client.close() - - -async def test_redis_semantic_async_paths_and_store_batch_share_one_layout( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - - await binding.async_store(semantic_request("async", "name a primary color"), {"answer": "blue"}) - hash_key: Final = f"{index}:{semantic_entry_id('name a primary color', 'async')}" - decoded: Final = cast(dict[str, object], json.loads(cast(bytes, client.hget(hash_key, "response")))) - python_read: Final = await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "async", messages=semantic_messages("name a primary color") - ) - assert python_read == decoded - - await binding.async_store_batch( - [ - semantic_request("batch-one", "first batch prompt"), - semantic_request("batch-two", "second batch prompt"), - ], - [{"answer": 1}, {"answer": 2}], - ) - expected: Final = { - key: json.loads(cast(bytes, client.hget(f"{index}:{semantic_entry_id(prompt, key)}", "response"))) - for key, prompt in ( - ("batch-one", "first batch prompt"), - ("batch-two", "second batch prompt"), - ) - } - for key, prompt in ( - ("batch-one", "first batch prompt"), - ("batch-two", "second batch prompt"), - ): - assert ( - cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - key, messages=semantic_messages(prompt) - ) - == expected[key] - ), key - - cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "async-python", - json.dumps({"timestamp": 1700000000.0, "response": {"answer": "python"}}), - messages=semantic_messages("python written prompt"), - ) - assert await binding.async_lookup(semantic_request("async-python", "python written prompt")) == {"answer": "python"} - client.close() - - -async def test_native_semantic_async_embedding_runs_inline_in_the_callers_task( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - assert binding.kind == "native" - caller: Final = asyncio.current_task() - SEMANTIC_CONTEXT.set("caller-sentinel") - response: Final = {"choices": [{"text": "paris"}]} - - await binding.async_store(semantic_request("inline", "what is the capital of france"), response) - assert ( - await binding.async_lookup(semantic_request("inline", f"what is the capital of france{PARAPHRASE_MARKER}")) - == response - ) - assert await binding.async_lookup(semantic_request("inline", "python written prompt")) is None - assert SEMANTIC_CONTEXT.get() == "written-in-aembedding" - assert semantic_embedding.async_calls == [ - { - "model": "deterministic", - "input": ["what is the capital of france"], - "task": caller, - "context": "caller-sentinel", - }, - { - "model": "deterministic", - "input": [f"what is the capital of france{PARAPHRASE_MARKER}"], - "task": caller, - "context": "written-in-aembedding", - }, - { - "model": "deterministic", - "input": ["python written prompt"], - "task": caller, - "context": "written-in-aembedding", - }, - ], semantic_embedding.async_calls - - -async def test_native_semantic_cancellation_during_embedding_skips_the_backend( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - assert binding.kind == "native" - semantic_embedding.gate = asyncio.Event() - - async def lookup() -> object: - return await binding.async_lookup(semantic_request("cancel", "cancelled prompt")) - - task: Final = asyncio.create_task(lookup()) - await semantic_embedding.entered.wait() - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - semantic_embedding.gate.set() - - assert len(semantic_embedding.async_calls) == 1 - assert ( - await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "cancel", messages=semantic_messages("cancelled prompt") - ) - is None - ) - - -def test_redis_semantic_similarity_tag_and_threshold_boundaries( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - - binding.store(semantic_request("sim", "tell me a joke"), {"answer": "haha"}) - paraphrase: Final = f"tell me a joke{PARAPHRASE_MARKER}" - assert binding.lookup(semantic_request("sim", paraphrase)) == {"answer": "haha"} - assert binding.lookup(semantic_request("sim", "an unrelated question about spreadsheets")) is None - assert binding.lookup(semantic_request("other-key", "tell me a joke")) is None - - strict: Final = semantic_facade(url, index, similarity_threshold=0.99) - strict_binding: Final = _CacheTestResolver(SimpleNamespace(cache=strict)).resolve() - assert strict_binding.lookup(semantic_request("sim", paraphrase)) is None - assert strict_binding.lookup(semantic_request("sim", "tell me a joke")) == {"answer": "haha"} - - -def test_redis_semantic_ttl_is_written_only_when_requested( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - - binding.store({**semantic_request("ttl", "ttl prompt"), "ttl_seconds": 12.0}, {"answer": 1}) - expiring: Final = f"{index}:{semantic_entry_id('ttl prompt', 'ttl')}" - assert 0 < client.ttl(expiring) <= 12 - - binding.store(semantic_request("ttl-none", "untimed prompt"), {"answer": 2}) - persistent: Final = f"{index}:{semantic_entry_id('untimed prompt', 'ttl-none')}" - assert client.ttl(persistent) == -1 - - binding.store( - {**semantic_request("ttl-fraction", "fractional prompt"), "ttl_seconds": 1.5}, - {"answer": 3}, - ) - fractional: Final = f"{index}:{semantic_entry_id('fractional prompt', 'ttl-fraction')}" - assert client.ttl(fractional) == 2 - client.close() - - -def test_redis_semantic_malformed_response_is_a_miss_for_both_readers( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - - binding.store(semantic_request("bad", "corrupt me"), {"answer": 1}) - hash_key: Final = f"{index}:{semantic_entry_id('corrupt me', 'bad')}" - client.hset(hash_key, "response", b"{not json") - assert binding.lookup(semantic_request("bad", "corrupt me")) is None - assert ( - cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class - "bad", messages=semantic_messages("corrupt me") - ) - is None - ) - client.close() - - -async def test_redis_semantic_unsupported_operations_raise_not_implemented( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - - with pytest.raises(NotImplementedError): - binding.lookup_batch([semantic_request("batch", "prompt one")]) - with pytest.raises(NotImplementedError): - await binding.async_lookup_batch([semantic_request("batch", "prompt one")]) - with pytest.raises(NotImplementedError): - await binding.async_flush() - with pytest.raises(NotImplementedError): - await binding.ping() - - -def test_redis_semantic_requests_without_prompt_are_noops( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - - binding.store(request("plain"), {"answer": 1}) - assert binding.lookup(request("plain")) is None - assert semantic_embedding.calls == [] - assert client.keys(f"{index}:*") == [] - client.close() - - -def test_redis_semantic_scope_overrides_the_tag_and_isolates_entries( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - client: Final = redis.Redis.from_url(url) - - scoped: Final = {**semantic_request("scoped", "scoped prompt"), "scope": "team-a"} - binding.store(scoped, {"answer": "kept"}) - hash_key: Final = f"{index}:{semantic_entry_id('scoped prompt', 'team-a')}" - assert client.hget(hash_key, "litellm_cache_key") == b"team-a" - assert binding.lookup(scoped) == {"answer": "kept"} - assert binding.lookup(semantic_request("scoped", "scoped prompt")) is None - assert binding.lookup({**scoped, "scope": "team-b"}) is None - client.close() - - -def test_redis_semantic_configuration_drift_falls_back_to_python( - redis_stack: tuple[str, str], - semantic_embedding: DeterministicEmbedding, - monkeypatch: pytest.MonkeyPatch, -) -> None: - url, index = redis_stack - facade: Final = semantic_facade(url, index) - resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) - assert resolver.resolve().kind == "native" - - with rebound(facade.cache, "similarity_threshold", 0.5): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "semantic_cache_scope", "end_user"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "embedding_model", "other-model"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "_index_name", "other-index"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "CACHE_KEY_FIELD_NAME", "other-field"): - assert resolver.resolve().kind == "python_callback" - - def patched_embedding(self: object, prompt: str, metadata: object = None) -> list[float]: - return _semantic_embedding(prompt) - - monkeypatch.setattr(RedisSemanticCache, "_get_embedding", patched_embedding) - assert resolver.resolve().kind == "python_callback" - - -def test_redis_semantic_handle_rejects_wrong_backends( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding -) -> None: - url, index = redis_stack - - class CustomSemanticCache(RedisSemanticCache): - pass - - with pytest.raises(TypeError, match="built-in RedisSemanticCache"): - _CacheTestHandle.redis_semantic(object()) - with pytest.raises(TypeError, match="built-in RedisSemanticCache"): - _CacheTestHandle.redis_semantic( - CustomSemanticCache( - redis_url=url, - similarity_threshold=0.8, - embedding_model=SEMANTIC_EMBEDDING_MODEL, - index_name=f"{index}_subclass", - ) - ) - - facade: Final = semantic_facade(url, index) - with pytest.raises(TypeError, match="backend types must match"): - _CacheTestHandle.redis(url)._bind_facade(facade) - - subclassed_facade: Final = Cache( - type=LiteLLMCacheType.REDIS_SEMANTIC, - redis_url=url, - similarity_threshold=0.8, - redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, - redis_semantic_cache_index_name=index, - ) - subclassed_facade.cache = CustomSemanticCache( # pyright: ignore[reportAttributeAccessIssue] # facade backend slot is not declared - redis_url=url, - similarity_threshold=0.8, - embedding_model=SEMANTIC_EMBEDDING_MODEL, - index_name=index, - ) - with pytest.raises(TypeError): - _CacheTestHandle.redis_semantic(subclassed_facade.cache)._bind_facade(subclassed_facade) - - replacement_facade: Final = Cache( - type=LiteLLMCacheType.REDIS_SEMANTIC, - redis_url=url, - similarity_threshold=0.8, - redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, - redis_semantic_cache_index_name=index, - ) - with pytest.raises(TypeError, match="must be the native embedder"): - _CacheTestHandle.redis_semantic(facade.cache)._bind_facade(replacement_facade) - - -def qdrant_facade(qdrant_url: str, collection_name: str) -> Cache: - return Cache( - type=LiteLLMCacheType.QDRANT_SEMANTIC, - qdrant_api_base=qdrant_url, - qdrant_collection_name=collection_name, - similarity_threshold=0.99, - qdrant_semantic_cache_embedding_model="text-embedding-3-small", - qdrant_semantic_cache_vector_size=8, - ) - - -def test_qdrant_semantic_facade_binds_native_and_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None: - del fake_embedding_endpoint - messages: Final = [{"role": "user", "content": "shared prompt"}] - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - facade.cache.set_cache( - "python-key", - {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, - messages=messages, - ) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - assert binding.kind == "native" - assert binding.lookup(qdrant_request("python-key", messages)) == {"id": "py"} - binding.store(qdrant_request("native-key", messages), {"id": "native"}) - python_value: Final = facade.cache.get_cache("native-key", messages=messages) - assert isinstance(python_value, dict) - assert python_value["response"] == {"id": "native"} - unrelated: Final = [{"role": "user", "content": "unrelated prompt"}] - assert binding.lookup(qdrant_request("native-key", unrelated)) is None - assert facade.cache.get_cache("native-key", messages=unrelated) is None - assert binding.lookup(qdrant_request("different-key", messages)) is None - assert facade.cache.get_cache("different-key", messages=messages) is None - - -async def test_qdrant_semantic_async_parity(qdrant_url: str, fake_embedding_endpoint: str) -> None: - del fake_embedding_endpoint - messages: Final = [{"role": "user", "content": "async prompt"}] - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - await facade.cache.async_set_cache( - "python-key", - {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, - messages=messages, - ) - assert await binding.async_lookup(qdrant_request("python-key", messages)) == {"id": "py"} - await binding.async_store(qdrant_request("native-key", messages), {"id": "native"}) - python_value: Final = await facade.cache.async_get_cache("native-key", messages=messages) - assert isinstance(python_value, dict) - assert python_value["response"] == {"id": "native"} - - -async def test_qdrant_semantic_async_store_batch_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None: - del fake_embedding_endpoint - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - entries: Final = [ - qdrant_request("batch-one", [{"role": "user", "content": "first batch prompt"}]), - qdrant_request("batch-two", [{"role": "user", "content": "second batch prompt"}]), - ] - await binding.async_store_batch(entries, [{"id": "one"}, {"id": "two"}]) - - assert binding.lookup(entries[0]) == {"id": "one"} - assert binding.lookup(entries[1]) == {"id": "two"} - assert (await facade.cache.async_get_cache("batch-one", messages=entries[0]["messages"]))["response"] == { - "id": "one" - } - assert (await facade.cache.async_get_cache("batch-two", messages=entries[1]["messages"]))["response"] == { - "id": "two" - } - - -async def test_qdrant_semantic_malformed_entries_and_unsupported_operations( - qdrant_url: str, fake_embedding_endpoint: str -) -> None: - del fake_embedding_endpoint - messages: Final = [{"role": "user", "content": "malformed prompt"}] - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - key: Final = "malformed-key" - response: Final = { - "points": [ - { - "id": str(uuid4()), - "vector": embedding_vector("malformed prompt"), - "payload": { - "litellm_cache_key": key, - "text": "malformed prompt", - "response": "not json", - }, - } - ] - } - facade.cache.sync_client.put( - url=f"{qdrant_url}/collections/{collection}/points", - headers=facade.cache.headers, - json=response, - ) - assert binding.lookup(qdrant_request(key, messages)) is None - with pytest.raises(RuntimeError, match="operation is not supported"): - binding.lookup_batch([qdrant_request(key, messages)]) - with pytest.raises(RuntimeError, match="operation is not supported"): - await binding.async_flush() - with pytest.raises(RuntimeError, match="operation is not supported"): - await binding.ping() - - -def test_qdrant_semantic_ignores_request_expiry(qdrant_url: str, fake_embedding_endpoint: str) -> None: - del fake_embedding_endpoint - messages: Final = [{"role": "user", "content": "persistent prompt"}] - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - binding.store(qdrant_request("persistent-key", messages, ttl_seconds=1.0), {"id": "persistent"}) - time.sleep(1.2) - assert binding.lookup(qdrant_request("persistent-key", messages)) == {"id": "persistent"} - python_value: Final = facade.cache.get_cache("persistent-key", messages=messages) - assert isinstance(python_value, dict) - assert python_value["response"] == {"id": "persistent"} - - -def test_qdrant_semantic_mutation_and_projection_fallback(qdrant_url: str, fake_embedding_endpoint: str) -> None: - del fake_embedding_endpoint - collection: Final = f"cache_{uuid4().hex}" - facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = _native._CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) - facade.cache.qdrant_api_key = "rotated" - assert _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" - facade.cache.similarity_threshold = 0.5 - assert _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" - unsupported: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") - unsupported.cache.embedding_max_input_tokens = 100 - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(unsupported) - unsupported.cache.embedding_max_input_tokens = None - unsupported.cache.qdrant_api_base = "http://127.0.0.1:7777" - with pytest.raises(TypeError, match="gRPC"): - handle._bind_facade(unsupported) - - -CacheFactory: TypeAlias = Callable[[], Cache] - - -def require_rust(monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType) -> None: - monkeypatch.setattr(catalog, "RULES", (CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({backend})),)) - - -def native_runtime(facade: Cache) -> ResponseCacheRuntime: - runtime: Final = facade._native_cache # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor - assert isinstance(runtime, ResponseCacheRuntime) - assert runtime.kind == "native" - return runtime - - -@pytest.fixture -def cache_factory(request: pytest.FixtureRequest, tmp_path: Path) -> CacheFactory: - backend: Final = cast(LiteLLMCacheType, request.param) - match backend: - case LiteLLMCacheType.LOCAL: - return lambda: Cache(type=backend) - case LiteLLMCacheType.DISK: - return lambda: Cache(type=backend, disk_cache_dir=str(tmp_path)) - case LiteLLMCacheType.REDIS: - parsed: Final = urlparse(cast(str, request.getfixturevalue("redis_url"))) - return lambda: Cache(type=backend, host=parsed.hostname, port=str(parsed.port)) - case LiteLLMCacheType.S3: - stub: Final = cast(S3Stub, request.getfixturevalue("s3_stub")) - return lambda: Cache( - type=backend, - s3_bucket_name="cache-bucket", - s3_region_name="us-east-1", - s3_endpoint_url=stub.url, - s3_aws_access_key_id="key", - s3_aws_secret_access_key="secret", - s3_path="team", - ) - case LiteLLMCacheType.GCS: - return lambda: Cache(type=backend, gcs_bucket_name="bucket", gcs_path="cache/") - case LiteLLMCacheType.REDIS_SEMANTIC: - return lambda: Cache( - type=backend, - redis_url="redis://127.0.0.1:6379", - similarity_threshold=0.8, - redis_semantic_cache_embedding_model="text-embedding-3-small", - ) - case LiteLLMCacheType.VALKEY_SEMANTIC: - return lambda: Cache(type=backend, redis_url="redis://127.0.0.1:6390/0", similarity_threshold=0.8) - case _: - raise AssertionError(f"no local factory for {backend}") - - -ROUND_TRIP_BACKENDS: Final = ( - LiteLLMCacheType.LOCAL, - LiteLLMCacheType.DISK, - LiteLLMCacheType.REDIS, - LiteLLMCacheType.S3, -) -SHARED_STORE_BACKENDS: Final = (LiteLLMCacheType.DISK, LiteLLMCacheType.REDIS, LiteLLMCacheType.S3) - - -def completion_kwargs(label: str) -> dict[str, object]: - return {"model": "gpt-4o", "messages": [{"role": "user", "content": f"{label} {uuid4().hex}"}]} - - -@pytest.mark.parametrize("backend", list(LiteLLMCacheType)) -def test_shipped_rules_keep_every_backend_on_python(backend: LiteLLMCacheType) -> None: - assert resolve_response_cache(cast(Cache, SimpleNamespace(type=backend))) is None - - -@pytest.mark.parametrize( - "cache_factory", - [ - LiteLLMCacheType.LOCAL, - LiteLLMCacheType.DISK, - LiteLLMCacheType.REDIS, - LiteLLMCacheType.S3, - LiteLLMCacheType.GCS, - LiteLLMCacheType.REDIS_SEMANTIC, - LiteLLMCacheType.VALKEY_SEMANTIC, - ], - indirect=True, -) -def test_shipped_rules_construct_python_backed_facades(cache_factory: CacheFactory) -> None: - assert cache_factory()._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor - - -@pytest.mark.parametrize( - "cache_factory", - [ - LiteLLMCacheType.LOCAL, - LiteLLMCacheType.DISK, - LiteLLMCacheType.REDIS, - LiteLLMCacheType.S3, - LiteLLMCacheType.GCS, - LiteLLMCacheType.REDIS_SEMANTIC, - LiteLLMCacheType.VALKEY_SEMANTIC, - ], - indirect=True, -) -def test_rust_required_rule_activates_the_native_backend( - cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest -) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - native_runtime(cache_factory()) - - -@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True) -async def test_facade_storage_calls_round_trip_through_the_native_backend( - cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest -) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - facade: Final = cache_factory() - native_runtime(facade) - - sync_kwargs: Final = completion_kwargs("sync") - facade.add_cache({"answer": 1}, **sync_kwargs) - assert facade.get_cache(**sync_kwargs) == {"answer": 1} - - async_kwargs: Final = completion_kwargs("async") - await facade.async_add_cache({"answer": 2}, **async_kwargs) - assert await facade.async_get_cache(**async_kwargs) == {"answer": 2} - assert facade.get_cache(**completion_kwargs("absent")) is None - - -async def test_memory_facade_writes_bypass_the_python_backend(monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.LOCAL) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - native_runtime(facade) - kwargs: Final = completion_kwargs("memory") - facade.add_cache({"answer": 1}, **kwargs) - assert facade.cache.get_cache(facade.get_cache_key(**kwargs)) is None - assert facade.get_cache(**kwargs) == {"answer": 1} - - -@pytest.mark.parametrize("cache_factory", SHARED_STORE_BACKENDS, indirect=True) -async def test_native_and_python_facades_share_one_wire_format( - cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest -) -> None: - python_facade: Final = cache_factory() - assert python_facade._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - native_facade: Final = cache_factory() - native_runtime(native_facade) - - native_written: Final = completion_kwargs("native") - native_facade.add_cache({"writer": "native"}, **native_written) - assert python_facade.get_cache(**native_written) == {"writer": "native"} - - python_written: Final = completion_kwargs("python") - python_facade.add_cache({"writer": "python"}, **python_written) - assert native_facade.get_cache(**python_written) == {"writer": "python"} - - async_native: Final = completion_kwargs("async-native") - await native_facade.async_add_cache({"writer": "async-native"}, **async_native) - assert await python_facade.async_get_cache(**async_native) == {"writer": "async-native"} - - async_python: Final = completion_kwargs("async-python") - await python_facade.async_add_cache({"writer": "async-python"}, **async_python) - assert await native_facade.async_get_cache(**async_python) == {"writer": "async-python"} - - -@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True) -async def test_embedding_pipeline_stores_one_native_entry_per_input( - cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest -) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - facade: Final = cache_factory() - native_runtime(facade) - inputs: Final = [f"alpha {uuid4().hex}", f"beta {uuid4().hex}"] - result: Final = EmbeddingResponse( - model="text-embedding-3-small", - data=[ - {"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}, - {"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}, - ], - ) - await facade.async_add_cache_pipeline(result, model="text-embedding-3-small", input=inputs) - - keys: Final = [facade.get_cache_key(model="text-embedding-3-small", input=text) for text in inputs] - assert len(set(keys)) == len(inputs) - for text, expected in zip(inputs, ([0.1, 0.2], [0.3, 0.4]), strict=True): - cached = await facade.async_get_cache(model="text-embedding-3-small", input=text) - assert isinstance(cached, dict) - assert cached["embedding"] == expected - assert await facade.async_get_cache(model="text-embedding-3-small", input=inputs) is None - - -def redis_facade(redis_url: str, **settings: object) -> Cache: - parsed: Final = urlparse(redis_url) - return Cache(type=LiteLLMCacheType.REDIS, host=parsed.hostname, port=str(parsed.port), **settings) - - -@pytest.mark.parametrize( - ("settings", "message"), - [ - pytest.param({"max_connections": 10}, "max_connections requires Python", id="pool-size"), - pytest.param({"socket_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="socket-timeout"), - pytest.param( - {"socket_connect_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="connect-timeout" - ), - pytest.param({"socket_keepalive": True}, "does not support socket_keepalive", id="keepalive"), - pytest.param({"health_check_interval": 5}, "does not support health_check_interval", id="health-check"), - pytest.param({"client_name": "litellm"}, "does not support client_name", id="client-name"), - pytest.param({"ssl": True}, "ssl_check_hostname=false require Python", id="tls-default-hostname-check"), - pytest.param({"ssl": True, "ssl_cert_reqs": "none"}, "ssl_cert_reqs=none", id="tls-without-verification"), - pytest.param( - {"ssl": True, "ssl_check_hostname": True, "ssl_ca_certs": "/ca.pem"}, - "does not support ssl_ca_certs", - id="tls-custom-ca", - ), - pytest.param( - {"ssl": True, "ssl_check_hostname": True, "ssl_certfile": "/client.pem", "ssl_keyfile": "/client.key"}, - "does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or ssl_keyfile", - id="tls-client-certificate", - ), - ], -) -def test_redis_settings_the_native_client_cannot_honor_decline( - redis_url: str, monkeypatch: pytest.MonkeyPatch, settings: dict[str, object], message: str -) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - with pytest.raises(RuntimeError, match=f"declined the cache: native Redis.*{message}"): - redis_facade(redis_url, **settings) - - -def test_redis_verified_tls_activates_natively(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - native_runtime(redis_facade(redis_url, ssl=True, ssl_check_hostname=True)) - - -async def test_redis_flush_size_buffers_native_facade_writes(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - facade: Final = redis_facade(redis_url, redis_flush_size=2, namespace="team") - native_runtime(facade) - client: Final = redis.Redis.from_url(redis_url) - first: Final = completion_kwargs("first") - await facade.async_add_cache({"value": 1}, **first) - first_key: Final = facade.get_cache_key(**first) - assert first_key.startswith("team:") - assert client.get(first_key) is None - second: Final = completion_kwargs("second") - await facade.async_add_cache({"value": 2}, **second) - assert client.get(first_key) is not None - assert client.get(facade.get_cache_key(**second)) is not None - client.close() - - -@pytest.mark.parametrize( - ("backend", "settings", "message"), - [ - pytest.param( - LiteLLMCacheType.VALKEY_SEMANTIC, - {"redis_url": "rediss://127.0.0.1:6390/0", "similarity_threshold": 0.8}, - "native Valkey semantic cache does not support TLS connections", - id="valkey-tls", - ), - pytest.param( - LiteLLMCacheType.VALKEY_SEMANTIC, - {"redis_url": "redis://127.0.0.1:6390/0?socket_timeout=1", "similarity_threshold": 0.8}, - "native Redis uses fixed socket timeouts; socket_timeout and socket_connect_timeout require Python", - id="valkey-socket-timeout", - ), - pytest.param( - LiteLLMCacheType.REDIS_SEMANTIC, - {"redis_url": "rediss://127.0.0.1:6380", "similarity_threshold": 0.8}, - "native Redis semantic cache does not support TLS or query options in redis_url", - id="redis-semantic-tls", - ), - pytest.param( - LiteLLMCacheType.REDIS_SEMANTIC, - {"redis_url": "redis://127.0.0.1:6379?socket_timeout=1", "similarity_threshold": 0.8}, - "native Redis semantic cache does not support TLS or query options in redis_url", - id="redis-semantic-query", - ), - ], -) -def test_semantic_settings_the_native_client_cannot_honor_decline( - monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType, settings: dict[str, object], message: str -) -> None: - require_rust(monkeypatch, backend) - with pytest.raises(RuntimeError, match=f"declined the cache: {message}"): - Cache(type=backend, **settings) - - -def test_rust_with_fallback_keeps_python_when_the_native_client_declines( - redis_url: str, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.setattr( - catalog, - "RULES", - (CacheRule(Rollout.RUST_OPT_OUT, backends=frozenset({LiteLLMCacheType.REDIS})),), - ) - assert redis_facade(redis_url, socket_timeout=1.0)._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor - - -def test_qdrant_semantic_rust_required_rule_activates_natively( - qdrant_url: str, fake_embedding_endpoint: str, monkeypatch: pytest.MonkeyPatch -) -> None: - del fake_embedding_endpoint - require_rust(monkeypatch, LiteLLMCacheType.QDRANT_SEMANTIC) - facade: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") - native_runtime(facade) - kwargs: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "qdrant activation"}]} - facade.add_cache({"answer": "qdrant"}, **kwargs) - assert facade.get_cache(**kwargs) == {"answer": "qdrant"} - - -async def test_redis_semantic_rust_required_rule_activates_natively( - redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding, monkeypatch: pytest.MonkeyPatch -) -> None: - del semantic_embedding - url, index = redis_stack - require_rust(monkeypatch, LiteLLMCacheType.REDIS_SEMANTIC) - facade: Final = Cache( - type=LiteLLMCacheType.REDIS_SEMANTIC, - redis_url=url, - similarity_threshold=0.8, - redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, - redis_semantic_cache_index_name=index, - ) - native_runtime(facade) - kwargs: Final = {"model": "gpt-4o", "messages": semantic_messages("name a primary color")} - await facade.async_add_cache({"answer": "blue"}, **kwargs) - assert await facade.async_get_cache(**kwargs) == {"answer": "blue"} - - -async def test_azure_blob_rust_required_rule_activates_natively(monkeypatch: pytest.MonkeyPatch) -> None: - account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") - if account_url is None: - pytest.skip( - "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" - ) - require_rust(monkeypatch, LiteLLMCacheType.AZURE_BLOB) - facade: Final = Cache( - type=LiteLLMCacheType.AZURE_BLOB, - azure_account_url=account_url, - azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", - ) - backend: Final = facade.cache - assert isinstance(backend, AzureBlobCache) - try: - native_runtime(facade) - kwargs: Final = completion_kwargs("azure") - await facade.async_add_cache({"answer": "azure"}, **kwargs) - assert await facade.async_get_cache(**kwargs) == {"answer": "azure"} - assert backend.get_cache(facade.get_cache_key(**kwargs))["response"] == {"answer": "azure"} - finally: - backend.container_client.delete_container() - await backend.disconnect() - - -class _SemanticHit: - """A native semantic runtime that answers every lookup with one cached response.""" - - kind: Final = "native" - - def lookup_semantic(self, request: object) -> tuple[object, float | None]: - return {"answer": 42}, 0.97 - - async def async_lookup_semantic(self, request: object) -> tuple[object, float | None]: - return {"answer": 42}, 0.97 - - -@pytest.mark.parametrize("semantic_type", [LiteLLMCacheType.QDRANT_SEMANTIC, LiteLLMCacheType.REDIS_SEMANTIC]) -@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) -def test_native_semantic_hit_stamps_similarity_on_request_metadata( - semantic_type: LiteLLMCacheType, use_async: bool -) -> None: - """Python semantic backends write `metadata["semantic-similarity"]` on every lookup, and the - facade copies it to the caller's metadata; the native path must report it the same way.""" - facade: Final = Cache() - facade.type = semantic_type - facade._native_cache = ResponseCacheRuntime(cast(NativeResponseCacheRuntime, _SemanticHit())) # pyright: ignore[reportPrivateUsage] # the native path under test has no public setter - metadata: Final[dict[str, object]] = {} - kwargs: Final = { - "cache_key": "semantic-key", - "messages": [{"role": "user", "content": "hello"}], - "metadata": metadata, - } - - result: Final = asyncio.run(facade.async_get_cache(**kwargs)) if use_async else facade.get_cache(**kwargs) - - assert result == {"answer": 42} - assert metadata["semantic-similarity"] == 0.97 diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py deleted file mode 100644 index 2fbf9817a53..00000000000 --- a/tests/test_litellm_rust/test_ocr.py +++ /dev/null @@ -1,134 +0,0 @@ -import json -import threading -from collections.abc import Generator -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from io import BytesIO -from typing import Final - -import pytest - -import litellm - -pytestmark = pytest.mark.requires_rust_extension - - -@pytest.fixture -def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[dict[str, object]]]]: - requests: Final[list[dict[str, object]]] = [] - - class Handler(BaseHTTPRequestHandler): - def do_POST(self) -> None: - requests.append( - { - "headers": {name.lower(): value for name, value in self.headers.items()}, - "body": json.loads(self.rfile.read(int(self.headers["Content-Length"]))), - } - ) - if self.headers.get("x-test-stall") == "true": - self.connection.settimeout(2) - try: - self.rfile.read(1) - except TimeoutError: - pass - return - if self.headers.get("User-Agent", "").startswith("python-httpx"): - self.send_response(418) - self.end_headers() - return - status = int(self.headers.get("x-test-status", "200")) - if status != 200: - body = b'{"error":"provider unavailable"}' - self.send_response(status) - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - return - response: Final = json.dumps( - { - "pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}], - "model": "mistral-ocr-latest", - "usage_info": {"pages_processed": 1, "doc_size_bytes": 3}, - } - ).encode() - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(response))) - self.end_headers() - self.wfile.write(response) - - def log_message(self, format: str, *args: object) -> None: - pass - - server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler) - thread: Final = threading.Thread(target=lambda: server.serve_forever(poll_interval=0.01), daemon=True) - thread.start() - try: - yield server, requests - finally: - server.shutdown() - server.server_close() - thread.join() - - -def test_native_lifecycle_core_encodes_python_file_input(ocr_server): - server, requests = ocr_server - litellm.rust(True) - response = litellm.ocr( - model="mistral/mistral-ocr-latest", - document={"type": "file", "file": BytesIO(b"abc"), "mime_type": "image/png"}, - api_key="test-key", - api_base=f"http://127.0.0.1:{server.server_port}", - opaque_extension=object(), - ) - assert response.pages[0].markdown == "native OCR response" - assert requests[0]["body"]["document"] == {"type": "image_url", "image_url": "data:image/png;base64,YWJj"} - assert "opaque_extension" not in requests[0]["body"] - - -@pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.asyncio -async def test_native_ocr_failures_do_not_retry_on_python(ocr_server, asynchronous): - server, requests = ocr_server - arguments = { - "model": "mistral-ocr-latest", - "custom_llm_provider": "mistral", - "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "api_key": "test-key", - "api_base": f"http://127.0.0.1:{server.server_port}", - "extra_headers": {"x-test-status": "503"}, - "num_retries": 0, - } - litellm.rust(True) - with pytest.raises(litellm.ServiceUnavailableError) as caught: - await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments) - assert caught.value.status_code == 503 - assert len(requests) == 1 - assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") - - -@pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.asyncio -async def test_native_ocr_enforces_request_deadline_without_fallback(ocr_server, asynchronous): - import asyncio - import time - - server, requests = ocr_server - litellm.rust(True) - arguments = { - "model": "mistral/mistral-ocr-latest", - "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "api_key": "test-key", - "api_base": f"http://127.0.0.1:{server.server_port}", - "extra_headers": {"x-test-stall": "true"}, - "timeout": 0.1, - "num_retries": 0, - } - started = time.monotonic() - with pytest.raises(litellm.Timeout): - await asyncio.wait_for( - litellm.aocr(**arguments) if asynchronous else asyncio.to_thread(litellm.ocr, **arguments), - timeout=3, - ) - assert 0.09 <= time.monotonic() - started < 3 - assert len(requests) == 1 - assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") diff --git a/tests/test_litellm/a2a_protocol/__init__.py b/tests/test_litellm_rust/tokenizer/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/__init__.py rename to tests/test_litellm_rust/tokenizer/__init__.py diff --git a/tests/test_litellm_rust/test_tokenizer.py b/tests/test_litellm_rust/tokenizer/test_fast_count.py similarity index 51% rename from tests/test_litellm_rust/test_tokenizer.py rename to tests/test_litellm_rust/tokenizer/test_fast_count.py index 98d5259b652..2902b79dca8 100644 --- a/tests/test_litellm_rust/test_tokenizer.py +++ b/tests/test_litellm_rust/tokenizer/test_fast_count.py @@ -12,67 +12,6 @@ from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOK pytestmark = pytest.mark.requires_rust_extension -def test_tiktoken_codec_round_trips_and_counts() -> None: - tokenizer: Final = _native.Tokenizer.from_tiktoken("cl100k_base") - encoded: Final = tokenizer.encode("hello world") - - assert tokenizer.name == "cl100k_base" - assert tokenizer.count("hello world") == len(encoded) - assert tokenizer.decode(encoded) == "hello world" - - -def test_huggingface_codec_skips_special_tokens() -> None: - tokenizer: Final = _native.Tokenizer.from_json(claude_json_str) - encoded: Final = tokenizer.encode("hello") - - assert "" in tokenizer.decode(encoded, skip_special_tokens=False) - assert tokenizer.decode(encoded, skip_special_tokens=True) == "hello" - - -def test_tiktoken_codec_keeps_the_requested_encoding_name() -> None: - assert _native.Tokenizer.from_tiktoken("gpt2").name == "gpt2" - assert _native.Tokenizer.from_tiktoken("r50k_base").name == "r50k_base" - assert _native.Tokenizer.from_tiktoken("gpt2").encode("hi") == _native.Tokenizer.from_tiktoken("r50k_base").encode( - "hi" - ) - - -def test_tiktoken_codec_exposes_its_vocabulary() -> None: - reference: Final = tiktoken.get_encoding("cl100k_base") - tokenizer: Final = _native.Tokenizer.from_tiktoken("cl100k_base") - - assert tokenizer.special_tokens() == reference._special_tokens - assert tokenizer.max_token_value() == reference.max_token_value - assert tokenizer.token_byte_values() == reference.token_byte_values() - assert tokenizer.encode_single_token(b"hello") == reference.encode_single_token("hello") - assert tokenizer.is_special_token(reference.eot_token) and not tokenizer.is_special_token(0) - with pytest.raises(KeyError): - tokenizer.encode_single_token(b"<|not-a-token|>") - - -def test_huggingface_codec_rejects_tiktoken_only_calls() -> None: - tokenizer: Final = _native.Tokenizer.from_json(claude_json_str) - with pytest.raises(ValueError, match="requires a tiktoken encoding"): - tokenizer.token_byte_values() - with pytest.raises(ValueError, match="requires a Hugging Face tokenizer"): - _native.Tokenizer.from_tiktoken("cl100k_base").get_vocab() - - -def test_unknown_tiktoken_encoding_raises_value_error() -> None: - with pytest.raises(ValueError, match="unsupported tokenizer"): - _native.Tokenizer.from_tiktoken("unknown-encoding") - - -def test_tiktoken_codec_decodes_truncated_unicode_like_python() -> None: - reference: Final = tiktoken.get_encoding("cl100k_base") - tokenizer: Final = _native.Tokenizer.from_tiktoken(reference.name) - encoded: Final = reference.encode("🙂漢字") - - assert tuple(tokenizer.decode(encoded[:end]) for end in range(1, len(encoded) + 1)) == tuple( - reference.decode(encoded[:end]) for end in range(1, len(encoded) + 1) - ) - - FAST_TEXTS: Final = ( "", "hello world <|endoftext|>", diff --git a/tests/test_litellm_rust/tokenizer/test_huggingface.py b/tests/test_litellm_rust/tokenizer/test_huggingface.py new file mode 100644 index 00000000000..05c5989c676 --- /dev/null +++ b/tests/test_litellm_rust/tokenizer/test_huggingface.py @@ -0,0 +1,24 @@ +from typing import Final + +import pytest + +from litellm.rust_bridge import _native +from litellm.utils import claude_json_str + +pytestmark = pytest.mark.requires_rust_extension + + +def test_huggingface_codec_skips_special_tokens() -> None: + tokenizer: Final = _native.Tokenizer.from_json(claude_json_str) + encoded: Final = tokenizer.encode("hello") + + assert "" in tokenizer.decode(encoded, skip_special_tokens=False) + assert tokenizer.decode(encoded, skip_special_tokens=True) == "hello" + + +def test_huggingface_codec_rejects_tiktoken_only_calls() -> None: + tokenizer: Final = _native.Tokenizer.from_json(claude_json_str) + with pytest.raises(ValueError, match="requires a tiktoken encoding"): + tokenizer.token_byte_values() + with pytest.raises(ValueError, match="requires a Hugging Face tokenizer"): + _native.Tokenizer.from_tiktoken("cl100k_base").get_vocab() diff --git a/tests/test_litellm_rust/tokenizer/test_tiktoken.py b/tests/test_litellm_rust/tokenizer/test_tiktoken.py new file mode 100644 index 00000000000..c204d7eaf2f --- /dev/null +++ b/tests/test_litellm_rust/tokenizer/test_tiktoken.py @@ -0,0 +1,53 @@ +from typing import Final + +import pytest +import tiktoken + +from litellm.rust_bridge import _native + +pytestmark = pytest.mark.requires_rust_extension + + +def test_tiktoken_codec_round_trips_and_counts() -> None: + tokenizer: Final = _native.Tokenizer.from_tiktoken("cl100k_base") + encoded: Final = tokenizer.encode("hello world") + + assert tokenizer.name == "cl100k_base" + assert tokenizer.count("hello world") == len(encoded) + assert tokenizer.decode(encoded) == "hello world" + + +def test_tiktoken_codec_keeps_the_requested_encoding_name() -> None: + assert _native.Tokenizer.from_tiktoken("gpt2").name == "gpt2" + assert _native.Tokenizer.from_tiktoken("r50k_base").name == "r50k_base" + assert _native.Tokenizer.from_tiktoken("gpt2").encode("hi") == _native.Tokenizer.from_tiktoken("r50k_base").encode( + "hi" + ) + + +def test_tiktoken_codec_exposes_its_vocabulary() -> None: + reference: Final = tiktoken.get_encoding("cl100k_base") + tokenizer: Final = _native.Tokenizer.from_tiktoken("cl100k_base") + + assert tokenizer.special_tokens() == reference._special_tokens + assert tokenizer.max_token_value() == reference.max_token_value + assert tokenizer.token_byte_values() == reference.token_byte_values() + assert tokenizer.encode_single_token(b"hello") == reference.encode_single_token("hello") + assert tokenizer.is_special_token(reference.eot_token) and not tokenizer.is_special_token(0) + with pytest.raises(KeyError): + tokenizer.encode_single_token(b"<|not-a-token|>") + + +def test_unknown_tiktoken_encoding_raises_value_error() -> None: + with pytest.raises(ValueError, match="unsupported tokenizer"): + _native.Tokenizer.from_tiktoken("unknown-encoding") + + +def test_tiktoken_codec_decodes_truncated_unicode_like_python() -> None: + reference: Final = tiktoken.get_encoding("cl100k_base") + tokenizer: Final = _native.Tokenizer.from_tiktoken(reference.name) + encoded: Final = reference.encode("🙂漢字") + + assert tuple(tokenizer.decode(encoded[:end]) for end in range(1, len(encoded) + 1)) == tuple( + reference.decode(encoded[:end]) for end in range(1, len(encoded) + 1) + ) 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/providers/__init__.py b/tests/unit/completion_extras/litellm_responses_transformation/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/__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..653b2c9914a 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -1,7 +1,11 @@ +import asyncio +import importlib import os -from collections.abc import Iterator +from collections.abc import Coroutine, Iterator +from pathlib import Path from typing import Final +import boto3 import pytest from pytest_socket import enable_socket, socket_allow_hosts @@ -10,6 +14,14 @@ 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.litellm_core_utils.prompt_templates import ( # noqa: E402 # same import-time dependency + image_handling as image_handling_module, +) +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 +32,63 @@ 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") def _allow_loopback_only() -> None: @@ -29,11 +98,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 +215,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) @@ -68,4 +243,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/bedrock_agentcore/__init__.py b/tests/unit/containers/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/__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/pydantic_ai_agents/__init__.py b/tests/unit/embeddings/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/pydantic_ai_agents/__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/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 99% rename from tests/test_litellm/experimental_mcp_client/test_mcp_client.py rename to tests/unit/experimental_mcp_client/test_mcp_client.py index ae30b086c6e..368e34c455d 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -2762,6 +2762,7 @@ async def test_cancellation_delivers_termination_over_tcp( cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool ) -> None: started: Final = asyncio.Event() + scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future() terminations: Final[list[bytes]] = [] starts: Final[list[bytes]] = [] stop: Final = asyncio.Event() @@ -2858,13 +2859,16 @@ async def test_cancellation_delivers_termination_over_tcp( async def invoke(): if cancel_mode == "scope": - with anyio.fail_after(0.2): + with anyio.fail_after(None) as scope: + scope_ready.set_result(scope) return await calls() return await calls() try: task: Final = asyncio.create_task(invoke()) await asyncio.wait_for(started.wait(), 3) + if cancel_mode == "scope": + (await scope_ready).deadline = anyio.current_time() + 0.2 if cancel_mode == "task": task.cancel() expected_error: Final = ( 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/unit/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py b/tests/unit/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py index 3846a94c9fe..5d97beeb3fc 100644 --- a/tests/unit/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py +++ b/tests/unit/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py @@ -1,5 +1,6 @@ import asyncio import base64 +import copy import json import uuid from types import SimpleNamespace @@ -1010,3 +1011,102 @@ def test_bedrock_chat_invoke_eager_input_streaming_beta_not_duplicated_with_clie ) assert result["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] + + +def _mid_conversation_system_conversation() -> list[dict]: + return [ + {"role": "system", "content": [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}]}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + {"role": "system", "content": "Answer with exactly one word."}, + {"role": "assistant", "content": "Second answer"}, + {"role": "user", "content": "Third question"}, + ] + + +def test_chat_unflagged_model_converts_mid_conversation_system_instead_of_hoisting(local_model_cost_map): + """A hoisted reminder rewrites the top-level system block and invalidates the + prompt cache for the whole conversation (#36559).""" + result = AmazonAnthropicClaudeConfig().transform_request( + model="invoke/us.anthropic.claude-opus-4-7", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "assistant", "user"] + texts = [b["text"] for b in result["messages"][2]["content"] if b.get("type") == "text"] + assert texts[0] == "Second question" + assert texts[-1] == "Answer with exactly one word." + + +def test_chat_flagged_model_keeps_mid_conversation_system_role_in_place(local_model_cost_map): + result = AmazonAnthropicClaudeConfig().transform_request( + model="invoke/us.anthropic.claude-opus-4-8", + messages=_mid_conversation_system_conversation(), + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert result["system"] == [{"type": "text", "text": "You are terse.", "cache_control": {"type": "ephemeral"}}] + assert [m["role"] for m in result["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] + assert result["messages"][3] == { + "role": "system", + "content": [{"type": "text", "text": "Answer with exactly one word."}], + } + + +def _thinking_reply(text: str) -> dict: + return { + "role": "assistant", + "content": text, + "thinking_blocks": [{"type": "thinking", "thinking": "Working it out.", "signature": f"sig-{text}"}], + } + + +def _preserved_thinking_turns(reminder_after_user: bool) -> tuple[list[dict], list[dict], list[dict]]: + turn_n = [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "First question"}] + reminder = {"role": "system", "content": "Answer with exactly one word."} + second_question = {"role": "user", "content": "Second question"} + second_turn = [second_question, reminder] if reminder_after_user else [reminder, second_question] + turn_n_plus_one = [*turn_n, _thinking_reply("First answer"), *second_turn] + turn_n_plus_two = [*turn_n_plus_one, _thinking_reply("Second answer"), {"role": "user", "content": "Third question"}] + return turn_n, turn_n_plus_one, turn_n_plus_two + + +def _replayed_prefix(request: dict, message_count: int) -> str: + replayed = { + "system": request.get("system"), + "tools": request.get("tools"), + "messages": request["messages"][:message_count], + } + return json.dumps(replayed, sort_keys=True) + + +def _assert_prefix_stable(requests: list[dict]) -> None: + for earlier, later in zip(requests, requests[1:]): + count = len(earlier["messages"]) + assert _replayed_prefix(later, count) == _replayed_prefix(earlier, count) + + +@pytest.mark.parametrize("reminder_after_user", [True, False]) +def test_chat_flagged_model_replays_a_byte_identical_prefix_around_a_mid_conversation_reminder( + local_model_cost_map, reminder_after_user +): + """Preserved thinking binds each signed block to the request prefix it was created + under (``system``, ``tools`` and the earlier messages), so turn N's transformed + request must be a byte-identical prefix of turn N+1's or the block is dropped.""" + requests = [ + AmazonAnthropicClaudeConfig().transform_request( + model="invoke/us.anthropic.claude-fable-5-1", messages=copy.deepcopy(turn), optional_params={}, litellm_params={}, headers={} + ) + for turn in _preserved_thinking_turns(reminder_after_user) + ] + + _assert_prefix_stable(requests) + assert [m["role"] for m in requests[1]["messages"]] == ["user", "assistant", "user", "system"] + assert [m["role"] for m in requests[2]["messages"]] == ["user", "assistant", "user", "system", "assistant", "user"] diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py index 3bca51ec6b3..05e4e36edd7 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py @@ -11,6 +11,39 @@ from litellm.llms.vertex_ai.vertex_ai_partner_models.llama3.transformation impor ) +OPENAI_PLATFORM_PARAMS = ( + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", +) + +SELF_DEPLOYED_ENDPOINT_MODELS = ( + "gemma/gemma-2-2b-it", + "vertex_ai/gemma/gemma-2-2b-it", + "openai/mg-endpoint-lit8592", + "vertex_ai/openai/mg-endpoint-lit8592", + "openai/5464397967697903616", +) + +MAAS_MODELS = ( + "meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "moonshotai/kimi-k2-thinking-maas", + "qwen/qwen3-next-80b-a3b-instruct-maas", + "google/gemma-4-26b-a4b-it-maas", + "xai/grok-4.1-fast-non-reasoning", + "openai/xai/grok-4.1-fast-reasoning", + "1984786713414729728", + "llama3", +) + + class TestVertexAILlama3Config: def test_transform_choices(self): """ @@ -56,6 +89,52 @@ class TestVertexAILlama3Config: assert response[0].message.tool_calls is not None assert response[0].finish_reason == "tool_calls" + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_omits_platform_params_for_self_deployed_endpoints( + self, model: str, param: str + ): + assert param not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", MAAS_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_keeps_platform_params_for_maas_models(self, model: str, param: str): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + def test_get_supported_openai_params_never_lists_max_retries(self, model: str): + assert "max_retries" not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + @pytest.mark.parametrize( + "param", + ["max_completion_tokens", "tools", "tool_choice", "response_format", "seed", "logprobs", "parallel_tool_calls"], + ) + def test_get_supported_openai_params_keeps_params_every_vertex_openai_endpoint_accepts( + self, model: str, param: str + ): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + def test_map_openai_params_drops_prompt_cache_key_for_self_deployed_endpoints(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"max_tokens": 10} + + @pytest.mark.parametrize("model", MAAS_MODELS) + def test_map_openai_params_forwards_prompt_cache_key_for_maas_models(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"prompt_cache_key": "session-lit8592", "max_tokens": 10} + class TestVertexAILlama3StreamingHandler: def test_first_chunk_has_role_assistant_when_missing(self): diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 5f74f0f602f..e5ca31833ce 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -273,6 +273,165 @@ class TestVertexGemmaCompletion: # Verify the error message contains the original error assert "missing 'predictions' field" in str(exc_info.value) + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True]) + async def test_acompletion_surfaces_container_error_object_as_its_own_status_and_message(self, stream): + """ + A serving container can reject the request with its own OpenAI-shaped error object, + which Vertex still wraps in an HTTP 200 :predict response. The container's status and + message must reach the caller instead of a 500 "no 'choices'". + """ + from litellm.exceptions import BadRequestError + + container_message = '"auto" tool choice requires --enable-auto-tool-choice and --tool-call-parser to be set' + vertex_response = { + "deployedModelId": "123", + "predictions": { + "code": 400, + "message": container_message, + "object": "error", + "param": None, + "type": "BadRequestError", + }, + } + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = vertex_response + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + with pytest.raises(BadRequestError) as exc_info: + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "What is the weather in Paris?"}], + tools=[{"type": "function", "function": {"name": "get_weather", "parameters": {}}}], + stream=stream, + api_base="https://test.prediction.vertexai.goog/v1/projects/test/locations/us-central1/endpoints/123:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + + assert exc_info.value.status_code == 400 + assert container_message in str(exc_info.value) + assert "no 'choices'" not in str(exc_info.value) + + @pytest.mark.asyncio + async def test_acompletion_keeps_container_error_status_beyond_400(self): + """The container's status is forwarded as is, not collapsed to 400.""" + from litellm.exceptions import RateLimitError + + vertex_response = { + "deployedModelId": "123", + "predictions": {"code": 429, "message": "engine overloaded", "object": "error", "type": "RateLimitError"}, + } + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = vertex_response + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + with pytest.raises(RateLimitError) as exc_info: + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + api_base="https://test.prediction.vertexai.goog/v1/projects/test/locations/us-central1/endpoints/123:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + + assert exc_info.value.status_code == 429 + assert "engine overloaded" in str(exc_info.value) + + @pytest.mark.asyncio + async def test_acompletion_error_object_without_http_status_keeps_generic_handling(self): + """An error-shaped body whose code is not an HTTP error status is not trusted as one.""" + from litellm.exceptions import APIError + + vertex_response = { + "deployedModelId": "123", + "predictions": {"code": 0, "message": "unknown failure", "object": "error"}, + } + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = vertex_response + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + with pytest.raises(APIError) as exc_info: + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + api_base="https://test.prediction.vertexai.goog/v1/projects/test/locations/us-central1/endpoints/123:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + + assert exc_info.value.status_code == 500 + + def test_sync_completion_surfaces_container_error_object_as_its_own_status_and_message(self): + """The synchronous path unwraps the same container error object.""" + from litellm.exceptions import BadRequestError + + container_message = '"auto" tool choice requires --enable-auto-tool-choice and --tool-call-parser to be set' + vertex_response = { + "deployedModelId": "123", + "predictions": {"code": 400, "message": container_message, "object": "error", "type": "BadRequestError"}, + } + + with ( + patch("litellm.llms.vertex_ai.vertex_gemma_models.transformation._get_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = vertex_response + mock_client.post = Mock(return_value=mock_response) + mock_get_client.return_value = mock_client + + with pytest.raises(BadRequestError) as exc_info: + litellm.completion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "What is the weather in Paris?"}], + tools=[{"type": "function", "function": {"name": "get_weather", "parameters": {}}}], + api_base="https://test.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + assert exc_info.value.status_code == 400 + assert container_message in str(exc_info.value) + @pytest.mark.asyncio async def test_acompletion_fake_streaming(self): """ @@ -535,6 +694,89 @@ class TestVertexGemmaCompletion: assert instance["@requestFormat"] == "chatCompletions" assert "messages" in instance + @pytest.mark.parametrize( + "param", + [ + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", + "max_retries", + ], + ) + def test_get_supported_openai_params_omits_params_the_predict_endpoint_rejects(self, param: str): + from litellm.llms.vertex_ai.vertex_gemma_models.transformation import ( + VertexGemmaConfig, + ) + + assert param not in VertexGemmaConfig().get_supported_openai_params(model="gemma-2-2b-it") + + @pytest.mark.asyncio + async def test_acompletion_drops_prompt_cache_key_when_drop_params_is_set(self): + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = _make_gemma_vertex_response() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + service_tier="default", + max_completion_tokens=16, + drop_params=True, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + instance = mock_client.post.call_args.kwargs["json"]["instances"][0] + assert "prompt_cache_key" not in instance + assert "service_tier" not in instance + assert instance["max_tokens"] == 16 + assert instance["messages"] == [{"role": "user", "content": "Test"}] + + @pytest.mark.asyncio + async def test_acompletion_rejects_prompt_cache_key_before_calling_vertex(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "drop_params", False) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_client.post = AsyncMock() + mock_get_client.return_value = mock_client + + with pytest.raises(litellm.UnsupportedParamsError, match="prompt_cache_key"): + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + drop_params=False, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + mock_client.post.assert_not_called() + def test_transform_request_strips_context_management(self): """ Direct unit test for VertexGemmaConfig.transform_request: verify that 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/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py index efe41e1da9a..f194e43c74a 100644 --- a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py @@ -19,6 +19,7 @@ from litellm.proxy.hooks.model_max_budget_limiter import ( resolve_model_budget, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import BudgetConfig as GenericBudgetInfo @@ -487,6 +488,23 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config mock_push.assert_awaited_once() +@pytest.mark.asyncio +async def test_model_budget_limiter_initializes_redis_increment_queue_lock(): + dual_cache = DualCache() + limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + spend_key = "virtual_key_spend:test-key:gpt-4:1d" + + await limiter._increment_spend_in_current_window( + spend_key=spend_key, response_cost=0.01, ttl=86400 + ) + + assert limiter.redis_increment_operation_queue == [ + RedisPipelineIncrementOperation( + key=spend_key, increment_value=0.01, ttl=86400 + ) + ] + + @pytest.mark.asyncio async def test_get_fallback_model_within_budget_returns_none_without_fallbacks( budget_limiter, 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/test_litellm/files/__init__.py b/tests/unit/rerank_api/__init__.py similarity index 100% rename from tests/test_litellm/files/__init__.py rename to tests/unit/rerank_api/__init__.py 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 82% rename from tests/test_litellm/test_check_type_discipline.py rename to tests/unit/test_check_type_discipline.py index b0c8d2d5d56..aee73825d63 100644 --- a/tests/test_litellm/test_check_type_discipline.py +++ b/tests/unit/test_check_type_discipline.py @@ -685,6 +685,169 @@ def test_writable_ok_without_reason_is_lit005_and_does_not_suppress(tmp_path): assert "LIT012" in codes +# --------------------------------------------------------------------------- # +# Stacked comprehension clauses (LIT014) +# --------------------------------------------------------------------------- # + + +def test_two_for_clauses_are_flagged(tmp_path: Path): + assert "LIT014" in _codes(tmp_path, "y = [x for a in xs for x in a]\n") + + +def test_two_ifs_on_one_generator_are_flagged(tmp_path: Path): + assert "LIT014" in _codes(tmp_path, "y = [x for x in xs if x if x > 1]\n") + + +def test_one_if_on_each_of_two_generators_is_flagged(tmp_path: Path): + assert "LIT014" in _codes(tmp_path, "y = [x for a in xs if a for x in a if x]\n") + + +def test_one_for_and_one_if_is_clean(tmp_path: Path): + assert "LIT014" not in _codes(tmp_path, "y = tuple(x for x in xs if x)\n") + + +def test_dict_set_and_generator_two_fors_are_each_flagged(tmp_path: Path): + assert "LIT014" in _codes(tmp_path, "d = {k: v for a in xs for k, v in a}\n") + assert "LIT014" in _codes(tmp_path, "s = {x for a in xs for x in a}\n") + assert "LIT014" in _codes(tmp_path, "g = (x for a in xs for x in a)\n") + + +def test_nested_comprehension_in_element_is_judged_separately(tmp_path: Path): + assert "LIT014" not in _codes(tmp_path, "y = [[v for v in a] for a in xs]\n") + + +def test_comprehension_ok_with_reason_suppresses_lit014(tmp_path: Path): + codes = _codes( + tmp_path, + "y = [x for a in xs for x in a] # comprehension-ok: flattens a stream of pairs, hot path\n", + ) + assert "LIT014" not in codes + + +def test_comprehension_ok_on_any_spanned_line_suppresses_lit014(tmp_path: Path): + src = ( + "y = [\n" + " x for a in xs\n" + " for x in a\n" + "] # comprehension-ok: cartesian product is the clearest form\n" + ) + assert "LIT014" not in _codes(tmp_path, src) + + +def test_comprehension_ok_after_the_closing_line_does_not_suppress(tmp_path: Path): + src = ( + "y = [\n" + " x for a in xs\n" + " for x in a\n" + "]\n" + "# comprehension-ok: cartesian product is the clearest form\n" + ) + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert [v.line for v in violations if v.code == "LIT014"] == [1] + assert [v.line for v in violations if v.code == "LIT013"] == [5] + + +def test_comprehension_ok_on_a_compliant_comprehension_is_an_unused_marker(tmp_path: Path): + f = tmp_path / "snippet.py" + f.write_text( + "y = tuple(x for x in xs if x) # comprehension-ok: kept for readability\n", + encoding="utf-8", + ) + violations = checker.check_file(f) + assert [v.line for v in violations if v.code == "LIT013"] == [1] + assert "LIT014" not in [v.code for v in violations] + + +def test_comprehension_ok_without_reason_is_lit005_and_does_not_suppress(tmp_path: Path): + codes = _codes(tmp_path, "y = [x for a in xs for x in a] # comprehension-ok\n") + assert "LIT005" in codes + assert "LIT014" in codes + + +def test_suppression_inside_inner_comprehension_does_not_silence_the_outer(tmp_path: Path): + src = ( + "y = [\n" + " x\n" + " for a in [\n" + " z for i in ys\n" + " for z in i\n" + " ] # comprehension-ok: inner flatten is the clearest form\n" + " for x in a\n" + "]\n" + ) + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert [v.line for v in violations if v.code == "LIT014"] == [1] + assert [v.code for v in violations if v.code == "LIT013"] == [] + + +def test_suppression_on_outer_closing_line_does_not_silence_the_inner(tmp_path: Path): + src = ( + "y = [\n" + " x\n" + " for a in [z for i in ys for z in i]\n" + " for x in a\n" + "] # comprehension-ok: outer flatten is the clearest form\n" + ) + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + flagged = [v for v in checker.check_file(f) if v.code == "LIT014"] + assert [v.line for v in flagged] == [3] + + +def test_equal_span_marker_suppresses_every_violating_comprehension_on_its_line(tmp_path: Path): + src = "y = [x for a in [z for i in ys for z in i] if a if x] # comprehension-ok: inner flatten is fine\n" + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert "LIT014" not in [v.code for v in violations] + assert "LIT013" not in [v.code for v in violations] + + +def test_single_line_outer_with_violating_inner_is_suppressed(tmp_path: Path): + src = "y = [x for a in [z for i in ys for z in i] for x in a] # comprehension-ok: nested flatten is fine\n" + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert "LIT014" not in [v.code for v in violations] + assert "LIT013" not in [v.code for v in violations] + + +def test_one_marker_suppresses_two_violating_sibling_comprehensions_on_its_line(tmp_path: Path): + src = "y = [x for a in xs for x in a] + [x for a in ys for x in a] # comprehension-ok: paired flattens\n" + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert "LIT014" not in [v.code for v in violations] + assert "LIT013" not in [v.code for v in violations] + + +def test_marker_on_a_non_violating_inner_line_suppresses_the_violating_outer(tmp_path: Path): + src = ( + "y = [\n" + " x\n" + " for a in [z for z in ys if z] # comprehension-ok: flatten stays readable\n" + " for x in a\n" + "]\n" + ) + f = tmp_path / "snippet.py" + f.write_text(src, encoding="utf-8") + violations = checker.check_file(f) + assert "LIT014" not in [v.code for v in violations] + assert "LIT013" not in [v.code for v in violations] + + +def test_violation_message_names_the_clause_counts(tmp_path: Path): + f = tmp_path / "snippet.py" + f.write_text("y = [x for a in xs for x in a if x]\n", encoding="utf-8") + messages = [v.message for v in checker.check_file(f) if v.code == "LIT014"] + assert len(messages) == 1 + assert "2 `for` clauses and 1 `if` clause" in messages[0] + + # --------------------------------------------------------------------------- # # Budget integrity: every emittable LIT rule (bar the LIT000 read/parse error) is gated # --------------------------------------------------------------------------- # 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 87% rename from tests/test_litellm/test_cost_calculator.py rename to tests/unit/test_cost_calculator.py index f7d6cfaf079..99dea6366f9 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -309,6 +309,243 @@ def test_realtime_logging_object_does_not_validate_unknown_event_types(): assert len(dumped["results"]) == len(results) +def test_realtime_transcription_honors_deployment_pricing_override(monkeypatch: pytest.MonkeyPatch) -> None: + """A deployment's pricing override must reach transcription events too. + + Transcription is billed separately from response usage inside the same realtime + session, so a deployment registered at zero rates has to zero both. Resolving + transcription against the public ASR model instead billed a zero-rated + deployment for every .completed event. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-zero-rated-asr" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.0, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + public_rate_cost = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert public_rate_cost > 0, "the public ASR rate must be non-zero for this test to mean anything" + + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + ) + assert abs(without_override - public_rate_cost) < 1e-9 + + with_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + assert with_override == 0.0, "the zero-rated deployment must not be billed for transcription" + + +def test_realtime_transcription_partial_override_keeps_unset_rates(monkeypatch: pytest.MonkeyPatch) -> None: + """An override must not blank the rates it does not set. + + A deployment that prices tokens but omits input_cost_per_second would otherwise + bill duration-based transcription at nothing, because the cost helpers read + `.get(key) or 0.0`. Only the fields the operator actually set may win. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-tokens-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + + expected = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert expected > 0, "the public ASR per-second rate must be non-zero for this test to mean anything" + assert cost == pytest.approx(expected, rel=1e-9), ( + "duration must keep the ASR per-second rate the override left unset" + ) + + +@pytest.mark.parametrize( + "label,override,expected_audio_rate,expected_per_second", + [ + ("tokens only", {"input_cost_per_token": 0.0}, 0.0, 0.017 / 60), + ("audio zeroed", {"input_cost_per_audio_token": 0.0}, 0.0, 0.017 / 60), + ("per second only", {"input_cost_per_second": 0.001}, 6e-06, 0.001), + ("empty override", {}, 6e-06, 0.017 / 60), + ("no override", None, 6e-06, 0.017 / 60), + ], +) +def test_transcription_rate_precedence( + monkeypatch: pytest.MonkeyPatch, + label: str, + override: dict[str, float] | None, + expected_audio_rate: float, + expected_per_second: float, +) -> None: + """Rates resolve within one entry before moving to the next, and zero is a real value. + + An override that prices only tokens must apply its own token rate to audio rather + than reaching past itself for the public audio rate, a deliberate zero must win + instead of being treated as unset, and a rate the override never mentions must keep + the base entry's value. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "asr-precedence-base" + deployment_id = "asr-precedence-deployment" + litellm.register_model( + model_cost={ + base_model: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_audio_token": 6e-06, + "input_cost_per_token": 2.5e-06, + "input_cost_per_second": 0.017 / 60, + } + } + ) + if override is not None: + litellm.register_model( + model_cost={deployment_id: {"litellm_provider": "openai", "mode": "audio_transcription", **override}} + ) + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[ + {"type": "transcription_session.created", "session": {"model": base_model}}, + {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}, + ], + custom_llm_provider="openai", + litellm_model_name=base_model, + custom_pricing_model=deployment_id if override is not None else None, + ) + + audio_cost = cost_for({"type": "tokens", "input_token_details": {"audio_tokens": 100}}) + assert audio_cost == pytest.approx(100 * expected_audio_rate, rel=1e-9), f"{label}: audio rate" + + per_second_cost = cost_for({"type": "duration", "seconds": 120.0}) + assert per_second_cost == pytest.approx(120.0 * expected_per_second, rel=1e-9), ( + f"{label}: an override must never blank a rate it does not set" + ) + + +def test_realtime_transcription_per_second_override_keeps_public_token_rates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A per-second override must not zero the token rates ``get_model_info`` synthesizes. + + ``get_model_info`` defaults input_cost_per_token and output_cost_per_token to 0 for entries + that omit them, so a deployment priced only per second looked like it had declared token + rates of 0. Token-shaped transcription then billed nothing instead of falling through to the + public ASR rates, while the per-second rate the operator did set stayed in force. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + asr_model = "gpt-4o-transcribe" + per_second_rate = 0.001 + deployment_id = "deployment-hash-per-second-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_second": per_second_rate, + } + } + ) + + public = litellm.model_cost[asr_model] + session_event = {"type": "transcription_session.created", "session": {"model": asr_model}} + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[session_event, {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}], + custom_llm_provider="openai", + litellm_model_name=asr_model, + custom_pricing_model=deployment_id, + ) + + token_cost = cost_for( + { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + } + ) + expected_token_cost = ( + 400 * public["input_cost_per_audio_token"] + + 12 * public["input_cost_per_token"] + + 30 * public["output_cost_per_token"] + ) + assert expected_token_cost > 0, "the public ASR token rates must be non-zero for this test to mean anything" + assert token_cost == pytest.approx(expected_token_cost, rel=1e-9), ( + "an override that prices only seconds must leave the public token rates in place" + ) + + assert cost_for({"type": "duration", "seconds": 120.0}) == pytest.approx(120.0 * per_second_rate, rel=1e-9) + + def test_realtime_transcription_no_completed_events_is_zero(monkeypatch): """A realtime stream without transcription completed events adds no extra cost.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") @@ -4635,6 +4872,391 @@ def test_gemini_live_native_audio_limits_and_capabilities_match_vendor_model_car assert info["supports_pdf_input"] is False +def test_realtime_honours_deployment_custom_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a deployment's pricing override never reached realtime costing. + + `model_info` overrides are registered under the deployment's own model_id, and + only `_select_model_name_for_cost_calc` knows to look there. The realtime branch + discarded that result and priced by the model the session reported, so a config + that zeroes a realtime deployment was billed at the public rate anyway. Audio is + the bulk of a voice call, so the gap was most of the cost. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-3.1-flash-live-preview" + deployment_key = "deployment-id-for-a-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + "cache_read_input_token_cost": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 10, "output_tokens": 200, "total_tokens": 210}}, + }, + ] + usage = Usage( + prompt_tokens=10, + completion_tokens=200, + total_tokens=210, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=10, cached_tokens=0), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=20, audio_tokens=180), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + ) + expected_paid = ( + 10 * paid["input_cost_per_token"] + + 20 * paid["output_cost_per_token"] + + 180 * paid["output_cost_per_audio_token"] + ) + assert paid_cost == pytest.approx(expected_paid, rel=1e-9) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + custom_pricing_model=deployment_key, + ) + assert zero_rated_cost == 0.0 + + +def test_realtime_honours_a_provider_prefixed_zero_rated_deployment(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: the override arrived provider-prefixed and was read as pricing nothing. + + `_select_model_name_for_cost_calc` hands back `/`, so the name reaching + the pricing guard carries a prefix the raw cost-map lookups cannot strip. The rates resolved + correctly through `get_model_info`, then the guard rejected them as undeclared and the session + billed the public rates. A zero-rated deployment must stay at zero however its name arrives. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-a-prefixed-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + assert zero_rated_cost == 0.0 + + +def test_unpriced_deployment_entry_still_falls_through_to_the_session_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The guard's own purpose must survive: an entry that prices nothing is not an override. + + Deployments are auto-registered under their model_id with no rates at all, and those must + keep billing at the session model's public rates rather than silently costing nothing. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-with-no-declared-rates" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + {key: value for key, value in litellm.model_cost[model].items() if "cost_per" not in key}, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + with_unpriced_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert with_unpriced_override == pytest.approx(without_override, rel=1e-9) + assert with_unpriced_override > 0 + + +def test_realtime_audio_only_override_bills_audio_at_the_deployment_rate( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression: an audio-only pricing override was never selected as the pricing key. + + The deployment-selection guard recognised only text, per-second, per-query and + tiered rates, so a deployment that priced just the audio meters was passed over + and the session kept billing the public rates for the exact tokens it priced. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-an-audio-only-realtime-group" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + "litellm_provider": "vertex_ai", + "mode": "realtime", + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=203, + completion_tokens=58, + total_tokens=261, + prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(audio_tokens=58), + ), + results=[ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 203, "output_tokens": 58, "total_tokens": 261}}, + }, + ], + ) + + public_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert public_cost > 0 + + overridden_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + custom_pricing=True, + router_model_id=deployment_key, + ) + assert overridden_cost == pytest.approx(0.0) + + +def test_realtime_session_falls_back_to_base_model_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a priced base_model was discarded for realtime sessions. + + The resolved base model only reached the realtime cost path when custom pricing + was on, so a session reporting an alias unmapped in the cost map recorded zero + instead of the base model's published price. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + base_model = "gemini-live-2.5-flash-native-audio" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ), + results=[ + { + "type": "session.created", + "session": {"model": "my-voice-alias"}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ], + ) + + aliased_cost = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + base_model=base_model, + ) + base_cost = completion_cost( + completion_response=logging_object, + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert aliased_cost == pytest.approx(base_cost, rel=1e-9) + assert aliased_cost > 0 + + +def test_base_model_does_not_override_transcription_rates(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "gpt-realtime-2" + asr_model = "gpt-4o-transcribe" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage(), + results=[ + { + "type": "session.created", + "session": { + "model": "my-voice-alias", + "audio": {"input": {"transcription": {"model": asr_model}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + }, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + asr_priced = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + realtime_card = litellm.model_cost[base_model] + billed_at_realtime = ( + 400 * realtime_card["input_cost_per_audio_token"] + + 12 * realtime_card["input_cost_per_token"] + + 30 * realtime_card["output_cost_per_audio_token"] + ) + assert billed_at_realtime != pytest.approx(asr_priced, rel=1e-9) + assert with_base_model == pytest.approx(asr_priced, rel=1e-9) + assert with_base_model > 0 + + +def test_realtime_base_model_outranks_the_session_reported_model(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.types.utils import CompletionTokensDetailsWrapper + + session_model = "gpt-realtime-mini" + base_model = "gpt-realtime-2" + + def logging_object_for(session: str) -> LiteLLMRealtimeStreamLoggingObject: + return LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=120, + completion_tokens=60, + total_tokens=180, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=20, audio_tokens=100), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=10, audio_tokens=50), + ), + results=[ + { + "type": "session.created", + "session": {"model": session}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 120, "output_tokens": 60, "total_tokens": 180}}, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + base_priced = completion_cost( + completion_response=logging_object_for(base_model), + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + session_priced = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + assert base_priced != pytest.approx(session_priced, rel=1e-9) + assert with_base_model == pytest.approx(base_priced, rel=1e-9) + + def test_baseten_glm_5_3_fast_is_priced_from_registry(_local_model_cost_map: None) -> None: model: Final = "baseten/zai-org/GLM-5.3-Fast" prompt_tokens: Final = 1000 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/unit/test_read_rc_version.py b/tests/unit/test_read_rc_version.py new file mode 100644 index 00000000000..7b8848916ff --- /dev/null +++ b/tests/unit/test_read_rc_version.py @@ -0,0 +1,49 @@ +"""Tests for .github/scripts/read_rc_version.py.""" + +import importlib.util +import sys +from pathlib import Path +from typing import Final + +import pytest + +_REPO_ROOT: Final = Path(__file__).resolve().parents[2] +_MODULE_PATH: Final = _REPO_ROOT / ".github" / "scripts" / "read_rc_version.py" +_spec: Final = importlib.util.spec_from_file_location("read_rc_version", _MODULE_PATH) +read_rc_version: Final = importlib.util.module_from_spec(_spec) +sys.modules[_spec.name] = read_rc_version +_spec.loader.exec_module(read_rc_version) + + +def _run(tmp_path: Path, version: str, capsys: pytest.CaptureFixture[str]) -> tuple[int, str, str]: + pyproject: Final = tmp_path / "pyproject.toml" + pyproject.write_text(f'[project]\nname = "litellm"\nversion = "{version}"\n', encoding="utf-8") + code: Final = read_rc_version.main(["read_rc_version.py", str(pyproject)]) + captured: Final = capsys.readouterr() + return code, captured.out, captured.err + + +@pytest.mark.parametrize("version", ["1.104.0", "2.0.0", "10.250.0"]) +def test_an_x_y_0_version_is_printed_as_a_github_output_line( + tmp_path: Path, capsys: pytest.CaptureFixture[str], version: str +) -> None: + code, out, err = _run(tmp_path, version, capsys) + assert (code, out, err) == (0, f"version={version}\n", "") + + +@pytest.mark.parametrize("version", ["1.104.1", "1.104.0rc1", "1.104", "v1.104.0", "1.104.0.dev1"]) +def test_a_non_release_version_exits_1_without_printing_a_version( + tmp_path: Path, capsys: pytest.CaptureFixture[str], version: str +) -> None: + code, out, err = _run(tmp_path, version, capsys) + assert code == 1 + assert out == "" + assert err == f"::error::pyproject.toml version {version} is not an X.Y.0 release version\n" + + +def test_the_repo_pyproject_is_read_when_no_path_is_given( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + monkeypatch.chdir(_REPO_ROOT) + assert read_rc_version.main(["read_rc_version.py"]) == 0 + assert capsys.readouterr().out == f"version={read_rc_version.read_version(_REPO_ROOT / 'pyproject.toml')}\n" 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 98% rename from tests/test_litellm/test_router.py rename to tests/unit/test_router/test_router.py index 82122da15dc..80131534183 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -25,6 +25,7 @@ from litellm import Router from litellm.caching.caching import DualCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.exceptions import MidStreamFallbackError +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -49,6 +50,7 @@ from litellm.router import ( from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments +from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -14392,6 +14394,193 @@ async def test_anthropic_messages_fallback_also_catches_raised_midstream_error() assert mock_fallback.await_args.kwargs["e"] is raised_error +_MID_STREAM_OPT_OUT_SHAPES: Final = ( + pytest.param({"disable_fallbacks": True}, id="raw-kwarg"), + pytest.param({"metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="metadata-stamp"), + pytest.param({"litellm_metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="litellm_metadata-stamp"), +) + + +def _mid_stream_opt_out_router() -> Router: + return Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/gpt-5.4", "api_key": "k1"}}, + {"model_name": "fallback", "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "k2"}}, + ], + fallbacks=[{"primary": ["fallback"]}], + ) + + +def _mid_stream_opt_out_primary_error() -> litellm.InternalServerError: + return litellm.InternalServerError(message="primary failed at stream start", llm_provider="openai", model="primary") + + +def _mid_stream_opt_out_trigger(primary_error: Exception) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=str(primary_error), + model="primary", + llm_provider="openai", + original_exception=primary_error, + is_pre_first_chunk=True, + ) + + +class _MidStreamOptOutChatStream(CustomStreamWrapper): + """A chat deployment stream, as the router sees one, that dies before its first chunk.""" + + def __init__(self, error: Exception, model: str = "primary") -> None: + super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()) + self._error: Final = error + + def __aiter__(self): + return self + + async def __anext__(self) -> object: + raise self._error + + def __iter__(self): + return self + + def __next__(self) -> object: + raise self._error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_acompletion_streaming_iterator_honors_disable_fallbacks(opt_out): + """A chat stream that fails before its first chunk on a request that opted out of fallbacks + surfaces the primary's own error and never tries the fallback deployment.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_acompletion", new=AsyncMock(return_value=_AsyncList([]))) as fallback_attempt: + wrapped = await router._acompletion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +def test_completion_streaming_iterator_honors_disable_fallbacks(opt_out): + """Sync counterpart of test_acompletion_streaming_iterator_honors_disable_fallbacks.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_completion", new=MagicMock(return_value=iter([]))) as fallback_attempt: + wrapped = router._completion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(wrapped) + + assert raised.value is primary_error + fallback_attempt.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_aresponses_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Responses API mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _make_responses_iterator(error=_mid_stream_opt_out_trigger(primary_error), model="primary") + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_responses_attempt", + new=AsyncMock(return_value=_AsyncList([])), + ) as fallback_attempt: + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={ + "model": "primary", + "stream": True, + "input": "Hi", + "original_generic_function": litellm.aresponses, + **copy.deepcopy(opt_out), + }, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_anthropic_messages_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Anthropic Messages mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _AnthropicMessagesRaisingByteStream([], _mid_stream_opt_out_trigger(primary_error)) + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_anthropic_messages_attempt", + new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), + ) as fallback_attempt: + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source, + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop(): + """`disable_fallbacks=True` sent to the public entrypoint survives the fallback wrapper's handoff + into the stream: the primary's own error surfaces and no fallback deployment is ever called.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + async def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.acompletion", side_effect=primary_stream) as provider_calls: + response = await router.acompletion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in response] + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + +def test_completion_disable_fallbacks_reaches_the_mid_stream_hop(): + """Sync counterpart of test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.completion", side_effect=primary_stream) as provider_calls: + response = router.completion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(response) + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + @pytest.mark.asyncio @pytest.mark.parametrize( "raised_error", 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 97% rename from tests/test_litellm/test_unit_shard_missing_paths.py rename to tests/unit/test_unit_shard_missing_paths.py index b91c2cff764..4fa9c5bd3c1 100644 --- a/tests/test_litellm/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -36,6 +36,7 @@ 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, }, 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 99% rename from tests/test_litellm/test_utils.py rename to tests/unit/test_utils.py index 5dc09db4535..768d8955b8e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/unit/test_utils.py @@ -642,6 +642,10 @@ def validate_model_cost_values(model_data, exceptions=None): "output_cost_per_image_512", "output_cost_per_image_1024", "output_cost_per_image_1536", + "output_cost_per_image_0.5K", + "output_cost_per_image_1K", + "output_cost_per_image_2K", + "output_cost_per_image_4K", "input_cost_per_pixel", "output_cost_per_pixel", "input_cost_per_second", @@ -875,6 +879,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_image_512": {"type": "number"}, "output_cost_per_image_1024": {"type": "number"}, "output_cost_per_image_1536": {"type": "number"}, + "output_cost_per_image_0.5K": {"type": "number"}, + "output_cost_per_image_1K": {"type": "number"}, + "output_cost_per_image_2K": {"type": "number"}, + "output_cost_per_image_4K": {"type": "number"}, "output_cost_per_image_token": {"type": "number"}, "output_cost_per_video_token": {"type": "number"}, "output_cost_per_pixel": {"type": "number"}, @@ -3722,6 +3730,23 @@ def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> N assert filtered == {"provider_option": "kept"} +@pytest.mark.parametrize( + "provider_filter", + [ + litellm.utils.get_non_default_completion_params, + litellm.utils.get_non_default_transcription_params, + litellm.utils.filter_out_litellm_params, + ], +) +@pytest.mark.parametrize("setting", [("tag_regex", ["^team-a$"]), ("max_file_size_mb", 5)]) +def test_deployment_only_settings_copied_by_the_router_stay_out_of_provider_params( + provider_filter: Callable[[dict[str, object]], Mapping[str, object]], setting: tuple[str, object] +) -> None: + name, value = setting + filtered: Final = provider_filter({"provider_option": "kept", name: value}) + assert filtered == {"provider_option": "kept"}, filtered + + class TestGetOptionalParamsTencent: """Tests that tencent provider uses TencentChatConfig for parameter mapping.""" 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/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py new file mode 100644 index 00000000000..e421321aaaa --- /dev/null +++ b/tests/unit/types/test_litellm_params.py @@ -0,0 +1,655 @@ +import inspect +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field, fields +from operator import attrgetter +from types import MappingProxyType +from typing import Final, TypeAlias, cast, get_type_hints + +import httpx +import pytest +from aiohttp import ClientSession +from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm.caching.caching import Cache +from litellm.litellm_core_utils.get_litellm_params import ( + get_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy carrier +) +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.router_strategy.complexity_router.context_compaction import CompactionState +from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets +from litellm.types import litellm_params +from litellm.types import utils as types_utils +from litellm.types.caching import DynamicCacheControl +from litellm.types.litellm_params import ( + ADDRESSED_RESPONSE_ID_FIELD, + LITELLM_OWNED_ROOTS, + TRUSTED_CALLBACK_VARS_FIELD, + CachingOptions, + owned_wire_names, + wire, + wire_names, +) +from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage +from litellm.types.proxy.litellm_pre_call_utils import SecretFields +from litellm.types.router import ( + ConfigurableClientsideParamsCustomAuth, + CredentialLiteLLMParams, + DeploymentTypedDict, + RetryPolicy, + RouterConfig, + UpdateRouterConfig, +) +from litellm.types.router_weights import RouterWeights +from litellm.types.utils import ( + CustomPricingLiteLLMParams, + ModelResponse, + ModelResponseStream, + ProviderSpecificHeader, + StandardCallbackDynamicParams, + agentic_loop_internal_litellm_params, + all_litellm_params, + bedrock_batch_litellm_params, +) +from litellm.utils import ( + filter_out_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_completion_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_transcription_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier +) + +PROVIDER_KNOB: Final = "registry_test_provider_only_knob" + +CONNECTION_NAMES: Final = ( + "api_key", + "api_base", + "api_version", + "region_name", + "headers", + "provider_specific_header", + "client", + "shared_session", + "ssl_verify", + "request_timeout", + "force_timeout", + "stream_timeout", + "max_retries", + "tenant_id", + "client_id", + "client_secret", + "azure_username", + "azure_password", + "azure_scope", + "azure_ad_token_provider", + "litellm_credential_name", + "configurable_clientside_auth_params", + "use_xai_oauth", + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +OPTION_NAMES: Final = ( + "custom_llm_provider", + "azure", + "use_litellm_proxy", + "use_chat_completions_api", + "use_in_pass_through", + "allowed_openai_params", + "fallbacks", + "context_window_fallback_dict", + "num_retries", + "retry_policy", + "retry_strategy", + "routing_strategy", + "cooldown_time", + "allowed_model_region", + "enable_tag_filtering", + "fastest_response", + "provider_affinity_header", + "search_tool_name", + "model_list", + "model_info", + "rpm", + "tpm", + "itpm", + "otpm", + "default_api_key_rpm_limit", + "default_api_key_tpm_limit", + "max_parallel_requests", + "weight", + "order", + "tag_regex", + "max_file_size_mb", + "auto_router_config_path", + "auto_router_config", + "auto_router_default_model", + "auto_router_embedding_model", + "auto_router_max_input_chars", + "auto_router_routing_compression", + "auto_router_model_compression", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "adaptive_router_default_model", + "quality_router_config", + "quality_router_default_model", + "caching", + "cache", + "ttl", + "enable_prompt_caching", + "caching_groups", + "cost_per_query", + "base_model", + "max_budget", + "budget_duration", + "id", + "metadata", + "litellm_metadata", + "tags", + "litellm_trace_id", + "litellm_session_id", + "litellm_request_debug", + "logger_fn", + "verbose", + "no-log", + "max_agentic_loops", + "guardrails", + "prompt_id", + "prompt_variables", + "prompt_version", + "prompt_environment", + "prompt_label", + "litellm_system_prompt", + "custom_prompt_dict", + "roles", + "final_prompt_value", + "bos_token", + "eos_token", + "hf_model_name", + "supports_system_message", + "ensure_alternating_roles", + "user_continue_message", + "assistant_continue_message", + "disable_add_transform_inline_image_block", + "merge_reasoning_content_in_choices", + "enable_json_schema_validation", + "complete_response", + "stream_chunk_size", + "keepalive_seconds", + "allow_client_keepalive_override", + "mock_response", + "mock_timeout", +) + +AGENTIC_LOOP_STATE_NAMES: Final = ( + "_agentic_loop_depth", + "_agentic_loop_fingerprints", + "_agentic_loop_api_surface", + "_code_interpreter_interception_active", + "_code_interpreter_interception_sandbox_key", + "_code_interpreter_interception_session_scoped", + "_code_interpreter_interception_converted_stream", + "_websearch_interception_emit_native_blocks", + "_websearch_interception_converted_stream", + "_headroom_interception_converted_stream", +) + +INTERNAL_STATE_NAMES: Final = ( + "litellm_call_id", + "completion_call_id", + "model_alias_map", + "data_residency", + "litellm_logging_obj", + "preset_cache_key", + "cache_key", + "stream_response", + "_context_compaction_state", + *AGENTIC_LOOP_STATE_NAMES, + "_router_weights", + "fallback_depth", + "max_fallbacks", + "attempted_targets", + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + "acompletion", + "aembedding", + "aimg_generation", + "atext_completion", + "text_completion", + "allm_passthrough_route", + "async_call", +) + +BEDROCK_BATCH_NAMES: Final = ( + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +ARTIFACT_NAMES: Final = ("self", "use_client", "model_config", "rust") + +CALLBACK_VAR_NAMES: Final = tuple(StandardCallbackDynamicParams.__annotations__) + +PRICING_NAMES: Final = tuple(CustomPricingLiteLLMParams.model_fields) + +OWNED_NAMES: Final = ( + *CONNECTION_NAMES, + *OPTION_NAMES, + *INTERNAL_STATE_NAMES, + *ARTIFACT_NAMES, + *CALLBACK_VAR_NAMES, + *PRICING_NAMES, +) + +Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict + +CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType( + { # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers + "completion": get_non_default_completion_params, + "transcription": get_non_default_transcription_params, + "filter_out": filter_out_litellm_params, + } +) + + +@pytest.mark.parametrize("classifier_name", CLASSIFIERS) +@pytest.mark.parametrize("name", OWNED_NAMES) +def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: str) -> None: + provider_value: Final = object() + classify: Final = CLASSIFIERS[classifier_name] + + result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + assert result[PROVIDER_KNOB] is provider_value + + +def test_a_name_no_object_declares_reaches_the_provider() -> None: + result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type + + assert result == MappingProxyType({PROVIDER_KNOB: 1}) + + +def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str: + return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder + model=model_group, + messages=(MappingProxyType({"role": "user", "content": "shared prompt"}),), + metadata=MappingProxyType({"caching_groups": options.caching_groups, "model_group": model_group}), + ) + + +def test_caching_groups_is_a_flat_sequence_of_model_groups_that_share_one_cache_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + for callback_list in ("input_callback", "success_callback", "_async_success_callback"): + monkeypatch.setattr(litellm, callback_list, []) # mutable-ok: Cache() appends "cache" to these lists + options: Final = CachingOptions(caching_groups=(("gpt-4", "gpt-4o"), ("claude-3",))) + cache: Final = Cache() + + keys: Final = tuple(_cache_key_for_model_group(cache, group, options) for group in ("gpt-4", "gpt-4o", "claude-3")) + + assert (keys[0] == keys[1], keys[0] == keys[2]) == (True, False) + + +def test_all_litellm_params_is_exactly_the_owned_inventory() -> None: + assert frozenset(all_litellm_params) == frozenset(OWNED_NAMES) + assert frozenset(ARTIFACT_NAMES).isdisjoint(DECLARED_NAMES) + + +def test_every_owned_name_has_exactly_one_owner() -> None: + duplicated: Final = tuple(name for name in dict.fromkeys(all_litellm_params) if all_litellm_params.count(name) > 1) + + assert duplicated == () + + +@pytest.mark.parametrize( + ("exported", "declared"), + ( + pytest.param( + types_utils.TRUSTED_CALLBACK_VARS_FIELD, + litellm_params.TRUSTED_CALLBACK_VARS_FIELD, + id="TRUSTED_CALLBACK_VARS_FIELD", + ), + pytest.param( + types_utils.ADDRESSED_RESPONSE_ID_FIELD, + litellm_params.ADDRESSED_RESPONSE_ID_FIELD, + id="ADDRESSED_RESPONSE_ID_FIELD", + ), + ), +) +def test_types_utils_still_exports_the_field_constant(exported: str, declared: str) -> None: + assert exported == declared + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Leaf: + plain: int | None = None + renamed: int | None = field(default=None, metadata=wire("wire-name")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _OtherLeaf: + plain: int | None = None + trailing: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Root: + first: _Leaf + second: _OtherLeaf + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _RootDeclaringAKwargDirectly: + first: _Leaf + stray: int | None = None + + +def test_wire_names_are_the_field_names_in_declaration_order_unless_wire_renames_them() -> None: + assert wire_names(_Leaf) == ("plain", "wire-name") + + +def test_owned_wire_names_walk_leaves_in_declaration_order_and_keep_every_occurrence() -> None: + assert owned_wire_names(_Root) == ("plain", "wire-name", "plain", "trailing") + + +def test_owned_wire_names_refuse_a_root_that_declares_a_kwarg_outside_a_leaf() -> None: + with pytest.raises(TypeError): + owned_wire_names(_RootDeclaringAKwargDirectly) + + +def test_agentic_loop_names_concatenate_as_a_list() -> None: + extended: Final = agentic_loop_internal_litellm_params + ["caller_added"] # mutable-ok: list contract under test + + assert (type(extended), len(extended), frozenset(extended)) == ( + list, + len(AGENTIC_LOOP_STATE_NAMES) + 2, + frozenset((*AGENTIC_LOOP_STATE_NAMES, "max_agentic_loops", "caller_added")), + ) + + +def test_bedrock_batch_names_concatenate_as_a_tuple() -> None: + extended: Final = bedrock_batch_litellm_params + ("caller_added",) + + assert extended == (*BEDROCK_BATCH_NAMES, "caller_added") + + +def test_proxy_stamped_fields_keep_their_wire_names() -> None: + assert (TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD) == ( + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + ) + + +def test_all_litellm_params_concatenates_with_a_list_like_the_completion_entrypoint_does() -> None: + extended: Final = ["aembedding", "extra_headers"] + all_litellm_params # mutable-ok: list contract under test + + assert (type(extended), frozenset(extended)) == (list, frozenset(("aembedding", "extra_headers", *OWNED_NAMES))) + + +CARRIED_AND_FORWARDED: Final = frozenset(("drop_params", "hugging_face", "no_log", "replicate", "together_ai")) + +CARRIER_SIGNATURE: Final = inspect.signature(get_litellm_params) # pyright: ignore[reportUnknownArgumentType] # legacy + +CARRIED_PARAMS: Final = tuple( + name for name in CARRIER_SIGNATURE.parameters if name != "kwargs" and name not in CARRIED_AND_FORWARDED +) + + +@pytest.mark.parametrize("name", CARRIED_PARAMS) +def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None: + provider_value: Final = object() + + result: Final = CLASSIFIERS["completion"]( + {name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type + ) + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + + +TYPED_CONFIG_MODELS: Final[Mapping[str, tuple[type[BaseModel], ...]]] = MappingProxyType( + { + "credentials": (CredentialLiteLLMParams,), + "router": (RouterConfig, UpdateRouterConfig), + } +) + +DECLARED_NAMES: Final = frozenset(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) + +ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient +) +MockResponse: TypeAlias = str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + +TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = { + "ProviderClient": ProviderClient, + "ProviderSpecificHeader": ProviderSpecificHeader, + "ClientSession": ClientSession, + "AsyncAzureOpenAI": AsyncAzureOpenAI, + "AsyncOpenAI": AsyncOpenAI, + "AzureOpenAI": AzureOpenAI, + "OpenAI": OpenAI, + "AsyncHTTPHandler": AsyncHTTPHandler, + "HTTPHandler": HTTPHandler, + "ConfigurableClientsideParamsCustomAuth": ConfigurableClientsideParamsCustomAuth, + "RetryPolicy": RetryPolicy, + "DeploymentTypedDict": DeploymentTypedDict, + "DynamicCacheControl": DynamicCacheControl, + "ChatCompletionUserMessage": ChatCompletionUserMessage, + "ChatCompletionAssistantMessage": ChatCompletionAssistantMessage, + "MockResponse": MockResponse, + "ModelResponse": ModelResponse, + "ModelResponseStream": ModelResponseStream, + "Logging": Logging, + "SecretFields": SecretFields, + "CompactionState": CompactionState, + "RouterWeights": RouterWeights, + "AttemptedFallbackTargets": AttemptedFallbackTargets, +} + +LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)}, + litellm_params.DispatchOptions: {"custom_llm_provider": "openai"}, + litellm_params.RoutingOptions: { + "fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}], + "num_retries": 2, + "retry_strategy": "constant_retry", + "routing_strategy": "simple-shuffle", + }, + litellm_params.DeploymentOptions: {"model_info": {"region": "us"}, "rpm": 2}, + litellm_params.SpecializedRouterOptions: {"adaptive_router_default_model": "gpt-4o"}, + litellm_params.CachingOptions: {"ttl": 30.0, "caching_groups": (("gpt-4o", "gpt-4o-mini"),)}, + litellm_params.CostOptions: {"max_budget": 10.0}, + litellm_params.ObservabilityOptions: {"metadata": {"request": "test"}, "no_log": True}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2}, + litellm_params.GuardrailOptions: {"guardrails": ("default",)}, + litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}}, + litellm_params.ResponseOptions: {"stream_chunk_size": 64}, + litellm_params.MockOptions: {"mock_timeout": True}, + litellm_params.CallState: { + "completion_call_id": "call", + "model_alias_map": {"alias": "gpt-4o"}, + "data_residency": "us", + }, + litellm_params.AgenticLoopState: {"api_surface": "chat_completions", "depth": 1}, + litellm_params.RouterState: {"fallback_depth": 1}, + litellm_params.ProxyRequestState: { + "proxy_server_request": {"path": "/chat/completions"}, + "trusted_callback_vars": {"dd_api_key": "k"}, + }, + litellm_params.EntrypointState: {"acompletion": True}, +} + +LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": 1}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1}, + litellm_params.DispatchOptions: {"custom_llm_provider": 1}, + litellm_params.RoutingOptions: {"num_retries": "2"}, + litellm_params.DeploymentOptions: {"rpm": "2"}, + litellm_params.SpecializedRouterOptions: {"auto_router_max_input_chars": "2"}, + litellm_params.CachingOptions: {"ttl": "30"}, + litellm_params.CostOptions: {"max_budget": "10"}, + litellm_params.ObservabilityOptions: {"verbose": "true"}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"}, + litellm_params.GuardrailOptions: {"guardrails": (1,)}, + litellm_params.PromptOptions: {"prompt_id": 1}, + litellm_params.ResponseOptions: {"stream_chunk_size": "64"}, + litellm_params.MockOptions: {"mock_timeout": "true"}, + litellm_params.CallState: {"completion_call_id": 1}, + litellm_params.AgenticLoopState: {"depth": "1"}, + litellm_params.RouterState: {"fallback_depth": "1"}, + litellm_params.ProxyRequestState: {"proxy_server_request": "request"}, + litellm_params.EntrypointState: {"acompletion": "true"}, +} + +INVALID_LITERAL_SAMPLES: Final[tuple[tuple[type, Mapping[str, object]], ...]] = ( + (litellm_params.RoutingOptions, {"retry_strategy": "linear"}), + (litellm_params.RoutingOptions, {"routing_strategy": "random"}), + (litellm_params.AgenticLoopState, {"api_surface": "batches"}), +) + + +def _leaf_id(value: object) -> str: + return value.__name__ if isinstance(value, type) else "" + + +def _leaf_instance(leaf: type, sample: Mapping[str, object]) -> object: + constructor: Final = cast(Callable[..., object], leaf) + return constructor(**sample) + + +def _strict_leaf_validation(leaf: type, instance: object) -> object: + hints: Final[Mapping[str, object]] = cast( + Mapping[str, object], get_type_hints(type(instance), localns=TYPE_HINT_NAMESPACE) + ) + for field_info in fields(leaf): + value = cast(Callable[[object], object], attrgetter(field_info.name))(instance) + field_adapter: TypeAdapter[object] = TypeAdapter[object]( + hints[field_info.name], + config=ConfigDict(arbitrary_types_allowed=True), + ) + field_adapter.validate_python(value, strict=True) + return instance + + +@pytest.mark.parametrize("leaf,sample", LEAF_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + result: Final = _strict_leaf_validation(leaf, instance) + + assert result == instance + assert frozenset(sample) <= frozenset(field.name for field in fields(leaf)) + + +@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id) +def test_owned_leaf_literals_reject_unknown_values(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize( + "strategy", + [ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", + ], +) +def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) -> None: + instance: Final = _leaf_instance(litellm_params.RoutingOptions, {"routing_strategy": strategy}) + + assert _strict_leaf_validation(litellm_params.RoutingOptions, instance) is instance + + +NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "credentials": ( + "api_base", + "api_key", + "api_version", + "aws_batch_role_arn", + "azure_password", + "azure_scope", + "azure_username", + "bedrock_tags", + "client_id", + "client_secret", + "region_name", + "s3_access_key_id", + "s3_bucket_name", + "s3_bucket_owner", + "s3_encryption_key_id", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_region_name", + "s3_secret_access_key", + "tenant_id", + ), + "router": ( + "caching_groups", + "cooldown_time", + "enable_tag_filtering", + "fallbacks", + "max_retries", + "model_list", + "num_retries", + "retry_policy", + "routing_strategy", + ), + } +) + + +@pytest.mark.parametrize("source", TYPED_CONFIG_MODELS) +def test_names_a_typed_config_model_shares_with_the_owned_inventory_are_exactly_these(source: str) -> None: + model_names: Final = frozenset(name for model in TYPED_CONFIG_MODELS[source] for name in model.model_fields) + + assert DECLARED_NAMES & model_names == frozenset(NAMES_SHARED_WITH_TYPED_MODELS[source]) + + +@pytest.mark.parametrize("name", PRICING_NAMES) +def test_pricing_name_is_owned_by_the_pricing_model_alone(name: str) -> None: + assert name not in DECLARED_NAMES 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/test_litellm/messages/__init__.py b/tests/unit/vector_stores/__init__.py similarity index 100% rename from tests/test_litellm/messages/__init__.py rename to tests/unit/vector_stores/__init__.py 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 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index beeb44474da..ee7aa22f759 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -37,5 +37,8 @@ }, "LIT013": { "limit": 0 + }, + "LIT014": { + "limit": 369 } } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx index 1c39c6eb36d..b06986d01f0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx @@ -46,13 +46,13 @@ const dayWithModels = (date: string, models: Record [ name, { metrics: baseMetrics(m), metadata: {}, api_key_breakdown: {} }, ]), ), - model_groups: {}, mcp_servers: {}, providers: {}, api_keys: {}, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.test.ts index 9c3915c812f..20b64179ce8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.test.ts @@ -56,10 +56,10 @@ const modelDay = (date: string, models: Record>): date, metrics: metrics({}), breakdown: { - models: Object.fromEntries( + models: {}, + model_groups: Object.fromEntries( Object.entries(models).map(([name, m]) => [name, { metrics: metrics(m), metadata: {}, api_key_breakdown: {} }]), ), - model_groups: {}, mcp_servers: {}, providers: {}, entities: {}, @@ -244,6 +244,54 @@ describe("computeCacheLeakage by model", () => { expect(rows.map((r) => r.id)).toEqual(["gemini-2.5-flash"]); expect(rows[0].potentialSavings).toBeCloseTo(1.0, 6); }); + + it("merges rows logged under a deployment's resolved and requested names into one model group row", () => { + const day: DailyData = { + date: "2026-07-01", + metrics: metrics({}), + breakdown: { + models: { + "bedrock/global.anthropic.claude-sonnet-4-6": { + metrics: metrics({ prompt_tokens: 270000 }), + metadata: {}, + api_key_breakdown: {}, + }, + "bedrock/claude-sonnet-4-6": { + metrics: metrics({ prompt_tokens: 5000 }), + metadata: {}, + api_key_breakdown: {}, + }, + }, + model_groups: { + "bedrock/claude-sonnet-4-6": { + metrics: metrics({ prompt_tokens: 275000 }), + metadata: {}, + api_key_breakdown: {}, + }, + }, + mcp_servers: {}, + providers: {}, + entities: {}, + api_keys: {}, + }, + }; + const { rows } = computeCacheLeakage([day], "model"); + expect(rows.map((r) => r.id)).toEqual(["bedrock/claude-sonnet-4-6"]); + expect(rows[0].uncachedPromptTokens).toBe(275000); + }); + + it("sums a model group across days", () => { + const results = [ + modelDay("2026-07-01", { "bedrock/claude-sonnet-4-6": { prompt_tokens: 1000 } }), + modelDay("2026-07-02", { + "bedrock/claude-sonnet-4-6": { prompt_tokens: 2500, cache_read_input_tokens: 500 }, + }), + ]; + const { rows } = computeCacheLeakage(results, "model"); + expect(rows).toHaveLength(1); + expect(rows[0].uncachedPromptTokens).toBe(3000); + expect(rows[0].cacheHitRatio).toBeCloseTo(500 / 3500, 6); + }); }); describe("buildDailyToolSeries", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts index 2e6d8208989..5f16b1b04fd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts @@ -93,7 +93,7 @@ const aggregateByKey = (results: readonly DailyData[]): Map => { const byModel = new Map(); for (const day of results) { - for (const [model, entry] of Object.entries(day.breakdown?.models ?? {})) { + for (const [model, entry] of Object.entries(day.breakdown?.model_groups ?? {})) { const acc = byModel.get(model) ?? emptyAccumulator(); byModel.set(model, addMetrics(acc, entry.metrics, null, null)); } diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx index 149d231fff6..d3e3f204813 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx @@ -133,9 +133,10 @@ describe("MCPToolPermissions", () => { const selectAllButton = screen.getByRole("button", { name: "Select All" }); await userEvent.click(selectAllButton); - // Verify onChange was called with all tools selected + // Selecting every displayed tool writes the wildcard, which also covers tools the + // server adds later. expect(mockOnChange).toHaveBeenCalledWith({ - [mockServerId]: ["read_wiki_structure", "read_wiki_contents", "ask_question"], + [mockServerId]: ["*"], }); }); @@ -190,6 +191,77 @@ describe("MCPToolPermissions", () => { }); }); + describe("wildcard all-tools grant", () => { + const wildcardServerId = "server-1"; + const wildcardServer = { server_id: wildcardServerId, server_name: "Wildcard Server", alias: "Wildcard Server" }; + const wildcardTools = [ + { name: "read_wiki_structure", description: "Get documentation topics" }, + { name: "read_wiki_contents", description: "View documentation" }, + { name: "ask_question", description: "Ask questions" }, + ]; + + beforeEach(() => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([wildcardServer]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: wildcardTools, error: false }); + }); + + it("renders every tool checked with the future-tools note when the entry is the wildcard", async () => { + renderWithProviders( + , + ); + + expect(await screen.findByText("Wildcard Server")).toBeInTheDocument(); + expect(screen.getByText("All tools allowed, including tools added to this server later")).toBeInTheDocument(); + + await userEvent.click(screen.getByText("Flat List")); + for (const checkbox of screen.getAllByRole("checkbox")) { + expect(checkbox).toBeChecked(); + } + }); + + it("writes the wildcard when Select All covers every displayed tool", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Select All" })); + + expect(mockOnChange).toHaveBeenCalledWith({ [wildcardServerId]: ["*"] }); + }); + + it("converts back to an enumerated list when one tool is unchecked from a wildcard grant", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByText("Flat List")); + await userEvent.click(screen.getByRole("checkbox", { name: "ask_question" })); + + expect(mockOnChange).toHaveBeenCalledWith({ + [wildcardServerId]: ["read_wiki_structure", "read_wiki_contents"], + }); + }); + }); + describe("servers reached indirectly", () => { const groupServer = { server_id: "srv-group-1", @@ -428,6 +500,8 @@ describe("MCPToolPermissions", () => { expect(await screen.findByText("list_issues")).toBeInTheDocument(); await userEvent.click(screen.getByText("Select All")); + // A toolset-sourced server never writes the wildcard: that would create a standing direct + // grant outliving the toolset. The write keeps only the tools this level grants itself. expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["delete_issue"] }); }); @@ -849,7 +923,7 @@ describe("MCPToolPermissions", () => { const written = mockOnChange.mock.calls.at(-1)?.[0] as Record; expect(written["github_mcp"]).toEqual(["list_issues"]); - expect(written[twin.server_id]).toEqual(["list_issues", "create_issue", "delete_issue"]); + expect(written[twin.server_id]).toEqual(["*"]); }); it("says nothing about shared names when every key names one server", async () => { diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx index e26f1a6f511..7edeaedff2e 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx @@ -8,13 +8,14 @@ import { useMCPAccessGroups } from "../../app/(dashboard)/hooks/mcpServers/useMC import { useMCPToolsets } from "../../app/(dashboard)/hooks/mcpServers/useMCPToolsets"; import McpCrudPermissionPanel from "../mcp_tools/McpCrudPermissionPanel"; import { classifyToolOp } from "../../utils/mcpToolCrudClassification"; -import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import { MCP_ALL_TOOLS_WILDCARD, NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; import { EffectiveMcpServer, McpGrantSource, applyToolPermissionWrite, emptyMcpAccessGroups, mcpAllowedToolsFor, + mcpGrantsAllTools, resolveEffectiveMcpServers, } from "./effectiveMcpServers"; @@ -150,7 +151,12 @@ const MCPToolPermissions: React.FC = ({ // Every write goes through here so an edit is authoritative for the SERVER, not for one of the // equivalent keys that may name it. const writeAllowedTools = (entry: EffectiveMcpServer, allowed: string[]) => { - onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed })); + const names = (serverTools[entry.server.server_id] ?? []).map((t) => t.name); + const next = + entry.source.kind !== "toolset" && names.length > 0 && names.every((n) => allowed.includes(n)) + ? [MCP_ALL_TOOLS_WILDCARD] + : allowed; + onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed: next })); }; const handleSelectAll = (entry: EffectiveMcpServer) => { @@ -222,7 +228,8 @@ const MCPToolPermissions: React.FC = ({ const serverId = server.server_id; const serverName = server.server_name || server.alias || serverId; const tools = serverTools[serverId] || []; - const selectedTools = entry.allowedTools ?? tools.map((t) => t.name); + const grantsAll = mcpGrantsAllTools(entry.keyedTools); + const selectedTools = grantsAll ? tools.map((t) => t.name) : entry.allowedTools ?? tools.map((t) => t.name); const isLoading = loadingTools[serverId]; const error = toolErrors[serverId]; const viewMode = viewModes[serverId] ?? "crud"; @@ -247,6 +254,11 @@ const MCPToolPermissions: React.FC = ({ )} {server.description &&

{server.description}

} + {grantsAll && ( +

+ All tools allowed, including tools added to this server later +

+ )} {entry.ambiguousKeys.length > 0 && (

{`Also granted by ${entry.ambiguousKeys.map((key) => `"${key}"`).join(", ")}, which names another server too. Those tools stay allowed here until the servers no longer share that name`} diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts index 487c6f9f55e..07f2b0e2508 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts @@ -4,6 +4,7 @@ import { applyToolPermissionWrite, emptyMcpAccessGroups, mcpAllowedToolsFor, + mcpGrantsAllTools, mcpServersForIdentifier, mcpToolPermissionKeyFor, resolveEffectiveMcpServers, @@ -66,6 +67,16 @@ describe("mcpServersForIdentifier", () => { }); }); +describe("mcpGrantsAllTools", () => { + it("is true only when the union carries the wildcard, never for an absent grant", () => { + expect(mcpGrantsAllTools(["*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file", "*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file"])).toBe(false); + expect(mcpGrantsAllTools([])).toBe(false); + expect(mcpGrantsAllTools(undefined)).toBe(false); + }); +}); + describe("mcpToolPermissionKeyFor", () => { const target = server({ server_id: "uuid-1", server_name: "github_mcp", alias: "GitHub" }); diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts index b3ba24f3c59..c9e85fef31b 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts @@ -1,5 +1,6 @@ import { z } from "zod/v4"; import { MCPServer, MCPToolset } from "../mcp_tools/types"; +import { MCP_ALL_TOOLS_WILDCARD } from "../mcp_tools/constants"; // Mirrors the backend resolver's union (direct + access_group + tool_perm + toolset), so the // editor shows exactly the servers this permission level entitles. @@ -121,6 +122,12 @@ export const mcpAllowedToolsFor = ( return [...new Set(keys.flatMap((key) => toolPermissions[key] ?? []))]; }; +// An allowed-tools union carrying the wildcard grants every current and future tool on the +// server; `undefined` (no entry at all) is unrestricted for a different reason and is not a +// wildcard grant the editor should expand. +export const mcpGrantsAllTools = (allowed: readonly string[] | undefined): boolean => + allowed !== undefined && allowed.includes(MCP_ALL_TOOLS_WILDCARD); + // Tool names the given toolsets grant on this server, `undefined` when they grant none. const mcpToolsetToolsFor = ( server: MCPServer, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts index 66ab1a352f4..eef98383d2a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts +++ b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts @@ -3,5 +3,8 @@ export const NO_MCP_SERVERS_SENTINEL = "no-mcp-servers"; export const ALL_PROXY_MCP_SERVERS_SENTINEL = "all-proxy-mcpservers"; +// Must match the backend MCP_ALL_TOOLS_WILDCARD constant in litellm/constants.py. +export const MCP_ALL_TOOLS_WILDCARD = "*"; + export const MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE = "Tool preview is not available for submissions. Tools will be verified by an admin during review."; diff --git a/ui/litellm-dashboard/src/components/team/teamAdminEditAccess.ts b/ui/litellm-dashboard/src/components/team/teamAdminEditAccess.ts index 5706eeafe82..1d8cd90c517 100644 --- a/ui/litellm-dashboard/src/components/team/teamAdminEditAccess.ts +++ b/ui/litellm-dashboard/src/components/team/teamAdminEditAccess.ts @@ -48,6 +48,7 @@ const TEAM_ADMIN_FIELD_LABELS: ReadonlyMap = new Map([ ["rpm_limit", "Requests per minute Limit (RPM)"], ["max_budget", "Max Budget (USD)"], ["projects", "Create and update projects"], + ["member_key_budgets", "Update budgets on team members' keys"], ]); export const teamAdminFieldLabel = (field: string): string => TEAM_ADMIN_FIELD_LABELS.get(field) ?? field; diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.integration.test.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.integration.test.tsx index 4d0afde6f25..9e3e19f1e1b 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.integration.test.tsx @@ -303,6 +303,7 @@ describe("KeyEditView", () => { fallbacks: [{ "gpt-4": ["gpt-4o", "gpt-4o-mini"] }], }), }), + expect.any(Array), ); }); }); @@ -323,6 +324,7 @@ describe("KeyEditView", () => { fallbacks: null, }), }), + expect.any(Array), ); }); }); @@ -632,7 +634,10 @@ describe("KeyEditView", () => { await userEvent.click(screen.getByRole("button", { name: /save changes/i })); await waitFor(() => { - expect(onSubmitMock).toHaveBeenCalledWith(expect.objectContaining({ throttle_on_budget_exceeded: true })); + expect(onSubmitMock).toHaveBeenCalledWith( + expect.objectContaining({ throttle_on_budget_exceeded: true }), + expect.any(Array), + ); }); }); @@ -662,7 +667,10 @@ describe("KeyEditView", () => { await userEvent.click(screen.getByRole("button", { name: /save changes/i })); await waitFor(() => { - expect(onSubmitMock).toHaveBeenCalledWith(expect.objectContaining({ enable_prompt_caching: true })); + expect(onSubmitMock).toHaveBeenCalledWith( + expect.objectContaining({ enable_prompt_caching: true }), + expect.any(Array), + ); }); }); @@ -1526,7 +1534,10 @@ describe("KeyEditView", () => { await userEvent.click(screen.getByRole("button", { name: /save changes/i })); await waitFor(() => { - expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining({ organization_id: null, team_id: null })); + expect(onSubmit).toHaveBeenCalledWith( + expect.objectContaining({ organization_id: null, team_id: null }), + expect.any(Array), + ); }); expect(JSON.parse(JSON.stringify(onSubmit.mock.calls[0][0]))).toMatchObject({ organization_id: null, @@ -1568,7 +1579,9 @@ describe("KeyEditView", () => { await userEvent.click(await screen.findByRole("button", { name: "Detach from project" })); await userEvent.click(screen.getByRole("button", { name: /save changes/i })); const expectedDetach = { project_id: null, organization_id: "org-1", team_id: "group-maple", models: key.models }; - await waitFor(() => expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining(expectedDetach))); + await waitFor(() => + expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining(expectedDetach), expect.any(Array)), + ); expect(screen.getByRole("combobox", { name: "Team ID" })).toBeDisabled(); view.rerender(renderEditor({ ...key, project_id: null })); expect(screen.getByRole("combobox", { name: "Team ID" })).toBeEnabled(); @@ -1866,7 +1879,10 @@ describe("KeyEditView", () => { await save(); await waitFor(() => { - expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining({ end_user_budget_id: "svc-b-budget" })); + expect(onSubmit).toHaveBeenCalledWith( + expect.objectContaining({ end_user_budget_id: "svc-b-budget" }), + expect.any(Array), + ); }); }); @@ -1878,7 +1894,7 @@ describe("KeyEditView", () => { await save(); await waitFor(() => { - expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining({ end_user_budget_id: "" })); + expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining({ end_user_budget_id: "" }), expect.any(Array)); }); }); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index 9cd97f4ef98..1342b97d1a0 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -80,7 +80,7 @@ import VectorStoreSelector from "../vector_store_management/VectorStoreSelector" interface KeyEditViewProps { keyData: KeyResponse; onCancel: () => void; - onSubmit: (values: any) => Promise; + onSubmit: (values: any, dirtyFields: readonly string[]) => Promise; teams?: any[] | null; accessToken: string | null; userID: string | null; @@ -317,6 +317,7 @@ export function KeyEditView({ ...values, ...(detachProject && enableProjectsUI && canDetachProject ? { project_id: null } : {}), }), + [...Object.keys(form.formState.dirtyFields), ...(budgetLimitsUnchanged ? [] : ["budget_limits"])], ); } finally { setIsKeySaving(false); diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx index 63693fd1af5..7eb09926caf 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx @@ -49,6 +49,7 @@ import { RegenerateKeyModal } from "../organisms/RegenerateKeyModal"; import { parseErrorMessage } from "../shared/errorUtils"; import { InheritedBudgetHint, inheritedBudgetGates, keyOwnerBudgetSource } from "../shared/InheritedBudgetHint"; import { KeyEditView } from "./key_edit_view"; +import { isTeamAdminEditingMemberKey, teamAdminMemberKeyPayload } from "./teamAdminMemberKeyPayload"; export function needsLifetimeSpendBackfill(spend: number, totalSpend: number | null | undefined): boolean { return (totalSpend ?? 0) < spend; @@ -187,7 +188,7 @@ export default function KeyInfoView({ ); } - const handleKeyUpdate = async (formValues: Record) => { + const handleKeyUpdate = async (formValues: Record, dirtyFields: readonly string[] = []) => { try { if (!accessToken) return; @@ -359,6 +360,25 @@ export default function KeyInfoView({ formValues.budget_duration = wordToCanonical[formValues.budget_duration] ?? formValues.budget_duration; } + const memberKeyEditContext = { + userRole: userRole || "", + userId: userID || "", + keyUserId: currentKeyData.user_id, + keyTeamId: currentKeyData.team_id, + teamMembers: teamsData?.find((team) => team.team_id === currentKeyData.team_id)?.members_with_roles, + }; + const editingMemberKeyAsTeamAdmin = isTeamAdminEditingMemberKey(memberKeyEditContext); + if (editingMemberKeyAsTeamAdmin) { + const trimmed = teamAdminMemberKeyPayload(formValues, dirtyFields); + if (trimmed.kind === "blocked") { + toast.error( + `Team admins can only change budget fields on other members' keys, not ${trimmed.fields.join(", ")}`, + ); + return; + } + formValues = trimmed.payload; + } + const newKeyValues = await keyUpdateCall(accessToken, formValues); // Update local state diff --git a/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.test.ts b/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.test.ts new file mode 100644 index 00000000000..a43b41abeb0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.test.ts @@ -0,0 +1,96 @@ +import { describe, expect, it } from "vitest"; +import { Member } from "@/components/networking"; +import { isTeamAdminEditingMemberKey, KEY_BUDGET_FIELDS, teamAdminMemberKeyPayload } from "./teamAdminMemberKeyPayload"; + +const members = (role: string): Member[] => [{ user_id: "admin-user", role, user_email: null } as unknown as Member]; + +const baseArgs = { + userRole: "Internal User", + userId: "admin-user", + keyUserId: "member-user", + keyTeamId: "team-1", +}; + +describe("isTeamAdminEditingMemberKey", () => { + it("is false for a proxy admin", () => { + expect(isTeamAdminEditingMemberKey({ ...baseArgs, userRole: "Admin", teamMembers: members("admin") })).toBe(false); + }); + + it("is false when the caller owns the key", () => { + expect(isTeamAdminEditingMemberKey({ ...baseArgs, keyUserId: "admin-user", teamMembers: members("admin") })).toBe( + false, + ); + }); + + it("is false for a personal key with no team", () => { + expect(isTeamAdminEditingMemberKey({ ...baseArgs, keyTeamId: null, teamMembers: members("admin") })).toBe(false); + }); + + it("is false when the caller is not a team admin", () => { + expect(isTeamAdminEditingMemberKey({ ...baseArgs, teamMembers: members("user") })).toBe(false); + expect(isTeamAdminEditingMemberKey({ ...baseArgs, teamMembers: null })).toBe(false); + expect(isTeamAdminEditingMemberKey({ ...baseArgs, teamMembers: undefined })).toBe(false); + }); + + it("is true for a team admin editing another member's team key", () => { + expect(isTeamAdminEditingMemberKey({ ...baseArgs, teamMembers: members("admin") })).toBe(true); + }); +}); + +describe("teamAdminMemberKeyPayload", () => { + it("keeps only dirty budget fields from the form values plus the key", () => { + const formValues = { + key: "sk-1", + max_budget: 25, + soft_budget: 10, + key_alias: "renamed", + metadata: { tags: ["a"] }, + tpm_limit: null, + }; + const result = teamAdminMemberKeyPayload(formValues, ["max_budget", "soft_budget"]); + expect(result).toEqual({ + kind: "ok", + payload: { key: "sk-1", max_budget: 25, soft_budget: 10 }, + }); + }); + + it("drops budget fields present in the form but not dirty", () => { + const formValues = { + key: "sk-1", + max_budget: 25, + budget_duration: "30d", + budget_limits: [{ budget_duration: "1d", max_budget: 5 }], + }; + const result = teamAdminMemberKeyPayload(formValues, ["max_budget"]); + expect(result).toEqual({ kind: "ok", payload: { key: "sk-1", max_budget: 25 } }); + }); + + it("drops budget_duration when it is an empty string but keeps null", () => { + const cleared = teamAdminMemberKeyPayload({ key: "sk-1", budget_duration: "" }, ["budget_duration"]); + expect(cleared).toEqual({ kind: "ok", payload: { key: "sk-1" } }); + const kept = teamAdminMemberKeyPayload({ key: "sk-1", budget_duration: null }, ["budget_duration"]); + expect(kept).toEqual({ kind: "ok", payload: { key: "sk-1", budget_duration: null } }); + }); + + it("keeps budget_limits when present", () => { + const windows = [{ budget_duration: "1d", max_budget: 5 }]; + const result = teamAdminMemberKeyPayload({ key: "sk-1", budget_limits: windows }, ["budget_limits"]); + expect(result).toEqual({ kind: "ok", payload: { key: "sk-1", budget_limits: windows } }); + }); + + it("is blocked when a dirty field is not a budget field, naming it", () => { + const result = teamAdminMemberKeyPayload({ key: "sk-1", key_alias: "renamed" }, ["key_alias", "max_budget"]); + expect(result).toEqual({ kind: "blocked", fields: ["key_alias"] }); + }); + + it("is ok when every dirty field is a budget field and ignores token/key", () => { + const result = teamAdminMemberKeyPayload({ key: "sk-1", max_budget: 5 }, ["token", "key", "max_budget"]); + expect(result).toEqual({ kind: "ok", payload: { key: "sk-1", max_budget: 5 } }); + }); + + it("covers exactly the backend budget field set", () => { + expect([...KEY_BUDGET_FIELDS].sort()).toEqual( + ["budget_duration", "budget_limits", "max_budget", "soft_budget"].sort(), + ); + }); +}); diff --git a/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.ts b/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.ts new file mode 100644 index 00000000000..abd8653dde5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/templates/teamAdminMemberKeyPayload.ts @@ -0,0 +1,41 @@ +import { Member } from "@/components/networking"; +import { isProxyAdminRole, isUserTeamAdminForSingleTeam } from "@/utils/roles"; + +export const KEY_BUDGET_FIELDS = ["max_budget", "soft_budget", "budget_duration", "budget_limits"] as const; + +export const isTeamAdminEditingMemberKey = (args: { + userRole: string; + userId: string; + keyUserId: string | null | undefined; + keyTeamId: string | null | undefined; + teamMembers: Member[] | null | undefined; +}): boolean => { + if (isProxyAdminRole(args.userRole)) return false; + if (!args.keyTeamId) return false; + if (args.keyUserId === args.userId) return false; + return isUserTeamAdminForSingleTeam(args.teamMembers ?? null, args.userId); +}; + +export type TeamAdminMemberKeyPayload = + | { kind: "ok"; payload: Record } + | { kind: "blocked"; fields: readonly string[] }; + +export const teamAdminMemberKeyPayload = ( + formValues: Record, + dirtyFields: readonly string[], +): TeamAdminMemberKeyPayload => { + const disallowed = dirtyFields.filter( + (field) => field !== "token" && field !== "key" && !(KEY_BUDGET_FIELDS as readonly string[]).includes(field), + ); + if (disallowed.length > 0) { + return { kind: "blocked", fields: disallowed }; + } + const payload: Record = { key: formValues.key }; + for (const field of KEY_BUDGET_FIELDS) { + if (!dirtyFields.includes(field)) continue; + if (formValues[field] === undefined) continue; + if (field === "budget_duration" && formValues[field] === "") continue; + payload[field] = formValues[field]; + } + return { kind: "ok", payload }; +}; diff --git a/uv.lock b/uv.lock index 85e2b6d4e52..c235171ecb2 100644 --- a/uv.lock +++ b/uv.lock @@ -4256,21 +4256,22 @@ wheels = [ [[package]] name = "langfuse" -version = "2.59.7" +version = "4.15.2" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "anyio" }, { name = "backoff" }, { name = "httpx" }, - { name = "idna" }, + { name = "opentelemetry-api" }, + { name = "opentelemetry-exporter-otlp-proto-http" }, + { name = "opentelemetry-sdk" }, { name = "packaging" }, { name = "pydantic" }, - { name = "requests" }, + { name = "typing-extensions" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/d5/0e/8390bd3a4ad92ecb1ba0462ec8b7c7d328b2e2f31ae0e734bf2f50dbdc96/langfuse-2.59.7.tar.gz", hash = "sha256:f631981705177bf53d030d191397da9b864b99729a7273448afed10d76f78e23", size = 146608, upload-time = "2025-03-03T16:30:59.926Z" } +sdist = { url = "https://files.pythonhosted.org/packages/97/30/6a64dcf84de2f2eb4d03adbfd22cc7bdc95ce67e3e56cd6288405087fb8a/langfuse-4.15.2.tar.gz", hash = "sha256:7f818f38cc22daba88fdcec62d2addcee4e18d1af4529978b6d07501e86b6946", size = 391727, upload-time = "2026-09-09T16:01:25.73Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b7/f3/420518b9003c997cdcb0a86473bf0c111181578a95565823c333cb58eb7b/langfuse-2.59.7-py3-none-any.whl", hash = "sha256:2c6890f5b842257173eb54d08f2890c7fd7617859a48b3914ef73f13a6514473", size = 260468, upload-time = "2025-03-03T16:30:57.426Z" }, + { url = "https://files.pythonhosted.org/packages/02/de/e59da18cb5ca9cb8515a254199bd96ace8cf918788d13ded66aa2693e007/langfuse-4.15.2-py3-none-any.whl", hash = "sha256:98c27a3c06e18c4497045f2d4decce2c716ef11cb215cf5b27bbea6ee0877115", size = 705824, upload-time = "2026-09-09T16:01:23.696Z" }, ] [[package]] @@ -4794,7 +4795,7 @@ requires-dist = [ { name = "jinja2", specifier = ">=3.1.6,<4.0" }, { name = "jsonschema", specifier = ">=4.0.0,<5.0" }, { name = "keyring", marker = "extra == 'cli'", specifier = ">=25.6.0,<26.0" }, - { name = "langfuse", marker = "extra == 'proxy-runtime'", specifier = ">=2.59.7,<3.0" }, + { name = "langfuse", marker = "extra == 'proxy-runtime'", specifier = ">=4.7,<5.0" }, { name = "litellm-enterprise", marker = "extra == 'proxy'", editable = "enterprise" }, { name = "litellm-proxy-extras", marker = "extra == 'proxy'", editable = "litellm-proxy-extras" }, { name = "llm-sandbox", marker = "extra == 'proxy-runtime'", specifier = ">=0.3.39,<1.0" }, @@ -4806,10 +4807,10 @@ requires-dist = [ { name = "numpydoc", marker = "extra == 'utils'", specifier = ">=1.8.0,<2.0" }, { name = "nvidia-riva-client", marker = "extra == 'stt-nvidia-riva'", specifier = ">=2.15.0" }, { name = "openai", specifier = ">=2.20.0,<3.0.0" }, - { name = "opentelemetry-api", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" }, - { name = "opentelemetry-exporter-otlp", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" }, - { name = "opentelemetry-instrumentation-fastapi", marker = "extra == 'proxy-runtime'", specifier = "==0.49b0" }, - { name = "opentelemetry-sdk", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" }, + { name = "opentelemetry-api", marker = "extra == 'proxy-runtime'", specifier = "==1.33.1" }, + { name = "opentelemetry-exporter-otlp", marker = "extra == 'proxy-runtime'", specifier = "==1.33.1" }, + { name = "opentelemetry-instrumentation-fastapi", marker = "extra == 'proxy-runtime'", specifier = "==0.54b1" }, + { name = "opentelemetry-sdk", marker = "extra == 'proxy-runtime'", specifier = "==1.33.1" }, { name = "orjson", marker = "extra == 'proxy'", specifier = ">=3.11.6,<4.0" }, { name = "packaging", specifier = ">=24.0" }, { name = "polars", marker = "extra == 'proxy'", specifier = ">=1.38.1,<2.0" }, @@ -4889,7 +4890,7 @@ ci = [ { name = "pytest-codspeed", specifier = "==4.3.0" }, { name = "pytest-retry", specifier = "==1.7.0" }, { name = "tenacity", specifier = "==8.5.0" }, - { name = "traceloop-sdk", specifier = "==0.33.12" }, + { name = "traceloop-sdk", specifier = "==0.34.0" }, ] dev = [ { name = "basedpyright", specifier = "==1.39.7" }, @@ -4899,14 +4900,14 @@ dev = [ { name = "fastapi-offline", specifier = "==1.7.6" }, { name = "hypothesis", specifier = "==6.165.10" }, { name = "keyring", specifier = "==25.7.0" }, - { name = "langfuse", specifier = "==2.59.7" }, + { name = "langfuse", specifier = ">=4.7,<5.0" }, { name = "mypy", specifier = "==1.20.1" }, { name = "numpy", specifier = ">=1.26.0,<3.0" }, { name = "openapi-core", specifier = "==0.22.0" }, - { name = "opentelemetry-api", specifier = "==1.28.0" }, - { name = "opentelemetry-exporter-otlp", specifier = "==1.28.0" }, - { name = "opentelemetry-instrumentation-fastapi", specifier = "==0.49b0" }, - { name = "opentelemetry-sdk", specifier = "==1.28.0" }, + { name = "opentelemetry-api", specifier = "==1.33.1" }, + { name = "opentelemetry-exporter-otlp", specifier = "==1.33.1" }, + { name = "opentelemetry-instrumentation-fastapi", specifier = "==0.54b1" }, + { name = "opentelemetry-sdk", specifier = "==1.33.1" }, { name = "parameterized", specifier = "==0.9.0" }, { name = "psycopg", specifier = "==3.3.3" }, { name = "psycopg-binary", specifier = "==3.3.3" }, @@ -4949,22 +4950,22 @@ proxy-dev = [ { name = "a2a-sdk", specifier = "==1.1.0" }, { name = "azure-identity", specifier = "==1.25.2" }, { name = "hypercorn", specifier = "==0.17.3" }, - { name = "opentelemetry-api", specifier = "==1.28.0" }, - { name = "opentelemetry-exporter-otlp", specifier = "==1.28.0" }, - { name = "opentelemetry-instrumentation-fastapi", specifier = "==0.49b0" }, - { name = "opentelemetry-sdk", specifier = "==1.28.0" }, + { name = "opentelemetry-api", specifier = "==1.33.1" }, + { name = "opentelemetry-exporter-otlp", specifier = "==1.33.1" }, + { name = "opentelemetry-instrumentation-fastapi", specifier = "==0.54b1" }, + { name = "opentelemetry-sdk", specifier = "==1.33.1" }, { name = "prisma", specifier = "==0.11.0" }, { name = "prometheus-client", specifier = "==0.20.0" }, ] [[package]] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" source = { editable = "litellm-proxy-extras" } [[package]] @@ -6152,45 +6153,45 @@ wheels = [ [[package]] name = "opentelemetry-api" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "deprecated" }, { name = "importlib-metadata" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/79/36/260eaea0f74fdd0c0d8f22ed3a3031109ea1c85531f94f4fde266c29e29a/opentelemetry_api-1.28.0.tar.gz", hash = "sha256:578610bcb8aa5cdcb11169d136cc752958548fb6ccffb0969c1036b0ee9e5353", size = 62803, upload-time = "2024-11-05T19:14:45.497Z" } +sdist = { url = "https://files.pythonhosted.org/packages/9a/8d/1f5a45fbcb9a7d87809d460f09dc3399e3fbd31d7f3e14888345e9d29951/opentelemetry_api-1.33.1.tar.gz", hash = "sha256:1c6055fc0a2d3f23a50c7e17e16ef75ad489345fd3df1f8b8af7c0bbf8a109e8", size = 65002, upload-time = "2025-05-16T18:52:41.146Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/22/e4/3b25d8b856791c04d8a62b1257b5fc09dc41a057800db06885af8ddcdce1/opentelemetry_api-1.28.0-py3-none-any.whl", hash = "sha256:8457cd2c59ea1bd0988560f021656cecd254ad7ef6be4ba09dbefeca2409ce52", size = 64314, upload-time = "2024-11-05T19:14:21.659Z" }, + { url = "https://files.pythonhosted.org/packages/05/44/4c45a34def3506122ae61ad684139f0bbc4e00c39555d4f7e20e0e001c8a/opentelemetry_api-1.33.1-py3-none-any.whl", hash = "sha256:4db83ebcf7ea93e64637ec6ee6fabee45c5cbe4abd9cf3da95c43828ddb50b83", size = 65771, upload-time = "2025-05-16T18:52:17.419Z" }, ] [[package]] name = "opentelemetry-exporter-otlp" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-exporter-otlp-proto-grpc" }, { name = "opentelemetry-exporter-otlp-proto-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/eb/16/14e3fc163930ea68f0980a4cdd4ae5796e60aeb898965990e13263d64baf/opentelemetry_exporter_otlp-1.28.0.tar.gz", hash = "sha256:31ae7495831681dd3da34ac457f6970f147465ae4b9aae3a888d7a581c7cd868", size = 6170, upload-time = "2024-11-05T19:14:47.349Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b1/3f/c8ad4f1c3aaadcea2b0f1b4d7970e7b7898c145699769a789f3435143f69/opentelemetry_exporter_otlp-1.33.1.tar.gz", hash = "sha256:4d050311ea9486e3994575aa237e32932aad58330a31fba24fdba5c0d531cf04", size = 6189, upload-time = "2025-05-16T18:52:43.176Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c2/82/3f521b3c1f2a411ed60a24a8c9f486c1beeaf8c6c55337c87d3ae1642151/opentelemetry_exporter_otlp-1.28.0-py3-none-any.whl", hash = "sha256:1fd02d70f2c1b7ac5579c81e78de4594b188d3317c8ceb69e8b53900fb7b40fd", size = 7024, upload-time = "2024-11-05T19:14:24.534Z" }, + { url = "https://files.pythonhosted.org/packages/4d/32/b9add70dd4e845654fc9fcd1401a705477743880be6c3e62acb1ad0d8662/opentelemetry_exporter_otlp-1.33.1-py3-none-any.whl", hash = "sha256:9bcf1def35b880b55a49e31ebd63910edac14b294fd2ab884953c4deaff5b300", size = 7045, upload-time = "2025-05-16T18:52:21.022Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-common" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-proto" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c2/8d/5d411084ac441052f4c9bae03a1aec65ae5d16b439fea7b9c5ac3842c013/opentelemetry_exporter_otlp_proto_common-1.28.0.tar.gz", hash = "sha256:5fa0419b0c8e291180b0fc8430a20dd44a3f3236f8e0827992145914f273ec4f", size = 18505, upload-time = "2024-11-05T19:14:48.204Z" } +sdist = { url = "https://files.pythonhosted.org/packages/7a/18/a1ec9dcb6713a48b4bdd10f1c1e4d5d2489d3912b80d2bcc059a9a842836/opentelemetry_exporter_otlp_proto_common-1.33.1.tar.gz", hash = "sha256:c57b3fa2d0595a21c4ed586f74f948d259d9949b58258f11edb398f246bec131", size = 20828, upload-time = "2025-05-16T18:52:43.795Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e1/72/3c44aabc74db325aaba09361b6a0d80f6d601f0ff86ecea8ee655c9538fc/opentelemetry_exporter_otlp_proto_common-1.28.0-py3-none-any.whl", hash = "sha256:467e6437d24e020156dffecece8c0a4471a8a60f6a34afeda7386df31a092410", size = 18403, upload-time = "2024-11-05T19:14:25.798Z" }, + { url = "https://files.pythonhosted.org/packages/09/52/9bcb17e2c29c1194a28e521b9d3f2ced09028934c3c52a8205884c94b2df/opentelemetry_exporter_otlp_proto_common-1.33.1-py3-none-any.whl", hash = "sha256:b81c1de1ad349785e601d02715b2d29d6818aed2c809c20219f3d1f20b038c36", size = 18839, upload-time = "2025-05-16T18:52:22.447Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-grpc" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "deprecated" }, @@ -6201,14 +6202,14 @@ dependencies = [ { name = "opentelemetry-proto" }, { name = "opentelemetry-sdk" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/43/4d/f215162e58041afb4bdf5dbd0d8faf0b7fc9bf7b3d3fc0e44e06f9e7e869/opentelemetry_exporter_otlp_proto_grpc-1.28.0.tar.gz", hash = "sha256:47a11c19dc7f4289e220108e113b7de90d59791cb4c37fc29f69a6a56f2c3735", size = 26237, upload-time = "2024-11-05T19:14:49.026Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d8/5f/75ef5a2a917bd0e6e7b83d3fb04c99236ee958f6352ba3019ea9109ae1a6/opentelemetry_exporter_otlp_proto_grpc-1.33.1.tar.gz", hash = "sha256:345696af8dc19785fac268c8063f3dc3d5e274c774b308c634f39d9c21955728", size = 22556, upload-time = "2025-05-16T18:52:44.76Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1d/b5/afabc8106abc0f9cfeecf5b3e682622b3e04bba1d9b967dbfcd91b9c4ebe/opentelemetry_exporter_otlp_proto_grpc-1.28.0-py3-none-any.whl", hash = "sha256:edbdc53e7783f88d4535db5807cb91bd7b1ec9e9b9cdbfee14cd378f29a3b328", size = 18532, upload-time = "2024-11-05T19:14:26.853Z" }, + { url = "https://files.pythonhosted.org/packages/ba/ec/6047e230bb6d092c304511315b13893b1c9d9260044dd1228c9d48b6ae0e/opentelemetry_exporter_otlp_proto_grpc-1.33.1-py3-none-any.whl", hash = "sha256:7e8da32c7552b756e75b4f9e9c768a61eb47dee60b6550b37af541858d669ce1", size = 18591, upload-time = "2025-05-16T18:52:23.772Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-http" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "deprecated" }, @@ -6219,14 +6220,14 @@ dependencies = [ { name = "opentelemetry-sdk" }, { name = "requests" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f1/2a/555f2845928086cd51aa6941c7a546470805b68ed631ec139ce7d841763d/opentelemetry_exporter_otlp_proto_http-1.28.0.tar.gz", hash = "sha256:d83a9a03a8367ead577f02a64127d827c79567de91560029688dd5cfd0152a8e", size = 15051, upload-time = "2024-11-05T19:14:49.813Z" } +sdist = { url = "https://files.pythonhosted.org/packages/60/48/e4314ac0ed2ad043c07693d08c9c4bf5633857f5b72f2fefc64fd2b114f6/opentelemetry_exporter_otlp_proto_http-1.33.1.tar.gz", hash = "sha256:46622d964a441acb46f463ebdc26929d9dec9efb2e54ef06acdc7305e8593c38", size = 15353, upload-time = "2025-05-16T18:52:45.522Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b2/ce/80d5adabbf7ab4a0ca7b5e0f4039b24d273be370c3ba85fc05b13794411c/opentelemetry_exporter_otlp_proto_http-1.28.0-py3-none-any.whl", hash = "sha256:e8f3f7961b747edb6b44d51de4901a61e9c01d50debd747b120a08c4996c7e7b", size = 17228, upload-time = "2024-11-05T19:14:28.613Z" }, + { url = "https://files.pythonhosted.org/packages/63/ba/5a4ad007588016fe37f8d36bf08f325fe684494cc1e88ca8fa064a4c8f57/opentelemetry_exporter_otlp_proto_http-1.33.1-py3-none-any.whl", hash = "sha256:ebd6c523b89a2ecba0549adb92537cc2bf647b4ee61afbbd5a4c6535aa3da7cf", size = 17733, upload-time = "2025-05-16T18:52:25.137Z" }, ] [[package]] name = "opentelemetry-instrumentation" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6234,14 +6235,14 @@ dependencies = [ { name = "packaging" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/de/6b/6c25b15063c92a011cf3f68375971e2c58a9c764690847edc97df2d94eeb/opentelemetry_instrumentation-0.49b0.tar.gz", hash = "sha256:398a93e0b9dc2d11cc8627e1761665c506fe08c6b2df252a2ab3ade53d751c46", size = 26478, upload-time = "2024-11-05T19:21:41.402Z" } +sdist = { url = "https://files.pythonhosted.org/packages/c3/fd/5756aea3fdc5651b572d8aef7d94d22a0a36e49c8b12fcb78cb905ba8896/opentelemetry_instrumentation-0.54b1.tar.gz", hash = "sha256:7658bf2ff914b02f246ec14779b66671508125c0e4227361e56b5ebf6cef0aec", size = 28436, upload-time = "2025-05-16T19:03:22.223Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/93/61/e0d21e958d6072ce25c4f5e26a1d22835fc86f80836660adf6badb6038ce/opentelemetry_instrumentation-0.49b0-py3-none-any.whl", hash = "sha256:68364d73a1ff40894574cbc6138c5f98674790cae1f3b0865e21cf702f24dcb3", size = 30694, upload-time = "2024-11-05T19:20:38.584Z" }, + { url = "https://files.pythonhosted.org/packages/f4/89/0790abc5d9c4fc74bd3e03cb87afe2c820b1d1a112a723c1163ef32453ee/opentelemetry_instrumentation-0.54b1-py3-none-any.whl", hash = "sha256:a4ae45f4a90c78d7006c51524f57cd5aa1231aef031eae905ee34d5423f5b198", size = 31019, upload-time = "2025-05-16T19:02:15.611Z" }, ] [[package]] name = "opentelemetry-instrumentation-alephalpha" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6249,14 +6250,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/47/32/15048d7773f6018abcd5b85f5c346b44fad8322031f6b4ea5a6c5ada304a/opentelemetry_instrumentation_alephalpha-0.33.12.tar.gz", hash = "sha256:b474ac634cd1e12b30c8863a925320a01043af8c0f46fd58288e587073d6ddec", size = 3727, upload-time = "2024-11-13T20:27:50.425Z" } +sdist = { url = "https://files.pythonhosted.org/packages/64/12/b962c7fd3d29bc4ffe70f41fab8054d0221ebfecc28a344aef6fc749be67/opentelemetry_instrumentation_alephalpha-0.34.0.tar.gz", hash = "sha256:ed6647505963d53aed63b0b2ca84c989ca94ccc215ad19355a7de33e0b10f0ac", size = 3688, upload-time = "2024-12-12T21:02:01.771Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d3/77/e483e2fa14fddc87324b242d59992cfbd2d590563352aa044d11e1d200ed/opentelemetry_instrumentation_alephalpha-0.33.12-py3-none-any.whl", hash = "sha256:b3c7e3dd99121f5c52d7c7a3a82dd2d7a9ba7360f63ac6fcbdca187f58756e16", size = 5116, upload-time = "2024-11-13T20:27:12.818Z" }, + { url = "https://files.pythonhosted.org/packages/ef/1b/d37c9af6319ad64b182f77aec1154f5fab25b9123c9e04fe1a6d19e19e7e/opentelemetry_instrumentation_alephalpha-0.34.0-py3-none-any.whl", hash = "sha256:4e05e1b12edf30597e3cb6163d2e63f938fd3b061a3251940ac12783d1103ce6", size = 5101, upload-time = "2024-12-12T21:01:12.317Z" }, ] [[package]] name = "opentelemetry-instrumentation-anthropic" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6264,14 +6265,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/40/0a/cba0a6ac1832e3002158b5a9451268aebfe0150c7b8355068d1f2cea148b/opentelemetry_instrumentation_anthropic-0.33.12.tar.gz", hash = "sha256:0bc1fd9d4cf2feec4fe9f80c0bdfcbfab33ed9cf0edea850b6c198a8679b01ff", size = 8711, upload-time = "2024-11-13T20:27:52.005Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d1/56/57bbdb8907e14793d9831b220a6561a29033204acc60ce2ebc6387d29ad5/opentelemetry_instrumentation_anthropic-0.34.0.tar.gz", hash = "sha256:ab4336723de8cc3327aeacfab6e2fa085101f92614a402ee2822f8fb557ba7a6", size = 8693, upload-time = "2024-12-12T21:02:02.731Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ba/46/ba2dc8d18b04acae3d34facd8fe1e5e0cdc9fe64292d45eca9d1d4a8a298/opentelemetry_instrumentation_anthropic-0.33.12-py3-none-any.whl", hash = "sha256:b31618d12a429045db14ed982a142a25df0f0f1dbf03d756e8d597f25b9a053d", size = 11024, upload-time = "2024-11-13T20:27:14.622Z" }, + { url = "https://files.pythonhosted.org/packages/5c/8e/ef2782ecd3e2b03fb792f42ade5fea3c549ba28e6ebefdcf95a4c14412df/opentelemetry_instrumentation_anthropic-0.34.0-py3-none-any.whl", hash = "sha256:8fc397802033636eb74967ffc6a85344e575ea615b5de502386b0a004b07ba68", size = 11005, upload-time = "2024-12-12T21:01:13.846Z" }, ] [[package]] name = "opentelemetry-instrumentation-asgi" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "asgiref" }, @@ -6280,14 +6281,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/e8/55/693c3d0938ba5fead5c3aa4ac7022a992b4ff99a8e9979800d0feb843ff4/opentelemetry_instrumentation_asgi-0.49b0.tar.gz", hash = "sha256:959fd9b1345c92f20c6ef1d42f92ef6a76b3c3083fbc4104d59da6859b15b083", size = 24117, upload-time = "2024-11-05T19:21:46.769Z" } +sdist = { url = "https://files.pythonhosted.org/packages/20/f7/a3377f9771947f4d3d59c96841d3909274f446c030dbe8e4af871695ddee/opentelemetry_instrumentation_asgi-0.54b1.tar.gz", hash = "sha256:ab4df9776b5f6d56a78413c2e8bbe44c90694c67c844a1297865dc1bd926ed3c", size = 24230, upload-time = "2025-05-16T19:03:30.234Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/2c/0b/7900c782a1dfaa584588d724bc3bbdf8405a32497537dd96b3fcbf8461b9/opentelemetry_instrumentation_asgi-0.49b0-py3-none-any.whl", hash = "sha256:722a90856457c81956c88f35a6db606cc7db3231046b708aae2ddde065723dbe", size = 16326, upload-time = "2024-11-05T19:20:46.176Z" }, + { url = "https://files.pythonhosted.org/packages/20/24/7a6f0ae79cae49927f528ecee2db55a5bddd87b550e310ce03451eae7491/opentelemetry_instrumentation_asgi-0.54b1-py3-none-any.whl", hash = "sha256:84674e822b89af563b283a5283c2ebb9ed585d1b80a1c27fb3ac20b562e9f9fc", size = 16338, upload-time = "2025-05-16T19:02:22.808Z" }, ] [[package]] name = "opentelemetry-instrumentation-bedrock" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anthropic" }, @@ -6296,14 +6297,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/95/5a/346c17fca4dd929ce6be8cf402cd3580bb6e4da42ca8eadd2b6f2b4907e4/opentelemetry_instrumentation_bedrock-0.33.12.tar.gz", hash = "sha256:6f5a3f7044edff020d62b3e94f0ea543da4e5c23b7cdb72642692952843b0003", size = 7690, upload-time = "2024-11-13T20:27:53.497Z" } +sdist = { url = "https://files.pythonhosted.org/packages/aa/79/c384051d3e234ffb5f995ecb2245aef54083dc4919258601d9449c8c47bd/opentelemetry_instrumentation_bedrock-0.34.0.tar.gz", hash = "sha256:07f0ed84fa6d9e93c8cefee48ce171c59961c44708fcc11ec21fc1fbcdfb314d", size = 7695, upload-time = "2024-12-12T21:02:04.602Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d1/04/d93857519edd693e72e6d9ba08a6f0feda2ca21a08e3bc02cbfe242495f6/opentelemetry_instrumentation_bedrock-0.33.12-py3-none-any.whl", hash = "sha256:f9749898c52643d5027b45ac92bf4d3fd39b83adfaf68705a0ed9b4f04b8afae", size = 8982, upload-time = "2024-11-13T20:27:15.935Z" }, + { url = "https://files.pythonhosted.org/packages/56/2c/6d3e353d69407b308a254713728a613651bbe34138956f4f6b0104a5cc0a/opentelemetry_instrumentation_bedrock-0.34.0-py3-none-any.whl", hash = "sha256:1e521e33721e0fbcde2c2cb7cf788e2b8926063846800db777678be583bf1420", size = 8966, upload-time = "2024-12-12T21:01:16.457Z" }, ] [[package]] name = "opentelemetry-instrumentation-chromadb" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6311,14 +6312,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/62/05/ae78dd08c30203009815b35bce9b524458d73174e7dc4924a431e7b0b65b/opentelemetry_instrumentation_chromadb-0.33.12.tar.gz", hash = "sha256:eb4c591d398963504f82c20879030ea3694f10065ee62450da761c9b6e1792e7", size = 4598, upload-time = "2024-11-13T20:27:54.38Z" } +sdist = { url = "https://files.pythonhosted.org/packages/29/8e/0846e9c8846eee6f782767a1ee2f760ed5ca53cc95035189706c63027d58/opentelemetry_instrumentation_chromadb-0.34.0.tar.gz", hash = "sha256:ed0b4842db9bd35a0cff138d88d84d63a1529038ac11cf37eeba1dd294d4a2e8", size = 4596, upload-time = "2024-12-12T21:02:06.825Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/cc/2f/3fabec28e538fc0671c5c149e979e6f36f823078aa8c43ab0f69392185e1/opentelemetry_instrumentation_chromadb-0.33.12-py3-none-any.whl", hash = "sha256:2413426c3bf1f3714a95318e934f090fa778ab7b3d7bdd2cc8ee068cda216a06", size = 6322, upload-time = "2024-11-13T20:27:18.789Z" }, + { url = "https://files.pythonhosted.org/packages/9a/b6/132c1cdd8dea4f0e4e1cab910dabdacd9802fb3a8e802e0c825bf6e9691f/opentelemetry_instrumentation_chromadb-0.34.0-py3-none-any.whl", hash = "sha256:d95df8285405a23b82c3b6d0c1b7c439ec86793d21b3a23e51965853d3e9c4a6", size = 6303, upload-time = "2024-12-12T21:01:17.711Z" }, ] [[package]] name = "opentelemetry-instrumentation-cohere" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6326,14 +6327,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fe/96/f9cfc4f27c20deabfca237cb26734f060525b43e6993f753fad4ee0eded1/opentelemetry_instrumentation_cohere-0.33.12.tar.gz", hash = "sha256:4ea626d096fdf4c64e04a63b437e36f72a4341f818034ee6dc73ba1dba9ab341", size = 4235, upload-time = "2024-11-13T20:27:55.358Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b2/bc/d38c64d0e0f92fb8b8bcde241024dd5d4810c2fbe379fdfbbbb32dee957f/opentelemetry_instrumentation_cohere-0.34.0.tar.gz", hash = "sha256:80e27c6f86a73a2c0e89aa3c9ca1a37ff58a01b4c0eb7f249d7ae66568730477", size = 4227, upload-time = "2024-12-12T21:02:08.177Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/bf/08/dce2b7926ace0204ce7946563348e1ff755873e387833484791e4ed391c8/opentelemetry_instrumentation_cohere-0.33.12-py3-none-any.whl", hash = "sha256:3bee3f7f7105259c85145be8c3b68612421860c95ad170f4d03144a3b8c07418", size = 5589, upload-time = "2024-11-13T20:27:21.317Z" }, + { url = "https://files.pythonhosted.org/packages/e3/bb/5efa301486ad236777d15b515158224cb17ca4e1f138e1480ce8a9d5c369/opentelemetry_instrumentation_cohere-0.34.0-py3-none-any.whl", hash = "sha256:6238c84948d809ea5feb1ce603de2c8f9d72d7b8286d9f9115edf33b91202011", size = 5576, upload-time = "2024-12-12T21:01:20.273Z" }, ] [[package]] name = "opentelemetry-instrumentation-fastapi" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6342,14 +6343,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fe/bf/8e6d2a4807360f2203192017eb4845f5628dbeaf0597adf3d141cc5c24e1/opentelemetry_instrumentation_fastapi-0.49b0.tar.gz", hash = "sha256:6d14935c41fd3e49328188b6a59dd4c37bd17a66b01c15b0c64afa9714a1f905", size = 19230, upload-time = "2024-11-05T19:21:59.361Z" } +sdist = { url = "https://files.pythonhosted.org/packages/98/3b/9a262cdc1a4defef0e52afebdde3e8add658cc6f922e39e9dcee0da98349/opentelemetry_instrumentation_fastapi-0.54b1.tar.gz", hash = "sha256:1fcad19cef0db7092339b571a59e6f3045c9b58b7fd4670183f7addc459d78df", size = 19325, upload-time = "2025-05-16T19:03:45.359Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b1/f4/0895b9410c10abf987c90dee1b7688a8f2214a284fe15e575648f6a1473a/opentelemetry_instrumentation_fastapi-0.49b0-py3-none-any.whl", hash = "sha256:646e1b18523cbe6860ae9711eb2c7b9c85466c3c7697cd6b8fb5180d85d3fe6e", size = 12101, upload-time = "2024-11-05T19:21:01.805Z" }, + { url = "https://files.pythonhosted.org/packages/df/9c/6b2b0f9d6c5dea7528ae0bf4e461dd765b0ae35f13919cd452970bb0d0b3/opentelemetry_instrumentation_fastapi-0.54b1-py3-none-any.whl", hash = "sha256:fb247781cfa75fd09d3d8713c65e4a02bd1e869b00e2c322cc516d4b5429860c", size = 12125, upload-time = "2025-05-16T19:02:41.172Z" }, ] [[package]] name = "opentelemetry-instrumentation-google-generativeai" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6357,14 +6358,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/b0/39/d33585303893fec6d4e828b794b8b26188658e3bd905098e568031eb0698/opentelemetry_instrumentation_google_generativeai-0.33.12.tar.gz", hash = "sha256:9d09cd39afecf70063733b3f2f15200b7dc28addfa6384947a9514557f18d64b", size = 4302, upload-time = "2024-11-13T20:27:56.179Z" } +sdist = { url = "https://files.pythonhosted.org/packages/30/c8/4620090d09b3d450ac7069ad84b366b34c2488df290c7fc0af6582178812/opentelemetry_instrumentation_google_generativeai-0.34.0.tar.gz", hash = "sha256:b0ecc9cb840277d4040277158c4d77a48c171a64ca556c54ddc1c5ce5105ebd8", size = 4288, upload-time = "2024-12-12T21:02:10.891Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/11/9a/622ca1552d05b5b948c1f0e78d8456464697118c63a58d4f5aa01c105d45/opentelemetry_instrumentation_google_generativeai-0.33.12-py3-none-any.whl", hash = "sha256:0dcd71c38331c47663d7ba6237ddfe02c14e3d1e3a47524c1437e3ee56cd0036", size = 5889, upload-time = "2024-11-13T20:27:22.37Z" }, + { url = "https://files.pythonhosted.org/packages/37/6e/e20b5fce0020a1f3de78227610a7d018764e9b26e8766c4f948493c2485e/opentelemetry_instrumentation_google_generativeai-0.34.0-py3-none-any.whl", hash = "sha256:eb42d8d48e3d13e03363932b69f424d27e8d9a53c8cbd23f190c4a294a881edc", size = 5879, upload-time = "2024-12-12T21:01:22.735Z" }, ] [[package]] name = "opentelemetry-instrumentation-groq" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6372,14 +6373,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fa/39/71c87d595a312e2cfef83b006070e5d56c73895c59b596336a20aa43e79c/opentelemetry_instrumentation_groq-0.33.12.tar.gz", hash = "sha256:1460901e66c87b47198d639fb22ec25552281cdf7cafe13ae9605447661d6871", size = 5687, upload-time = "2024-11-13T20:28:00.703Z" } +sdist = { url = "https://files.pythonhosted.org/packages/af/1d/443944e52fc37a5e564525134dd86bee6a3f2db7be1c08f6459f056965ad/opentelemetry_instrumentation_groq-0.34.0.tar.gz", hash = "sha256:0c9162ce1a7b5b5a613dbf50f5f2ee8d5e6e175e0cc1758d53d71cb22c7aac1b", size = 5670, upload-time = "2024-12-12T21:02:12.256Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f7/24/631269741eabb0b028a313f15063871b28e700ba27feece154a3dd71f62d/opentelemetry_instrumentation_groq-0.33.12-py3-none-any.whl", hash = "sha256:4d239c73d689c046ab2c90a25b78d6c7406cef1e26f04633bc148464b66cc74c", size = 7270, upload-time = "2024-11-13T20:27:23.508Z" }, + { url = "https://files.pythonhosted.org/packages/a3/81/beb464fdd0d3f568b589b45629f74e0fb1a1e518a9df2f575bb68ea2096a/opentelemetry_instrumentation_groq-0.34.0-py3-none-any.whl", hash = "sha256:0f74c8b0df2984b27aadabebf3bed4443c0db45fbd851f956299207be12bb207", size = 7252, upload-time = "2024-12-12T21:01:24.069Z" }, ] [[package]] name = "opentelemetry-instrumentation-haystack" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6387,14 +6388,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/9c/06/067e4b2db2bc29d0a7e3a6cc8676d5f1971b0ecbaf7e5fa0c1e478e092af/opentelemetry_instrumentation_haystack-0.33.12.tar.gz", hash = "sha256:3d45df14aff1f2321066e55ecce632653d67c36249d3eaccbefa189f0daaba05", size = 4663, upload-time = "2024-11-13T20:28:01.551Z" } +sdist = { url = "https://files.pythonhosted.org/packages/c3/61/2aa5d850c1891fd99636a4ad724489ed792ac4aa560be75ab34af0ee26eb/opentelemetry_instrumentation_haystack-0.34.0.tar.gz", hash = "sha256:29739e9429a1a327dc72f743a0b37a3b7f26a742ac762791a75b1bc2f3ba43ff", size = 4645, upload-time = "2024-12-12T21:02:13.193Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/dc/ba/b8872dce7eb67bd6589d4cccfc97554851dad719067b24c7697a5ffef69a/opentelemetry_instrumentation_haystack-0.33.12-py3-none-any.whl", hash = "sha256:d2a3041a58e1027d8728e1a24430f7b01e8c27e9c04c22fa4204c590b12a95f6", size = 7513, upload-time = "2024-11-13T20:27:24.671Z" }, + { url = "https://files.pythonhosted.org/packages/58/15/682dfc4717e4ddbb668fdcb5a12a8b22a2f6c9402d78c26690528722e8e5/opentelemetry_instrumentation_haystack-0.34.0-py3-none-any.whl", hash = "sha256:2ae56f4abc7a2bafad7b2b3ec8e218edf2aa0daaa6570c692076c24682ee78ce", size = 7495, upload-time = "2024-12-12T21:01:25.723Z" }, ] [[package]] name = "opentelemetry-instrumentation-lancedb" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6402,14 +6403,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/38/52/16eef8e5c92627a82904f0112715ee95b2a6ee74ec958ec69f7db803a8d5/opentelemetry_instrumentation_lancedb-0.33.12.tar.gz", hash = "sha256:0aa9f6319374f532e2087949c15674f4d8036591ba70f91f9ce6996ea34508c3", size = 3198, upload-time = "2024-11-13T20:28:02.455Z" } +sdist = { url = "https://files.pythonhosted.org/packages/05/00/ad6383e2981308146e282da4d977ed61dd63321c6ac72751aa1c9eb26d74/opentelemetry_instrumentation_lancedb-0.34.0.tar.gz", hash = "sha256:5d081f36335d7b5dd3a8ae3b0fac0b895f4284941e3521f32332d3393b3b1178", size = 3185, upload-time = "2024-12-12T21:02:14.193Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/26/1d/218b74341471aa370999f7dbee09de44e32d8f822d69b6ca6d95f400be36/opentelemetry_instrumentation_lancedb-0.33.12-py3-none-any.whl", hash = "sha256:e1cdd55ef38d939d8af924478486e66c0cf65a7e5ac19c82f20ad5d06e682b9d", size = 4794, upload-time = "2024-11-13T20:27:25.825Z" }, + { url = "https://files.pythonhosted.org/packages/5b/65/db706f845a5ab861ee59feb6eb394843e27bf0f025bc1438e44a7af19f19/opentelemetry_instrumentation_lancedb-0.34.0-py3-none-any.whl", hash = "sha256:b8284453cb3d98fbe83bd286448eca4edbb779fc79ffc58bdb3a344137d82719", size = 4780, upload-time = "2024-12-12T21:01:26.96Z" }, ] [[package]] name = "opentelemetry-instrumentation-langchain" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6417,14 +6418,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/73/60/4fb638bc69cc63bbf7aad81a08650c99bd343a67f49c532261190e7ee7e4/opentelemetry_instrumentation_langchain-0.33.12.tar.gz", hash = "sha256:ff607742c76a1844211648415fa35da9eac22a42da2a9732c673bd13f2973994", size = 8518, upload-time = "2024-11-13T20:28:03.247Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ce/b4/a8bafbc727a874eb26e788b1fe667db85c7dbcb6d685a6b9da07f6ba231b/opentelemetry_instrumentation_langchain-0.34.0.tar.gz", hash = "sha256:2a25bc07ff8719d30b9a01acf29305c7de5418683c14334ad7ddef4608222911", size = 8508, upload-time = "2024-12-12T21:02:15.138Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b1/fe/215a5b5b52360c94b2223f4dfb339665a60838d8f53edd37dd1c40897de4/opentelemetry_instrumentation_langchain-0.33.12-py3-none-any.whl", hash = "sha256:7406ab7116fa43343f53602f7b530f9bb1552e20ad750dbfc5aa1761c027c2d3", size = 9749, upload-time = "2024-11-13T20:27:27.623Z" }, + { url = "https://files.pythonhosted.org/packages/a1/3f/01f5d6e5fc3e34e068b6ad650bde73facb9748df74de305a33702ad06820/opentelemetry_instrumentation_langchain-0.34.0-py3-none-any.whl", hash = "sha256:373c69adcf18e9d37cd47d96fad78c57959c3f8af7034aff50553103fdbf0ba8", size = 9734, upload-time = "2024-12-12T21:01:28.244Z" }, ] [[package]] name = "opentelemetry-instrumentation-llamaindex" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "inflection" }, @@ -6433,27 +6434,27 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/df/e7/9b9d43c7b78eea5ecf95b378ed4c12850f8745ff5b46ac5fa9042c89c941/opentelemetry_instrumentation_llamaindex-0.33.12.tar.gz", hash = "sha256:7a278dfe21fbba7dd1b8fe824c9baee0bfb3b4f7ccd71aae5f677412be45587e", size = 9285, upload-time = "2024-11-13T20:28:04.082Z" } +sdist = { url = "https://files.pythonhosted.org/packages/06/fe/b73490ee120672c81f78209a787feb1a5fbf19f2ec0657cf9b85277597ae/opentelemetry_instrumentation_llamaindex-0.34.0.tar.gz", hash = "sha256:f84eaa198873e856401fd8382f86d4f099e8edd712369579f4b74c24e0404933", size = 9274, upload-time = "2024-12-12T21:02:17.055Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d6/6a/aef813dff690cf06c62a86bf3f723ab1331b55f7c3dfd56690e93ac993df/opentelemetry_instrumentation_llamaindex-0.33.12-py3-none-any.whl", hash = "sha256:7f0d0700015f1e1576cf2de211a4062ae2d0ea899c298f4d7d44ce2a37226135", size = 16372, upload-time = "2024-11-13T20:27:28.724Z" }, + { url = "https://files.pythonhosted.org/packages/71/38/01d81a1bae3965031d612afe162a4fdee40d181b46d0f2aab7d2ac49d015/opentelemetry_instrumentation_llamaindex-0.34.0-py3-none-any.whl", hash = "sha256:0058a44a584ccb9046bed3d5da7bb64160c51d46f0a9946d2bb6517ffcd29fd0", size = 16354, upload-time = "2024-12-12T21:01:32.806Z" }, ] [[package]] name = "opentelemetry-instrumentation-logging" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-instrumentation" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c8/80/1d15f8afebc2b67ed47bfe45ee97c042808441586617d5aea8df1f1cbd96/opentelemetry_instrumentation_logging-0.49b0.tar.gz", hash = "sha256:d8058216b06c029785113a71428c6edbb3f0e3b9f69ee917050cb98cd8137fb2", size = 9731, upload-time = "2024-11-05T19:22:05.252Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d9/5b/88ed39f22e8c6eb4f6192ab9a62adaa115579fcbcadb3f0241ee645eea56/opentelemetry_instrumentation_logging-0.54b1.tar.gz", hash = "sha256:893a3cbfda893b64ff71b81991894e2fd6a9267ba85bb6c251f51c0419fbe8fa", size = 9976, upload-time = "2025-05-16T19:03:49.976Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a7/c4/0eedcf9ccce07a64baa002fae7001d84f2c032cf5b2ff1a9438bff0479dd/opentelemetry_instrumentation_logging-0.49b0-py3-none-any.whl", hash = "sha256:9f9405d2f8e6fd756d49da979710f7b5ba1b95bd534467f176aae756102eed58", size = 12150, upload-time = "2024-11-05T19:21:07.826Z" }, + { url = "https://files.pythonhosted.org/packages/96/0c/b441fb30d860f25040eaed61e89d68f4d9ee31873159ed18cbc1b92eba56/opentelemetry_instrumentation_logging-0.54b1-py3-none-any.whl", hash = "sha256:01a4cec54348f13941707d857b850b0febf9d49f45d0fcf0673866e079d7357b", size = 12579, upload-time = "2025-05-16T19:02:49.039Z" }, ] [[package]] name = "opentelemetry-instrumentation-marqo" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6461,14 +6462,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/82/fb/775a1a5f9b9f641b3c7aa7ea8a7d83cbb97a4ddf86e1c5a4dd2a3c42af8d/opentelemetry_instrumentation_marqo-0.33.12.tar.gz", hash = "sha256:802def00b35033055618dc137f81895496bb449ff405580ff9414eeda139b89f", size = 3479, upload-time = "2024-11-13T20:28:04.892Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e7/50/2585b0d337a15b7fe31ec0e7245c09154b7d3d7e7e172c3973d40ad313ee/opentelemetry_instrumentation_marqo-0.34.0.tar.gz", hash = "sha256:7bcc091b89717ac7b04c224dfc1429f200ba2b3e930d7a4de80bf9bc054fc0db", size = 3471, upload-time = "2024-12-12T21:02:17.979Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5b/b1/153592356d8cc6faab61f7e5c727ab3c382984cede3fbdc228872ba5ac7c/opentelemetry_instrumentation_marqo-0.33.12-py3-none-any.whl", hash = "sha256:6f939532f1f953a22eb2811dfdb439bb8bd813182d59d3444dcc4eca46d0805f", size = 5091, upload-time = "2024-11-13T20:27:31.133Z" }, + { url = "https://files.pythonhosted.org/packages/d4/51/9be4f5df62db6ff6e786933136541e3534656bb49a7d35324d96a5c07818/opentelemetry_instrumentation_marqo-0.34.0-py3-none-any.whl", hash = "sha256:dd342cfd4b70d4f65830708bc253397734d3da51d3773b677693e9007217e3ed", size = 5077, upload-time = "2024-12-12T21:01:35.643Z" }, ] [[package]] name = "opentelemetry-instrumentation-milvus" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6476,14 +6477,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/58/be/d86538d7b09c6ed77f137229ba418e402659e4aa9268ffff696e059ec8fb/opentelemetry_instrumentation_milvus-0.33.12.tar.gz", hash = "sha256:8720e8fd29ea3009dd0e5b8849b1d24657a5750d81f6c1af605e8aabe827f742", size = 3666, upload-time = "2024-11-13T20:28:06.004Z" } +sdist = { url = "https://files.pythonhosted.org/packages/3b/a8/18725e95cb5cf0d01c001698aa4b01198b5f205da4bb98acc8e0b0da1b6c/opentelemetry_instrumentation_milvus-0.34.0.tar.gz", hash = "sha256:6c19aa93c392f5c736390320b27e035761f36ead904277943d62ac6662d77f83", size = 3657, upload-time = "2024-12-12T21:02:18.849Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/6d/22/0972d94433358624e8f956228763650a2366921d4659035b260aad9775aa/opentelemetry_instrumentation_milvus-0.33.12-py3-none-any.whl", hash = "sha256:608783fa555aded64606cfda64f4aa6ace5c2f9790b2b920e114ca229ab00915", size = 5311, upload-time = "2024-11-13T20:27:32.219Z" }, + { url = "https://files.pythonhosted.org/packages/7b/fb/685282b0e0339d629d4fdb03af7d2904461f20c128ae88fa148847e8664a/opentelemetry_instrumentation_milvus-0.34.0-py3-none-any.whl", hash = "sha256:4c587c6031bc82d78189b31f6acd4f36a62ce3ae1f2b18bc7fabb667912cc2d7", size = 5294, upload-time = "2024-12-12T21:01:38.115Z" }, ] [[package]] name = "opentelemetry-instrumentation-mistralai" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6491,14 +6492,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f6/78/5cab3468d3885cc391f67ef0a4a845abb085aa8d786aaa565b9cd9c6612f/opentelemetry_instrumentation_mistralai-0.33.12.tar.gz", hash = "sha256:c22f7006a56180ab6384e47b4e49bde8597833f73955e48ac323cdbe107f06ae", size = 4383, upload-time = "2024-11-13T20:28:09.179Z" } +sdist = { url = "https://files.pythonhosted.org/packages/49/0e/3d86aa6b5a31a20ecadbd4423e83255fd09d9648f431ccc786665e0f98be/opentelemetry_instrumentation_mistralai-0.34.0.tar.gz", hash = "sha256:7c81d8602a16b37d698002a7b06233095fe5c17ddf2f0b9d973b78255cdf7547", size = 4387, upload-time = "2024-12-12T21:02:19.71Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/12/c8/f47d404273e7130f4e4360f33a093d73f8fbbf232ef536a761baeb91027c/opentelemetry_instrumentation_mistralai-0.33.12-py3-none-any.whl", hash = "sha256:66c5961a33492aaf4420ed1ab8c63e533162a06366108c3398c461d05b2a1154", size = 5858, upload-time = "2024-11-13T20:27:33.251Z" }, + { url = "https://files.pythonhosted.org/packages/16/c8/5644b1a821b60a34bebc58f96367571bcdcdf5ab1522137e13ae3a936360/opentelemetry_instrumentation_mistralai-0.34.0-py3-none-any.whl", hash = "sha256:f682b8d4011124fa326308e8fc4ce9e9fdbacfc72fc77431682d2ef950e636d8", size = 5842, upload-time = "2024-12-12T21:01:39.204Z" }, ] [[package]] name = "opentelemetry-instrumentation-ollama" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6506,14 +6507,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/d9/aa/e9c0f903b8ae750794688f82a00ac0b7ab00a42e57d7c279738b5a80a0ce/opentelemetry_instrumentation_ollama-0.33.12.tar.gz", hash = "sha256:4cd012503f8d692453645353231e216c756fc926bdd3142e8c97fc8e87cbe06f", size = 4512, upload-time = "2024-11-13T20:28:10.969Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ff/f2/4c8e16bb5b13a85d86f6a3c515bee05051dcc08f8023dd201a69c9c0580f/opentelemetry_instrumentation_ollama-0.34.0.tar.gz", hash = "sha256:c9cabfac35945eb9b167f174a9fcafe82ec7c70ae1ba04d462486d2ef4c20f70", size = 4491, upload-time = "2024-12-12T21:02:20.656Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/30/58/5f11976bc5fde11390e709d96c757e1b21ff73faf88d5b0fd97ace56b061/opentelemetry_instrumentation_ollama-0.33.12-py3-none-any.whl", hash = "sha256:ed5313f45f5d46e17d93096eba91ed6338abf535cf2d2eca648d0e2d621e9d6b", size = 5847, upload-time = "2024-11-13T20:27:35.323Z" }, + { url = "https://files.pythonhosted.org/packages/16/3b/1547e92c76b9dd3097a98a67a84b0641ecbba3348f07d5825ecfa3433c7d/opentelemetry_instrumentation_ollama-0.34.0-py3-none-any.whl", hash = "sha256:17beea413c78be8510409aa4b5a5f909ba9e9d14799fd6372b16d84cefb21120", size = 5832, upload-time = "2024-12-12T21:01:40.792Z" }, ] [[package]] name = "opentelemetry-instrumentation-openai" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6522,14 +6523,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions-ai" }, { name = "tiktoken" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/10/ec/2f9bb0a22ba916c10b2ef63ccde48f17348c49c2a651b8590a94076308e8/opentelemetry_instrumentation_openai-0.33.12.tar.gz", hash = "sha256:2c6dfd74d9d56ca393f9dbfc92883c7397d63408ff18b3d9a774ea1611a48ed9", size = 14631, upload-time = "2024-11-13T20:28:11.767Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b2/9a/04bb865c14d44111ccde056ffa994d5c29ee604c286d6248b2365f53d676/opentelemetry_instrumentation_openai-0.34.0.tar.gz", hash = "sha256:67fabd6b178837c3d115296654a0daaebeeec763789e3f7ffd9a3db6117b354e", size = 14967, upload-time = "2024-12-12T21:02:22.935Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c5/21/c4a2b70e9f3487ba7123fde8c55090ff4a3f227477261fac1d0b73d6d349/opentelemetry_instrumentation_openai-0.33.12-py3-none-any.whl", hash = "sha256:d5d0c83a469dbf7ab97d1c482ce78f7ba23c00015b01bc8be43cdc0e5d7c497f", size = 22089, upload-time = "2024-11-13T20:27:36.375Z" }, + { url = "https://files.pythonhosted.org/packages/fe/08/6b3c0404d53a2ca913a98fecb4228be9238965cf2f9092acf5c3e960cba0/opentelemetry_instrumentation_openai-0.34.0-py3-none-any.whl", hash = "sha256:22e902b1b830ca53a0a94ec523880a4d39a210e4ec34d0ce76605b726eed1aab", size = 22597, upload-time = "2024-12-12T21:01:43.669Z" }, ] [[package]] name = "opentelemetry-instrumentation-pinecone" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6537,14 +6538,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/9a/2f/f9e0d1d5eb6f3eee791dbeac13185f4b77bbdce774a01011e03a5d8e2a71/opentelemetry_instrumentation_pinecone-0.33.12.tar.gz", hash = "sha256:92ed3221bddb061ebe7f50cd4804c76c9f5d019e2afb967858c1be62c1a3ebf2", size = 4651, upload-time = "2024-11-13T20:28:14.402Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0e/a0/35791b83b78157f4bfb2f7d9c356a1f42a6ffc30f17f2e96210dabe090f7/opentelemetry_instrumentation_pinecone-0.34.0.tar.gz", hash = "sha256:573483686da9fd2be48c6de870b87515e479d6ee489ebc471d7c90e0de4106e1", size = 4649, upload-time = "2024-12-12T21:02:23.857Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/32/36/8458e916a2cca0378ac28e6d70347d5263160600fcb47b131f029bdd390a/opentelemetry_instrumentation_pinecone-0.33.12-py3-none-any.whl", hash = "sha256:8d6185bd2f5bf34f3983cad48bbaa86fc72925c02a07ab4a79eb805401556b19", size = 6377, upload-time = "2024-11-13T20:27:38.177Z" }, + { url = "https://files.pythonhosted.org/packages/9b/39/0f09e3de4fa72f17438a1ab7f2a8693a9c983a1e1ad6bf6e72c590616e4a/opentelemetry_instrumentation_pinecone-0.34.0-py3-none-any.whl", hash = "sha256:d81387e703bfd59ff03ed46acf0bf6a0c11cf4cdec16a440acbdea18987fda71", size = 6363, upload-time = "2024-12-12T21:01:46.926Z" }, ] [[package]] name = "opentelemetry-instrumentation-qdrant" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6552,14 +6553,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1e/0a/216d28c48dc8b9e37094c54cd1e3ef9a43609fc25b2eeb8bd939f0026047/opentelemetry_instrumentation_qdrant-0.33.12.tar.gz", hash = "sha256:ba34c6863c652f27ae28b9922b25d77617ece4a2233ad0c9c1ebd257605853a5", size = 3988, upload-time = "2024-11-13T20:28:18.929Z" } +sdist = { url = "https://files.pythonhosted.org/packages/da/03/bbf02439ba6c6077ac814957695846bc44edc4e631fce8f0cbb792aa5572/opentelemetry_instrumentation_qdrant-0.34.0.tar.gz", hash = "sha256:8d569b2d7ac70bbf7e75abe5f572ff9576fa175660d1ecdc98f60c6ea1d7010b", size = 3977, upload-time = "2024-12-12T21:02:24.795Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d2/dc/caa6f4951c84ecb11594ab8545a8692c8a166d618a67687e7f07694dcd97/opentelemetry_instrumentation_qdrant-0.33.12-py3-none-any.whl", hash = "sha256:e759fe49c67092197eaa547570d67e15f30675bfdc772839d0456d535297d4c0", size = 6317, upload-time = "2024-11-13T20:27:39.638Z" }, + { url = "https://files.pythonhosted.org/packages/2f/fe/1797190a4a6b81b50a5c83615590929a21775bd4cd8a855102703243fd51/opentelemetry_instrumentation_qdrant-0.34.0-py3-none-any.whl", hash = "sha256:34ef85e62f3039a2b61a68c6decdb9e5d05ab9f7303d08fc010d6d5ec8f144e0", size = 6302, upload-time = "2024-12-12T21:01:48.042Z" }, ] [[package]] name = "opentelemetry-instrumentation-replicate" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6567,14 +6568,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/89/1e/b4260513c9526b2113772e4bf10bee714555abfb6e7cf69f86326e5607d5/opentelemetry_instrumentation_replicate-0.33.12.tar.gz", hash = "sha256:5dafad1a7a20ba762f689f30c4f76bcb3817b617adb7da3288ac545d15a14565", size = 3767, upload-time = "2024-11-13T20:28:19.766Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0e/98/a36a16df6876396962071d8fbc6f0dc1c81bc1fdcb324e72871683b83e51/opentelemetry_instrumentation_replicate-0.34.0.tar.gz", hash = "sha256:124796ff8593cd211bfa05773f70e8f087a8c0522a544be39bed212a95c8dec3", size = 3767, upload-time = "2024-12-12T21:02:25.678Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/4d/6d/20dab686dce1ce491c97ea4a0feec86b9a2207891a2627b1f4a670d109b5/opentelemetry_instrumentation_replicate-0.33.12-py3-none-any.whl", hash = "sha256:dc527f470080248a57b738b63ed29eae2d82ef68a100fe6e3548f5de7677f2ed", size = 5189, upload-time = "2024-11-13T20:27:40.792Z" }, + { url = "https://files.pythonhosted.org/packages/69/4b/cb70ab819ec045c2e494deea99a542e4819c9d1c5f09ec99d6dffacb00ad/opentelemetry_instrumentation_replicate-0.34.0-py3-none-any.whl", hash = "sha256:c5f3d712702f3cbcfde619d08e83b1c2fd70e4ad36190d68575d576e27370c4d", size = 5175, upload-time = "2024-12-12T21:01:49.409Z" }, ] [[package]] name = "opentelemetry-instrumentation-requests" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6582,14 +6583,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c1/16/c71196d8f4cac30b6936c77567ae769f44ac97227255627f5277d825277d/opentelemetry_instrumentation_requests-0.49b0.tar.gz", hash = "sha256:b75a282b3641547272dc7d2fdc0dd68269d0c1e685e4d17579b7fbd34c19b6bb", size = 14123, upload-time = "2024-11-05T19:22:14.128Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e7/45/116da84930d3dc2f5cdd876283ca96e9b96547bccee7eaa0bd01ce6bf046/opentelemetry_instrumentation_requests-0.54b1.tar.gz", hash = "sha256:3eca5d697c5564af04c6a1dd23b6a3ffbaf11e64887c6051655cee03998f4654", size = 15148, upload-time = "2025-05-16T19:04:00.488Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/79/33/4b8a4a839401290c44c65a8ca926a60a86c5ee3ecdcf54de4575c288b5ac/opentelemetry_instrumentation_requests-0.49b0-py3-none-any.whl", hash = "sha256:bb39803359e226b8eb0d4c8aaba6fd8a883a7f869fc331ff861743173b33d26d", size = 12368, upload-time = "2024-11-05T19:21:22.387Z" }, + { url = "https://files.pythonhosted.org/packages/2b/b1/6e33d2c3d3cc9e3ae20a9a77625ec81a509a0e5d7fa87e09e7f879468990/opentelemetry_instrumentation_requests-0.54b1-py3-none-any.whl", hash = "sha256:a0c4cd5d946224f336d6bd73cdabdecc6f80d5c39208f84eb96eb15f16cd41a0", size = 12968, upload-time = "2025-05-16T19:03:03.131Z" }, ] [[package]] name = "opentelemetry-instrumentation-sagemaker" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6597,14 +6598,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c2/97/2ddbceba0f95f9b28e7ed75ee2d38214ea1e7d071585d9afea35d0e71619/opentelemetry_instrumentation_sagemaker-0.33.12.tar.gz", hash = "sha256:286bb0e7765967212e111274ca523084d8105a3f18d1dfc90873bca60f6ad766", size = 4508, upload-time = "2024-11-13T20:28:21.829Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e6/98/fc4a33b8800a4a774c50a4b3da83a95359b89b3744c259b2ebbafc3b2f3a/opentelemetry_instrumentation_sagemaker-0.34.0.tar.gz", hash = "sha256:b7c2be5ba9ea4f4b9705705cddc1c3474cf4cb4e6db9fdf6968ad97ec8e6f1df", size = 4506, upload-time = "2024-12-12T21:02:26.652Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c3/f3/21896275fb1b4954082c4b95277d8ce66d6e947f1c153b0f01182e9a852f/opentelemetry_instrumentation_sagemaker-0.33.12-py3-none-any.whl", hash = "sha256:da72e78a094106c3ce48e2410665016f161211976c577e92b4624dfbbc54e47e", size = 6296, upload-time = "2024-11-13T20:27:41.815Z" }, + { url = "https://files.pythonhosted.org/packages/89/c2/b60f211e51b3c8346073dde33e7053ba1027b943da50d34ec6f00afe7d78/opentelemetry_instrumentation_sagemaker-0.34.0-py3-none-any.whl", hash = "sha256:ed7a50a5a863bfc81bc792fd3bc7b33bbf0af9e279b6e527c79e93034deda1a0", size = 6282, upload-time = "2024-12-12T21:01:50.527Z" }, ] [[package]] name = "opentelemetry-instrumentation-sqlalchemy" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6613,28 +6614,28 @@ dependencies = [ { name = "packaging" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/a0/a7/24f6cce3808ae1802dd1b60d752fbab877db5655198929cf4ee8ea416923/opentelemetry_instrumentation_sqlalchemy-0.49b0.tar.gz", hash = "sha256:32658e520fc8b35823c722f5d8831d3a410b76dd2724adb2887befc041ddef04", size = 13194, upload-time = "2024-11-05T19:22:14.92Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ac/33/78a25ae4233d42058bb0b363ba4fea7d7210e53c24e5e31f16d5cf6cf957/opentelemetry_instrumentation_sqlalchemy-0.54b1.tar.gz", hash = "sha256:97839acf1c9b96ded857fca57a09b86a56cf8d9eb6d706b7ceaee9352a460e03", size = 14620, upload-time = "2025-05-16T19:04:01.215Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ec/6b/a1a3685fed593282999cdc374ece15efbd56f8d774bd368bf7ff2cf5923c/opentelemetry_instrumentation_sqlalchemy-0.49b0-py3-none-any.whl", hash = "sha256:d854052d2b02cd0562e5628a514c8153fceada7f585137e173165dfd0a46ef6a", size = 13358, upload-time = "2024-11-05T19:21:23.654Z" }, + { url = "https://files.pythonhosted.org/packages/c7/2b/1c954885815614ef5c1e8c7bbf57a5275e64cd6fb5946b65e17162a34037/opentelemetry_instrumentation_sqlalchemy-0.54b1-py3-none-any.whl", hash = "sha256:d2ca5edb4c7ecef120d51aad6793b7da1cc80207ccfd31c437ee18f098e7c4c4", size = 14169, upload-time = "2025-05-16T19:03:04.119Z" }, ] [[package]] name = "opentelemetry-instrumentation-threading" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-instrumentation" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/80/88/b19f064ebf1650a7291cb7fcb623129997a7d8af603ffe7cd1907fe469ba/opentelemetry_instrumentation_threading-0.49b0.tar.gz", hash = "sha256:b65ec668a3ee73fccb1432edf52556f374cb9d9e5b160a6da3a6f67890adf444", size = 8283, upload-time = "2024-11-05T19:22:18.778Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a0/bd/561245292e7cc78ac7a0a75537873aea87440cb9493d41371421b3308c2b/opentelemetry_instrumentation_threading-0.54b1.tar.gz", hash = "sha256:3a081085b59675baf7bd93126a681903e6304a5f283df5eaecdd44bcb66df578", size = 8774, upload-time = "2025-05-16T19:04:04.482Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e1/77/cf262caae1a8903bbe9c379dc6908fddc9f7bbd5c51866d7c7fbae2edb70/opentelemetry_instrumentation_threading-0.49b0-py3-none-any.whl", hash = "sha256:47a49931a2244c2b17db985c512e6c922328b891ff2b64d37b0cd3bd00fd00a9", size = 9072, upload-time = "2024-11-05T19:21:29.564Z" }, + { url = "https://files.pythonhosted.org/packages/81/10/d87ec07d69546adaad525ba5d40d27324a45cba29097d9854a53d9af5047/opentelemetry_instrumentation_threading-0.54b1-py3-none-any.whl", hash = "sha256:bc229e6cd3f2b29fafe0a8dd3141f452e16fcb4906bca4fbf52609f99fb1eb42", size = 9314, upload-time = "2025-05-16T19:03:09.527Z" }, ] [[package]] name = "opentelemetry-instrumentation-together" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6642,14 +6643,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/7f/5f/63ee7efc3de97e12eadc927ac4079caef454ee41f8e608bf7a83734024a9/opentelemetry_instrumentation_together-0.33.12.tar.gz", hash = "sha256:4ac8676560e93492bdd0540d67672424e26f1eb9a41a266d5248eb09b00dc4d2", size = 3907, upload-time = "2024-11-13T20:28:22.971Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1c/b5/f306f884fd775aff4195ca1d8c7a1426829cd7060ff50b293be23e1ad869/opentelemetry_instrumentation_together-0.34.0.tar.gz", hash = "sha256:f8968d2aaae123e556e9bd7ce9213f40888a180a8014382bb738cff0bc8de8a1", size = 3907, upload-time = "2024-12-12T21:02:28.939Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e4/6a/56f3a5abea0a3086d24a32398231e0a578c01b7fad174cd69cdd644fb87d/opentelemetry_instrumentation_together-0.33.12-py3-none-any.whl", hash = "sha256:6a1941e3d02b1505bd79a1ef3540d1fe15bf4f61c72cc162445d25d8715386b3", size = 5284, upload-time = "2024-11-13T20:27:42.798Z" }, + { url = "https://files.pythonhosted.org/packages/78/82/32bc20923c9ecd4495a01bc2dcabd377e4fd82c9cf334ed3ad3a81afaf02/opentelemetry_instrumentation_together-0.34.0-py3-none-any.whl", hash = "sha256:9b5069c3a294c161d8ad638a6d234484a2c600f77260902fb8e15afdd8dfdd33", size = 5267, upload-time = "2024-12-12T21:01:51.576Z" }, ] [[package]] name = "opentelemetry-instrumentation-transformers" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6657,14 +6658,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/3f/25/f73bfae73a466b0ee206168a0ec2212db694132f0af75ff7bf8d7da74488/opentelemetry_instrumentation_transformers-0.33.12.tar.gz", hash = "sha256:d7b9c0d4bd71b834a79c2522455799feb7e76148e1dd371408e9907e847e8d6a", size = 3714, upload-time = "2024-11-13T20:28:24.399Z" } +sdist = { url = "https://files.pythonhosted.org/packages/02/54/1ab4fb5409cf6c48f7b0c0a48b39cbea70b4083d994719ba0975ba9a9580/opentelemetry_instrumentation_transformers-0.34.0.tar.gz", hash = "sha256:586b146509a90900486039850f5f3d63256c7f1546e1a897912ba454aa14e5af", size = 3714, upload-time = "2024-12-12T21:02:29.913Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/12/ef/4305bdf6af7161c2b38d5341b25bb2a0817f8ba40bf370bb6eeb131a223c/opentelemetry_instrumentation_transformers-0.33.12-py3-none-any.whl", hash = "sha256:14c3f3831a892ae38f8bb85240c195ed95e8fa996f60930e2e4f00bb73073036", size = 5255, upload-time = "2024-11-13T20:27:43.888Z" }, + { url = "https://files.pythonhosted.org/packages/25/e9/081aeb69bf4170a5d88de48db11837cce136649b35304ab0ac7164fcc06a/opentelemetry_instrumentation_transformers-0.34.0-py3-none-any.whl", hash = "sha256:984cf5e0f4ef31662382019e3a18edf821f8ce3c20d53aeea68cee5718aad752", size = 5241, upload-time = "2024-12-12T21:01:53.001Z" }, ] [[package]] name = "opentelemetry-instrumentation-urllib3" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6673,14 +6674,14 @@ dependencies = [ { name = "opentelemetry-util-http" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fb/fd/79fa96e997a9ba9f90dd6fd9bd20c67db8b965dea035e54b864665a2508d/opentelemetry_instrumentation_urllib3-0.49b0.tar.gz", hash = "sha256:33db59eafc80877c225467bf71dfe098874dd7f4463a4f12c61fb7dbcd3b4e31", size = 15432, upload-time = "2024-11-05T19:22:23.261Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ed/6f/76a46806cd21002cac1bfd087f5e4674b195ab31ab44c773ca534b6bb546/opentelemetry_instrumentation_urllib3-0.54b1.tar.gz", hash = "sha256:0d30ba3b230e4100cfadaad29174bf7bceac70e812e4f5204e681e4b55a74cd9", size = 15697, upload-time = "2025-05-16T19:04:07.709Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/82/56/6339b51038142ffacc33821d1cf9a3cf91d9c9166088a5c7d862d40000bb/opentelemetry_instrumentation_urllib3-0.49b0-py3-none-any.whl", hash = "sha256:672855f033e608c857353b6e098551f70088664fbec227f4ea5d90463d602adc", size = 12847, upload-time = "2024-11-05T19:21:34.14Z" }, + { url = "https://files.pythonhosted.org/packages/ff/7a/d75bec41edb6deaf1d2859bab66a84c8ba03e822e7eafdb245da205e53f6/opentelemetry_instrumentation_urllib3-0.54b1-py3-none-any.whl", hash = "sha256:e87958c297ddd36d30e1c9069f34a9690e845e4ccc2662dd80e99ed976d4c03e", size = 13123, upload-time = "2025-05-16T19:03:14.053Z" }, ] [[package]] name = "opentelemetry-instrumentation-vertexai" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6688,14 +6689,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/70/99/355c73ba6fb1679f32caa5579d9956dd3e0d40fa2205b41932694bd54696/opentelemetry_instrumentation_vertexai-0.33.12.tar.gz", hash = "sha256:a4ff534f24d4e1caecc621bea1ad19905bafc8ebf2fd1506e9eb1ae8f2a7831a", size = 4356, upload-time = "2024-11-13T20:28:25.514Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f0/63/37f55389efdffcb167ec6e5cdb6b01cfeaaab01becac80d300aff547c2f5/opentelemetry_instrumentation_vertexai-0.34.0.tar.gz", hash = "sha256:4db963d487a4c26875c50dfeddfb589d998cc46b3cb89dc9a3f1083352b9e607", size = 4343, upload-time = "2024-12-12T21:02:32.037Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5c/e6/04cb7853674e4412d929d7c3172b073b402b3ee8f049e57dce042bdd5a2d/opentelemetry_instrumentation_vertexai-0.33.12-py3-none-any.whl", hash = "sha256:cf61cdc08bb6cb4dcbb0b59d1d0432cc1b0b7bee8fa25c69ae4886cf048204c4", size = 5789, upload-time = "2024-11-13T20:27:45.036Z" }, + { url = "https://files.pythonhosted.org/packages/4a/11/6dbf0defdfeeeaf4bb2037732bedbe020e048157642519cd022b33af84e6/opentelemetry_instrumentation_vertexai-0.34.0-py3-none-any.whl", hash = "sha256:d9206a65a416159597676ac60d1331abdc3844e98982126c155e1cacd939d395", size = 5773, upload-time = "2024-12-12T21:01:55.753Z" }, ] [[package]] name = "opentelemetry-instrumentation-watsonx" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6703,14 +6704,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/5b/f4/13359f1ef849d87010e494f503ce0d90c0b37f25b3cdf1ce58f0eaf0aa0b/opentelemetry_instrumentation_watsonx-0.33.12.tar.gz", hash = "sha256:98d537e3e9a919eab87f1f5f487679dd642d0742032635001b974c2154cedc0b", size = 6552, upload-time = "2024-11-13T20:28:26.346Z" } +sdist = { url = "https://files.pythonhosted.org/packages/2e/ed/78fafee5b64f728c5d8958f0fa558148674fb878b5585873e7d76708fa18/opentelemetry_instrumentation_watsonx-0.34.0.tar.gz", hash = "sha256:149a2ec1c6aa476c6258d7f00fc7951220ea8cc23be9a7a1273009377b9df0a4", size = 6552, upload-time = "2024-12-12T21:02:32.962Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a3/f6/74c9e14dc3324a9e5fb9a31c1ae29f92d25afe7930d1430dc55923867547/opentelemetry_instrumentation_watsonx-0.33.12-py3-none-any.whl", hash = "sha256:76bde9b15ca9be9fa6124b7e09606203103b77e7fa05227b8c9145fd2a782102", size = 7457, upload-time = "2024-11-13T20:27:46.861Z" }, + { url = "https://files.pythonhosted.org/packages/9c/ef/4b2189eda9ed49f4ea69e6b102351944d7e17ab90bfa6ff451ee20c1c97d/opentelemetry_instrumentation_watsonx-0.34.0-py3-none-any.whl", hash = "sha256:85d352880c8abccba92c728cbea7cab455a4acb454d43ed0037b6afecdb3a90c", size = 7442, upload-time = "2024-12-12T21:01:58.071Z" }, ] [[package]] name = "opentelemetry-instrumentation-weaviate" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -6718,48 +6719,48 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-semantic-conventions-ai" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1d/50/38b46f295c4f28301d6aea15aeddcbc9550bb51a005da95045310549191f/opentelemetry_instrumentation_weaviate-0.33.12.tar.gz", hash = "sha256:1d14949e2123e5a2bd0eb149d8281713b33623d3f09f7aa587d4fca130d11b70", size = 4635, upload-time = "2024-11-13T20:28:27.189Z" } +sdist = { url = "https://files.pythonhosted.org/packages/68/47/9f0fc2310ef155edd22ae8ee3444e76d91a100a1579b40d034d85d2b0806/opentelemetry_instrumentation_weaviate-0.34.0.tar.gz", hash = "sha256:b69294e0b6b2fc5b90cd389c1a2bc75d18ed09f075ab589a61a0bcbe049ef9db", size = 4654, upload-time = "2024-12-12T21:02:34.344Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/28/e7/16d9a936d716af84546045fa62e593079069e14c667d049602b6e19d6e31/opentelemetry_instrumentation_weaviate-0.33.12-py3-none-any.whl", hash = "sha256:afa500e59bd7059495c6190decb1dd57dc620e17181c9543bd91e26afce74dcd", size = 6428, upload-time = "2024-11-13T20:27:48.71Z" }, + { url = "https://files.pythonhosted.org/packages/6a/9f/f55c020a3619d31dd39d32e376d8d5f8f6322f82b1c27acd5666503d6643/opentelemetry_instrumentation_weaviate-0.34.0-py3-none-any.whl", hash = "sha256:79eaa9be4393702d7b3cc938f3d01d82371d4a236326b01819002bac3f118194", size = 6410, upload-time = "2024-12-12T21:02:00.604Z" }, ] [[package]] name = "opentelemetry-proto" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "protobuf" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c9/63/ac4cef4d30ea0ca1d2153ad2fc62d91d1cf3b89b0e4e5cbd61a8c567885f/opentelemetry_proto-1.28.0.tar.gz", hash = "sha256:4a45728dfefa33f7908b828b9b7c9f2c6de42a05d5ec7b285662ddae71c4c870", size = 34331, upload-time = "2024-11-05T19:14:59.503Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f6/dc/791f3d60a1ad8235930de23eea735ae1084be1c6f96fdadf38710662a7e5/opentelemetry_proto-1.33.1.tar.gz", hash = "sha256:9627b0a5c90753bf3920c398908307063e4458b287bb890e5c1d6fa11ad50b68", size = 34363, upload-time = "2025-05-16T18:52:52.141Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/86/94/c0b43d16e1d96ee1e699373aa59f14a3aa2e7126af3f11d6adc5dcc531cd/opentelemetry_proto-1.28.0-py3-none-any.whl", hash = "sha256:d5ad31b997846543b8e15504657d9a8cf1ad3c71dcbbb6c4799b1ab29e38f7f9", size = 55832, upload-time = "2024-11-05T19:14:40.446Z" }, + { url = "https://files.pythonhosted.org/packages/c4/29/48609f4c875c2b6c80930073c82dd1cafd36b6782244c01394007b528960/opentelemetry_proto-1.33.1-py3-none-any.whl", hash = "sha256:243d285d9f29663fc7ea91a7171fcc1ccbbfff43b48df0774fd64a37d98eda70", size = 55854, upload-time = "2025-05-16T18:52:36.269Z" }, ] [[package]] name = "opentelemetry-sdk" -version = "1.28.0" +version = "1.33.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-semantic-conventions" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/0c/5b/a509ccab93eacc6044591d5ec437d8266e76f893d0389bbf7e5592c7da32/opentelemetry_sdk-1.28.0.tar.gz", hash = "sha256:41d5420b2e3fb7716ff4981b510d551eff1fc60eb5a95cf7335b31166812a893", size = 156155, upload-time = "2024-11-05T19:15:00.451Z" } +sdist = { url = "https://files.pythonhosted.org/packages/67/12/909b98a7d9b110cce4b28d49b2e311797cffdce180371f35eba13a72dd00/opentelemetry_sdk-1.33.1.tar.gz", hash = "sha256:85b9fcf7c3d23506fbc9692fd210b8b025a1920535feec50bd54ce203d57a531", size = 161885, upload-time = "2025-05-16T18:52:52.832Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c3/fe/c8decbebb5660529f1d6ba65e50a45b1294022dfcba2968fc9c8697c42b2/opentelemetry_sdk-1.28.0-py3-none-any.whl", hash = "sha256:4b37da81d7fad67f6683c4420288c97f4ed0d988845d5886435f428ec4b8429a", size = 118692, upload-time = "2024-11-05T19:14:41.669Z" }, + { url = "https://files.pythonhosted.org/packages/df/8e/ae2d0742041e0bd7fe0d2dcc5e7cce51dcf7d3961a26072d5b43cc8fa2a7/opentelemetry_sdk-1.33.1-py3-none-any.whl", hash = "sha256:19ea73d9a01be29cacaa5d6c8ce0adc0b7f7b4d58cc52f923e4413609f670112", size = 118950, upload-time = "2025-05-16T18:52:37.297Z" }, ] [[package]] name = "opentelemetry-semantic-conventions" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "deprecated" }, { name = "opentelemetry-api" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ee/c8/433b0e54143f8c9369f5c4a7a83e73eec7eb2ee7d0b7e81a9243e78c8e80/opentelemetry_semantic_conventions-0.49b0.tar.gz", hash = "sha256:dbc7b28339e5390b6b28e022835f9bac4e134a80ebf640848306d3c5192557e8", size = 95227, upload-time = "2024-11-05T19:15:01.443Z" } +sdist = { url = "https://files.pythonhosted.org/packages/5b/2c/d7990fc1ffc82889d466e7cd680788ace44a26789809924813b164344393/opentelemetry_semantic_conventions-0.54b1.tar.gz", hash = "sha256:d1cecedae15d19bdaafca1e56b29a66aa286f50b5d08f036a145c7f3e9ef9cee", size = 118642, upload-time = "2025-05-16T18:52:53.962Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/25/05/20104df4ef07d3bf5c3fd6bcc796ef70ab4ea4309378a9ba57bc4b4d01fa/opentelemetry_semantic_conventions-0.49b0-py3-none-any.whl", hash = "sha256:0458117f6ead0b12e3221813e3e511d85698c31901cac84682052adb9c17c7cd", size = 159214, upload-time = "2024-11-05T19:14:43.047Z" }, + { url = "https://files.pythonhosted.org/packages/0a/80/08b1698c52ff76d96ba440bf15edc2f4bc0a279868778928e947c1004bdd/opentelemetry_semantic_conventions-0.54b1-py3-none-any.whl", hash = "sha256:29dab644a7e435b58d3a3918b58c333c92686236b30f7891d5e51f02933ca60d", size = 194938, upload-time = "2025-05-16T18:52:38.796Z" }, ] [[package]] @@ -6773,11 +6774,11 @@ wheels = [ [[package]] name = "opentelemetry-util-http" -version = "0.49b0" +version = "0.54b1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/a3/99/377ef446928808211b127b9ab31c348bc465c8da4514ebeec6e4a3de3d21/opentelemetry_util_http-0.49b0.tar.gz", hash = "sha256:02928496afcffd58a7c15baf99d2cedae9b8325a8ac52b0d0877b2e8f936dd1b", size = 7863, upload-time = "2024-11-05T19:22:26.973Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a8/9f/1d8a1d1f34b9f62f2b940b388bf07b8167a8067e70870055bd05db354e5c/opentelemetry_util_http-0.54b1.tar.gz", hash = "sha256:f0b66868c19fbaf9c9d4e11f4a7599fa15d5ea50b884967a26ccd9d72c7c9d15", size = 8044, upload-time = "2025-05-16T19:04:10.79Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/66/0e/ab0a89b315d0bacdd355a345bb69b20c50fc1f0804b52b56fe1c35a60e68/opentelemetry_util_http-0.49b0-py3-none-any.whl", hash = "sha256:8661bbd6aea1839badc44de067ec9c15c05eab05f729f496c856c50a1203caf1", size = 6945, upload-time = "2024-11-05T19:21:37.81Z" }, + { url = "https://files.pythonhosted.org/packages/a4/ef/c5aa08abca6894792beed4c0405e85205b35b8e73d653571c9ff13a8e34e/opentelemetry_util_http-0.54b1-py3-none-any.whl", hash = "sha256:b1c91883f980344a1c3c486cffd47ae5c9c1dd7323f9cbe9fdb7cadb401c87c9", size = 7301, upload-time = "2025-05-16T19:03:18.18Z" }, ] [[package]] @@ -9860,7 +9861,7 @@ wheels = [ [[package]] name = "traceloop-sdk" -version = "0.33.12" +version = "0.34.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aiohttp" }, @@ -9906,9 +9907,9 @@ dependencies = [ { name = "pydantic" }, { name = "tenacity" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/7e/0d/d7d413e9fe907a8abc33e6f93044484d158722b5ca0bfe22e1ef9ad4e729/traceloop_sdk-0.33.12.tar.gz", hash = "sha256:999ae50b1e5773b2802a8b3e8585c3826b7867bba032a88b6f30ec2727225dda", size = 19768, upload-time = "2024-11-13T20:29:26.67Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0a/b1/fd7360d97c651098da505e95600e067a7eedb1b78635b2f1d23545ee4a46/traceloop_sdk-0.34.0.tar.gz", hash = "sha256:4aa26003dfa2e417f73728bd847284a12d6da43a946dd588603a0966e753b3e6", size = 19808, upload-time = "2024-12-12T21:03:41.647Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ce/13/53c2ab6ac27804769314554a062e0651a44db2360be47e21cf0a29d202ee/traceloop_sdk-0.33.12-py3-none-any.whl", hash = "sha256:d47a474afbf4a68ff38a702dbaca7b17d2d4f0b0e14dc2f1560b6bdd3859ac75", size = 25932, upload-time = "2024-11-13T20:29:25.174Z" }, + { url = "https://files.pythonhosted.org/packages/c5/e8/c89cc77c272312930cc263c45fbd2a648536e93358611bf03dba6f176a0b/traceloop_sdk-0.34.0-py3-none-any.whl", hash = "sha256:1cc3e5be9dd2765212feaa5655e1f43ddc66739585d78d9c81134428a2a7d927", size = 25944, upload-time = "2024-12-12T21:03:39.565Z" }, ] [[package]]