diff --git a/.circleci/config.yml b/.circleci/config.yml index 84d8f48b4be..bb4ad0f4019 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -147,6 +147,9 @@ commands: db_name: type: string default: circle_test + image: + type: string + default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26 steps: - run: name: Start PostgreSQL @@ -157,7 +160,7 @@ commands: -e POSTGRES_PASSWORD=postgres \ -e POSTGRES_DB=<< parameters.db_name >> \ -p 5432:5432 \ - postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26 + << parameters.image >> - wait_for_service: url: tcp://localhost:5432 timeout: "60" @@ -2912,7 +2915,69 @@ jobs: exit 1 fi + provider_replay_harness: + docker: + - *python312_image + working_directory: ~/project + resource_class: medium + steps: + - setup_litellm_test_deps + - run: + name: Test provider replay harness + command: | + mkdir -p test-results/provider-replay-harness + uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \ + --junitxml=test-results/provider-replay-harness/junit.xml \ + tests/e2e/test_provider_edge.py tests/e2e/test_fixture_bundle.py \ + tests/e2e/test_fixture_canonical.py tests/e2e/test_fixture_mode.py \ + tests/code_coverage_tests/test_provider_replay_harness.py + - store_test_results: + path: test-results/provider-replay-harness + + integration_contracts: + parameters: + suite: + type: string + machine: + image: ubuntu-2204:2024.04.1 + resource_class: large + working_directory: ~/project + steps: + - setup_litellm_test_deps + - start_postgres: + image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5 + - start_redis + - run: + name: Run owned integration contracts + command: bash .circleci/scripts/run_integration.sh << parameters.suite >> + no_output_timeout: 15m + - run: + name: Stop owned database and Redis + when: always + command: | + mkdir -p test-results/integration-<< parameters.suite >> + docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true + docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true + docker rm -f postgres-db redis-cache + test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)" + - store_test_results: + path: test-results + - store_artifacts: + path: test-results + workflows: + integration: + jobs: + - integration_contracts: + name: integration-<< matrix.suite >> + matrix: + parameters: + suite: [management, accounting, providers] + filters: + branches: + only: + - main + - /litellm_.*/ build_and_test: jobs: - using_litellm_on_windows: @@ -2921,6 +2986,7 @@ workflows: only: - main - /litellm_.*/ + - provider_replay_harness - base_sdk_install: filters: *main_branches - local_testing_part1: diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh new file mode 100644 index 00000000000..9fd2e7c32df --- /dev/null +++ b/.circleci/scripts/run_integration.sh @@ -0,0 +1,142 @@ +#!/usr/bin/env bash +set -euo pipefail + +suite="${1:?integration suite required}" +results="test-results/integration-${suite}" +mkdir -p "$results" +integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')" +upstream_pid="" +proxy_pid="" +peer_pid="" +launched_pid="" +guard_created=false +guard_installed=false +guard6_created=false +guard6_installed=false +cleanup() { + original_status=$? + trap - EXIT INT TERM + sudo .venv/bin/python .circleci/scripts/stop_integration_processes.py \ + "$integration_identity" "$(id -u)" "$proxy_pid" "$peer_pid" "$upstream_pid" \ + > "$results/process-cleanup.txt" 2>&1 || original_status=1 + for owned_pid in "$peer_pid" "$proxy_pid" "$upstream_pid"; do + if [ -n "$owned_pid" ]; then + kill -- "-$owned_pid" 2>/dev/null || true + for _ in {1..50}; do + kill -0 -- "-$owned_pid" 2>/dev/null || break + sleep 0.1 + done + if kill -0 -- "-$owned_pid" 2>/dev/null; then + kill -KILL -- "-$owned_pid" 2>/dev/null || true + original_status=1 + fi + wait "$owned_pid" 2>/dev/null || true + fi + done + if [ "$guard_installed" = true ]; then + sudo iptables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1 + fi + if [ "$guard_created" = true ]; then + sudo iptables -F integration_only || original_status=1 + sudo iptables -X integration_only || original_status=1 + fi + if [ "$guard6_installed" = true ]; then + sudo ip6tables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1 + fi + if [ "$guard6_created" = true ]; then + sudo ip6tables -F integration_only || original_status=1 + sudo ip6tables -X integration_only || original_status=1 + fi + printf '%s\n' "$original_status" > "$results/exit-status.txt" + exit "$original_status" +} +trap cleanup EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +export PATH="$PWD/.venv/bin:$PATH" +export PYTHONPATH="$PWD:$PWD/tests:$PWD/tests/e2e" +export DATABASE_URL="postgresql://postgres:postgres@127.0.0.1:5432/circle_test" +export REDIS_HOST=127.0.0.1 REDIS_PORT=6379 +export LITELLM_MASTER_KEY=sk-integration-master LITELLM_SALT_KEY=sk-integration-salt +export LITELLM_MODE=PRODUCTION LITELLM_LOCAL_MODEL_COST_MAP=True +export STORE_MODEL_IN_DB=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 +export INTEGRATION_PROXY_URL=http://127.0.0.1:4000 +export INTEGRATION_PEER_URL="" +export INTEGRATION_UPSTREAM_URL=http://127.0.0.1:8190 +export INTEGRATION_MASTER_KEY="$LITELLM_MASTER_KEY" +export INTEGRATION_SEED="$((16#$(git rev-parse --short=8 HEAD)))" + +uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1 + +sudo iptables -N integration_only +guard_created=true +sudo iptables -A integration_only -o lo -j ACCEPT +sudo iptables -A integration_only -m conntrack --ctstate ESTABLISHED,RELATED -j ACCEPT +for service in postgres-db redis-cache; do + address="$(docker inspect --format '{{range .NetworkSettings.Networks}}{{.IPAddress}}{{end}}' "$service")" + sudo iptables -A integration_only -d "$address" -j ACCEPT +done +sudo iptables -A integration_only -j REJECT +sudo iptables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only +guard_installed=true +sudo ip6tables -N integration_only +guard6_created=true +sudo ip6tables -A integration_only -o lo -j ACCEPT +sudo ip6tables -A integration_only -j REJECT +sudo ip6tables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only +guard6_installed=true + +if curl --noproxy '*' --connect-timeout 2 -s http://198.51.100.1 >/dev/null 2>&1; then + echo "Unexpected outbound network access" >&2 + exit 1 +fi +sudo iptables -L integration_only -n -v -x > "$results/egress-guard.txt" +awk '$3 == "REJECT" && $1 > 0 { rejected=1 } END { exit !rejected }' "$results/egress-guard.txt" + +setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \ + .venv/bin/python -m integration._support.upstream > "$results/upstream.log" 2>&1 & +upstream_pid=$! +start_proxy() { + local port="$1" + local log_name="$2" + setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \ + DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \ + LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" \ + LITELLM_MODE=PRODUCTION LITELLM_LOCAL_MODEL_COST_MAP=True STORE_MODEL_IN_DB=True \ + AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \ + .venv/bin/python -m integration._support.proxy --config tests/integration/proxy_config.yaml \ + --host 127.0.0.1 --port "$port" --num_workers 1 --telemetry False \ + --use_prisma_db_push --enforce_prisma_migration_check \ + > "$results/$log_name" 2>&1 & + launched_pid=$! +} +start_proxy 4000 proxy.log +proxy_pid="$launched_pid" +.venv/bin/python .circleci/scripts/wait_integration_services.py +if [ "$suite" = management ]; then + export INTEGRATION_PEER_URL=http://127.0.0.1:4001 + start_proxy 4001 peer.log + peer_pid="$launched_pid" + .venv/bin/python .circleci/scripts/wait_integration_services.py +fi + +if [ "$suite" = providers ]; then + INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --noconftest -o addopts= \ + --strict-markers --strict-config -p no:pytest-retry -p no:rerunfailures --timeout=30 \ + tests/e2e/test_provider_edge.py::TestReplayMode::test_content_drift_returns_the_miss_status_naming_both_keys \ + tests/e2e/test_provider_edge.py::TestReplayMode::test_exhausted_key_returns_the_miss_status \ + tests/e2e/test_provider_edge.py::TestReplayLeftover::test_partially_consumed_recording_names_the_leftover \ + tests/e2e/test_provider_edge.py::TestStreamingFidelity::test_replay_of_a_stream_makes_no_provider_connection \ + --junitxml="$results/replay-controls.xml" +fi + +timeout --signal=TERM --kill-after=20s 11m env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \ + INTEGRATION_RUN_ID="$integration_identity" \ + DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \ + INTEGRATION_PROXY_URL="$INTEGRATION_PROXY_URL" INTEGRATION_PEER_URL="$INTEGRATION_PEER_URL" \ + INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \ + INTEGRATION_MASTER_KEY="$INTEGRATION_MASTER_KEY" LITELLM_MODE=PRODUCTION \ + INTEGRATION_SEED="$INTEGRATION_SEED" \ + LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \ + .venv/bin/python tests/integration/run.py "$suite" --results "$results" diff --git a/.circleci/scripts/stop_integration_processes.py b/.circleci/scripts/stop_integration_processes.py new file mode 100644 index 00000000000..8f11aaaddbd --- /dev/null +++ b/.circleci/scripts/stop_integration_processes.py @@ -0,0 +1,53 @@ +import sys +from typing import Final + +import psutil + + +def is_owned(process: psutil.Process, identity: str, owner_uid: int) -> bool: + try: + return process.uids().real == owner_uid and process.environ().get("INTEGRATION_RUN_ID") == identity + except psutil.NoSuchProcess: + return False + + +def owned_processes(identity: str, owner_uid: int) -> tuple[psutil.Process, ...]: + return tuple(process for process in psutil.process_iter() if is_owned(process, identity, owner_uid)) + + +def main(identity: str, owner_uid: int, root_pids: tuple[int, ...]) -> int: + assert owner_uid > 0, "The integration process owner must be a non-root UID" + owned: Final = owned_processes(identity, owner_uid) + roots: Final = tuple(process for process in owned if process.pid in root_pids) + for process in roots: + try: + process.terminate() + except psutil.NoSuchProcess: + continue + psutil.wait_procs(roots, timeout=30) + residual: Final = owned_processes(identity, owner_uid) + for process in residual: + try: + process.terminate() + except psutil.NoSuchProcess: + continue + psutil.wait_procs(residual, timeout=10) + remaining: Final = owned_processes(identity, owner_uid) + for process in remaining: + try: + process.kill() + except psutil.NoSuchProcess: + continue + psutil.wait_procs(remaining, timeout=2) + survivors: Final = owned_processes(identity, owner_uid) + print( + f"Owned integration processes: {len(owned)}, roots: {len(roots)}, " + f"residual: {len(residual)}, forced: {len(remaining)}, remaining: {len(survivors)}" + ) + for process in remaining: + print(f"Forced cleanup was required for PID {process.pid}") + return 1 if remaining or survivors else 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1], int(sys.argv[2]), tuple(int(value) for value in sys.argv[3:] if value))) diff --git a/.circleci/scripts/wait_integration_services.py b/.circleci/scripts/wait_integration_services.py new file mode 100644 index 00000000000..486e37cba00 --- /dev/null +++ b/.circleci/scripts/wait_integration_services.py @@ -0,0 +1,43 @@ +import os +import time +from typing import Final + +import httpx +from redis import Redis + + +def main() -> None: + primary: Final = os.environ["INTEGRATION_PROXY_URL"] + peer: Final = os.environ.get("INTEGRATION_PEER_URL") + proxies: Final = (primary, peer) if peer else (primary,) + deadline: Final = time.monotonic() + 90 + headers: Final = {"Authorization": f"Bearer {os.environ['INTEGRATION_MASTER_KEY']}"} + with httpx.Client(trust_env=False, timeout=2) as client, Redis( + host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), socket_timeout=2 + ) as cache: + while True: + try: + ready: Final = ( + client.get(f"{os.environ['INTEGRATION_UPSTREAM_URL']}/health").status_code == 200 + and all(client.get(f"{url}/health/readiness").status_code == 200 for url in proxies) + ) + if ready: + for url in proxies: + response: Final = client.get(f"{url}/cache/ping", headers=headers) + response.raise_for_status() + result: Final = response.json() + assert result["status"] == "healthy", result + assert result["cache_type"] == "redis", result + assert result["ping_response"] is True, result + assert result["set_cache_response"] == "success", result + if cache.pubsub_numsub("litellm_proxy.auth_cache_invalidation")[0][1] >= len(proxies): + return + except httpx.TransportError: + pass + if time.monotonic() >= deadline: + raise SystemExit("Integration services or auth-cache subscribers did not become ready") + time.sleep(0.2) + + +if __name__ == "__main__": + main() diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml index 36d70c1d746..6e15c1069a3 100644 --- a/.github/codeql/codeql-config.yml +++ b/.github/codeql/codeql-config.yml @@ -14,6 +14,17 @@ query-filters: id: py/clear-text-logging-sensitive-data # CWE-312 - exclude: id: py/polynomial-redos # CWE-730 + # Import resolution confuses stdlib types with management_endpoints/types.py. + # The generic cycle query also reports intentional deferred imports. + - exclude: + id: py/cyclic-import + - exclude: + id: py/unsafe-cyclic-import + # Known false positives on live settings and Protocol placeholders. + - exclude: + id: py/unused-global-variable + - exclude: + id: py/ineffectual-statement paths-ignore: - tests diff --git a/.github/e2e-stack/assert_tests_ran.py b/.github/e2e-stack/assert_tests_ran.py index c4348c20873..2303c42f4fb 100644 --- a/.github/e2e-stack/assert_tests_ran.py +++ b/.github/e2e-stack/assert_tests_ran.py @@ -3,6 +3,9 @@ import xml.etree.ElementTree as ET from pathlib import Path from typing import Final +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "tests/e2e")) +from coverage_registry.management_cases import MANAGEMENT_CASES + def main() -> int: selected: Final = tuple(sys.argv[2:]) @@ -16,6 +19,17 @@ def main() -> int: case.get("file") for case in cases if all(case.find(tag) is None for tag in ("skipped", "failure", "error")) ) missing: Final = tuple(path for path in selected if path not in passed) + required_nodes: Final = frozenset(case.node for case in MANAGEMENT_CASES if case.node.split("::", 1)[0] in selected) + passed_nodes: Final = frozenset( + prop.get("value") + for case in cases + if all(case.find(tag) is None for tag in ("skipped", "failure", "error")) + for prop in case.findall("./properties/property") + if prop.get("name") == "management_node" + ) + missing_nodes: Final = required_nodes - passed_nodes + for node in sorted(missing_nodes): + _ = sys.stdout.write(f"::error::required management case did not pass: {node}\n") for path in selected: collected: Final = sum(case.get("file") == path for case in cases) skipped: Final = sum(case.get("file") == path and case.find("skipped") is not None for case in cases) @@ -27,6 +41,7 @@ def main() -> int: if ( selected and not missing + and not missing_nodes and not any(case.find(tag) is not None for case in cases for tag in ("failure", "error")) ): return 0 diff --git a/.github/e2e-stack/oidc-profile.sh b/.github/e2e-stack/oidc-profile.sh new file mode 100755 index 00000000000..84eaaaf8051 --- /dev/null +++ b/.github/e2e-stack/oidc-profile.sh @@ -0,0 +1,5 @@ +#!/usr/bin/env bash +set -euo pipefail +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +cd "${REPO_ROOT}" +exec uv run --no-sync python tests/e2e/idp.py "$@" diff --git a/.github/e2e-stack/select_tests.py b/.github/e2e-stack/select_tests.py index 238818a0d36..982e93cf642 100644 --- a/.github/e2e-stack/select_tests.py +++ b/.github/e2e-stack/select_tests.py @@ -12,6 +12,8 @@ UNSUPPORTED: Final = re.compile( HARNESS: Final = re.compile( r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$" r"|^tests/e2e/idp_realm\.json$" + r"|^tests/e2e/management/(management_client|jwt_actors|conftest)\.py$" + r"|^tests/e2e/coverage_registry/management_cases\.py$" r"|^tests/e2e/gateway/" r"|^\.github/e2e-stack/" r"|^\.github/workflows/test-e2e-changed\.yml$" diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 2dc85fce05b..7a9883df356 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -101,7 +101,8 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac For bug fixes: Before shows the reproduction, After shows the same steps passing For new features: Before shows the capability missing, After shows it working end-to-end If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one - For UI changes: before/after screenshots under the same headings --> + For UI changes: before/after screenshots under the same headings + If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof --> ## Type diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index 411852acb98..4c66ab251de 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -1,6 +1,7 @@ from __future__ import annotations import ast +import json import operator import pathlib import re @@ -498,6 +499,69 @@ def _check_shards() -> int: return 0 +def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]: + manifest: Final = repo_root / "tests/integration/contracts.json" + if not manifest.exists(): + return frozenset(), () + entries: Final = json.loads(manifest.read_text()) + paths: Final = frozenset(node.split("::", 1)[0] for node in entries["tests"]) + circle_path: Final = repo_root / ".circleci/config.yml" + circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {} + steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ()) + invoked: Final = any( + ".circleci/scripts/run_integration.sh" in scalar.value + for scalar in _scalars(steps, "integration_contracts") + if scalar.key == "command" + ) + scheduled: Final = frozenset( + suite + for job in circle.get("workflows", {}).get("integration", {}).get("jobs", ()) + if isinstance(job, dict) and "integration_contracts" in job + for suite in job["integration_contracts"] + .get("matrix", {}) + .get("parameters", {}) + .get("suite", (job["integration_contracts"].get("suite"),)) + if isinstance(suite, str) + ) + required: Final = frozenset( + group + for group, folders in entries["groups"].items() + if any(any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for path in paths) + ) + ungrouped: Final = frozenset( + path + for path in paths + if sum( + any(path.startswith(f"tests/integration/{folder}/") for folder in folders) + for folders in entries["groups"].values() + ) + != 1 + ) + gha_tokens: Final = _invoked_test_tokens( + scalar + for path in (repo_root / ".github/workflows").glob("*.y*ml") + for scalar in _scalars(yaml.safe_load(path.read_text()), path.name) + ) + findings: Final = tuple( + Finding(path, "integration contract is also selected by GitHub Actions") + for path in paths + if any(_token_covers(token, path) for token in gha_tokens) + ) + tuple( + Finding(path, "canonical integration test file is missing") + for path in paths + if not (repo_root / path).is_file() + ) + group_findings: Final = tuple( + Finding(group, "canonical integration group is not scheduled by CircleCI") + for group in sorted(required - scheduled) + ) + tuple(Finding(path, "canonical node must have exactly one integration group") for path in sorted(ungrouped)) + if not paths or not invoked or not scheduled: + return frozenset(), findings + ( + Finding(str(manifest.relative_to(repo_root)), "dedicated CircleCI runner is missing"), + ) + return paths, findings + group_findings + + def main() -> int: if "--shards" in sys.argv[1:]: return _check_shards() @@ -507,7 +571,8 @@ def main() -> int: allowlist = _load_allowlist() scalars = _all_scalars() - test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars)) + integration_paths, ownership_findings = _integration_ownership() + test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | integration_paths) + ownership_findings dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars)) stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles()) diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 62790e23143..bbf0cb4e891 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -57,6 +57,7 @@ permissions: env: UV_PYTHON: "3.12" + LITELLM_LOCAL_MODEL_COST_MAP: "True" jobs: run: @@ -113,6 +114,7 @@ jobs: if: steps.changes.outputs.decision != 'skip' timeout-minutes: 8 run: | + diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]' diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 9c7e0db7065..987f66773f2 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -178,7 +178,7 @@ jobs: version: "0.10.9" - name: Install dependencies - run: uv sync --frozen --extra proxy --python 3.10 + run: uv sync --frozen --extra proxy --extra cli --python 3.10 - run: uv run --no-sync python --version @@ -187,3 +187,6 @@ jobs: - name: Check litellm CLI run: uv run --no-sync litellm --version + + - name: Check lite CLI + run: uv run --no-sync lite version diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml index 1db597ff673..c9f08deb36e 100644 --- a/.github/workflows/test-e2e-changed.yml +++ b/.github/workflows/test-e2e-changed.yml @@ -183,7 +183,7 @@ jobs: log="${RUNNER_TEMP}/e2e-pass-${pass}.log" echo "::group::pass ${pass} of 3" set +e - uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v -p no:cacheprovider \ + uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v --reruns 0 -p no:cacheprovider \ -o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1 status=$? uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}" diff --git a/litellm/__init__.py b/litellm/__init__.py index 261457d6889..ccfbf80369f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -538,7 +538,7 @@ context_window_fallbacks: Optional[List] = None content_policy_fallbacks: Optional[List] = None allowed_fails: int = 3 allow_dynamic_callback_disabling: bool = True -num_retries_per_request: Optional[int] = None # cap on Router retries of one model group; resets per fallback hop +num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries) ####### SECRET MANAGERS ##################### secret_manager_client: Optional[Any] = ( None # list of instantiated key management clients - e.g. azure kv, infisical, etc. diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 87f8fd3946e..26b4318da2d 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -9,6 +9,7 @@ import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details +from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details from litellm.types.llms.openai import Batch from litellm.types.utils import ModelInfo, Usage from litellm.utils import token_counter @@ -356,6 +357,7 @@ def calculate_vertex_ai_batch_cost_and_usage( prompt_tokens=_prompt, completion_tokens=_completion, total_tokens=_total, + prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata), ) try: diff --git a/litellm/caching/affinity_cache.py b/litellm/caching/affinity_cache.py new file mode 100644 index 00000000000..2712679b99b --- /dev/null +++ b/litellm/caching/affinity_cache.py @@ -0,0 +1,125 @@ +"""Atomic affinity claims shared by deployment and tier-model selection.""" + +import json +from collections.abc import Mapping +from typing import ( + Final, + cast, # noqa: TID251 # Redis script results are narrowed only to object, then validated +) + +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache + +_PIN_JSON_ADAPTER: Final = TypeAdapter[JsonValue](JsonValue) + +_CLAIM_PIN_SCRIPT: Final = """ +local current = redis.call('GET', KEYS[1]) +if current == false then + redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2]) + return ARGV[1] +end +if ARGV[3] then + local decoded, stored = pcall(cjson.decode, current) + if decoded and type(stored) == 'table' then + for _, eligible in ipairs(cjson.decode(ARGV[3])) do + local matches = true + for key, value in pairs(eligible) do + if stored[key] ~= value then matches = false; break end + end + for key, _ in pairs(stored) do + if eligible[key] == nil then matches = false; break end + end + if matches then + redis.call('EXPIRE', KEYS[1], ARGV[2]) + return current + end + end + end + redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2]) + return ARGV[1] +end +if current == ARGV[1] then + redis.call('EXPIRE', KEYS[1], ARGV[2]) +end +return current +""" + + +def set_local_affinity_pin(cache: DualCache, cache_key: str, value: object, ttl_seconds: int) -> None: + """Replace the entry because InMemoryCache.set_cache preserves a live key's expiry.""" + cache.in_memory_cache.delete_cache(cache_key) + cache.in_memory_cache.set_cache(cache_key, value, ttl=ttl_seconds) + + +def _legacy_pin_matches(stored: object, pin_value: Mapping[str, str]) -> bool: + if isinstance(stored, dict): + return all(stored.get(key) is not None and str(stored[key]) == value for key, value in pin_value.items()) + return isinstance(stored, str) and len(pin_value) == 1 and stored in pin_value.values() + + +def claim_affinity_pin_in_memory( + cache: DualCache, + cache_key: str, + pin_value: Mapping[str, str], + ttl_seconds: int, + *, + eligible_values: tuple[Mapping[str, str], ...] | None = None, +) -> object: + """No await between read and write, so same-loop claims agree during a Redis outage.""" + existing: Final[object] = cache.in_memory_cache.get_cache(cache_key) + if existing is not None and eligible_values is None: + if _legacy_pin_matches(existing, pin_value): + set_local_affinity_pin(cache, cache_key, pin_value, ttl_seconds) + return existing + winner: Final = existing if existing is not None and existing in (eligible_values or ()) else pin_value + set_local_affinity_pin(cache, cache_key, winner, ttl_seconds) + return winner + + +def _decode_pin(value: str) -> object: + try: + return _PIN_JSON_ADAPTER.validate_json(value) + except ValidationError: + return value + + +async def claim_affinity_pin( + cache: DualCache, + cache_key: str, + pin_value: Mapping[str, str], + ttl_seconds: int, + *, + eligible_values: tuple[Mapping[str, str], ...] | None = None, +) -> object: + """Return the authoritative first writer, replacing it only when it becomes ineligible. + + Eligible claims refresh the returned winner. Legacy deployment claims only refresh + a matching candidate. Resolve Redis per call because the proxy attaches it lazily. + """ + redis_cache: Final = cache.redis_cache + if redis_cache is not None: + try: + claim_script: Final = redis_cache.async_register_script(_CLAIM_PIN_SCRIPT) + args: Final = ( + json.dumps(dict(pin_value)), # mutable-ok: JSON serialization requires dict, not a generic Mapping + int(ttl_seconds), + *( + (json.dumps(tuple(dict(value) for value in eligible_values)),) # mutable-ok: JSON requires dict + if eligible_values is not None + else () + ), + ) + raw: Final = cast( # cast-ok: Redis scripts return heterogeneous values; only object is asserted here + object, await claim_script(keys=(cache_key,), args=args) + ) + decoded: Final = raw.decode("utf-8") if isinstance(raw, bytes) else raw + if not isinstance(decoded, str): + return pin_value + winner: Final = _decode_pin(decoded) + set_local_affinity_pin(cache, cache_key, winner, ttl_seconds) + return winner + except Exception as error: # noqa: BLE001 # Redis/Lua faults retain same-pod affinity through local claims + verbose_router_logger.debug("Affinity cache: Redis claim failed, using pod-local claim. error=%s", error) + return claim_affinity_pin_in_memory(cache, cache_key, pin_value, ttl_seconds, eligible_values=eligible_values) diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index c646baf9d9e..b80f78a50c1 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -205,21 +205,35 @@ def _extract_anthropic_tool_exchange_spans( return spans, None +def _message_has_cache_control(message: Mapping[str, object]) -> bool: + if message.get("cache_control") is not None: + return True + content: Final = message.get("content") + if isinstance(content, list): + return any(isinstance(part, Mapping) and part.get("cache_control") is not None for part in content) + return False + + def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int, ...]: """ Return indices of messages that must never be compressed: - All system messages - The last user message - The last assistant message + - Any message carrying an Anthropic cache_control breakpoint The last user message is what the model is being asked to act on right now, so compressing it replaces the live instruction with a marker. Compression - guardrails share this policy; see the Headroom guardrail. + guardrails share this policy; see the Headroom guardrail. A cache_control + breakpoint pins the provider's prompt-cache prefix to that row's exact + bytes, so rewriting a marked row anywhere in history turns the next + request's cache read into a cache write. """ system_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system") last_user: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:] - last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:] - return system_indices + last_user + last_assistant + assistant_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant") + cache_control_indices: Final = tuple(index for index, msg in enumerate(messages) if _message_has_cache_control(msg)) + return tuple(dict.fromkeys(system_indices + last_user + assistant_indices[-1:] + cache_control_indices)) def _combine_scores( @@ -421,7 +435,7 @@ def compress( combined_scores = bm25_scores # Protected messages are never compressed - protected_indices: Final = get_protected_indices(normalized_messages) + protected_indices: Final = get_protected_indices(original_messages) kept_indices: set[int] = set(protected_indices) tool_exchange_spans: list[set[int]] = [] diff --git a/litellm/constants.py b/litellm/constants.py index c106be688e4..09442d6151e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1976,6 +1976,8 @@ BROWSER_SECURITY_HEADERS: Final[frozenset[str]] = frozenset( UNSAFE_PROXY_RESPONSE_HEADERS: Final[frozenset[str]] = HTTP_FRAMING_HEADERS | BROWSER_SECURITY_HEADERS +STRINGIFIED_NONE: Final[str] = "None" + # A retrieved response replays the usage of the call that created it, so pricing these # read/management routes like inference bills the same tokens twice. NON_INFERENCE_CALL_TYPES: Final[frozenset[str]] = frozenset( diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index f5319776213..3dc6d81256b 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -2278,6 +2278,19 @@ def default_video_cost_calculator( return 0.0 +def _batch_rate( + model_info: ModelInfo, + key: Literal[ + "input_cost_per_audio_token_batches", + "input_cost_per_image_token_batches", + "input_cost_per_video_token_batches", + ], + fallback: float, +) -> float: + rate: Final = model_info.get(key) + return fallback if rate is None else rate + + def batch_cost_calculator( usage: Usage, model: str, @@ -2337,7 +2350,29 @@ def batch_cost_calculator( total_prompt_cost = 0.0 total_completion_cost = 0.0 if input_cost_per_token_batches is not None: - total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches + batch_details: Final = parse_prompt_tokens_details(usage) + audio_tokens, image_tokens, video_tokens = ( + batch_details["audio_tokens"], + batch_details["image_tokens"], + batch_details["video_tokens"], + ) + modality_rates: Final = ( + _batch_rate(model_info, "input_cost_per_audio_token_batches", input_cost_per_token_batches), + _batch_rate(model_info, "input_cost_per_image_token_batches", input_cost_per_token_batches), + _batch_rate(model_info, "input_cost_per_video_token_batches", input_cost_per_token_batches), + ) + total_prompt_cost = sum( + tokens * rate + for tokens, rate in zip( + ( + max((usage.prompt_tokens or 0) - audio_tokens - image_tokens - video_tokens, 0), + audio_tokens, + image_tokens, + video_tokens, + ), + (input_cost_per_token_batches, *modality_rates), + ) + ) elif input_cost_per_token: details: Final = parse_prompt_tokens_details(usage) cache_read_tokens: Final = details["cache_hit_tokens"] diff --git a/litellm/integrations/arize/arize_phoenix_prompt_manager.py b/litellm/integrations/arize/arize_phoenix_prompt_manager.py index 0c9e868c146..fbf25ea87fd 100644 --- a/litellm/integrations/arize/arize_phoenix_prompt_manager.py +++ b/litellm/integrations/arize/arize_phoenix_prompt_manager.py @@ -379,7 +379,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement): """Reload prompts from Arize Phoenix.""" if self.prompt_id: self._prompt_manager = None # Reset to force reload - self.prompt_manager # This will trigger reload + _ = self.prompt_manager # access triggers lazy reload def should_run_prompt_management( self, diff --git a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py index ff34bd91e31..e98fa77a562 100644 --- a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py +++ b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py @@ -406,7 +406,7 @@ class BitBucketPromptManager(CustomPromptManagement): """Reload prompts from BitBucket.""" if self.prompt_id: self._prompt_manager = None # Reset to force reload - self.prompt_manager # This will trigger reload + _ = self.prompt_manager # access triggers lazy reload def should_run_prompt_management( self, diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 39adea30828..1d00ad8c29a 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -884,7 +884,9 @@ class CustomGuardrail(CustomLogger): """logging_only: run apply_guardrail on copies of the logged request/response and record the verdict.""" from litellm.llms import get_guardrail_translation_mapping - if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks: + if not self.uses_apply_guardrail_interface(): + return kwargs, result + if not self._event_hook_is_event_type(GuardrailEventHooks.logging_only): return kwargs, result try: translation: Final = get_guardrail_translation_mapping(CallTypes(call_type))() @@ -901,8 +903,18 @@ class CustomGuardrail(CustomLogger): for key, value in (litellm_params.get("metadata") or {}).items() if key != "standard_logging_guardrail_information" } + response: Final = ( + kwargs.get("async_complete_streaming_response") or kwargs.get("complete_streaming_response") or result + ) + from litellm.types.utils import ModelResponse + + output_translation: Final = ( + get_guardrail_translation_mapping(CallTypes.acompletion)() + if isinstance(response, ModelResponse) + else translation + ) try: - await self._scan_logged_call(kwargs, result, translation, scratch_metadata) + await self._scan_logged_call(kwargs, response, translation, output_translation, scratch_metadata) except Exception as e: verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e) recorded: Final = scratch_metadata.get("standard_logging_guardrail_information") @@ -919,8 +931,9 @@ class CustomGuardrail(CustomLogger): async def _scan_logged_call( self, kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract - result: object, + response: object | None, translation: "BaseTranslation", + output_translation: "BaseTranslation", scratch_metadata: dict, # mutable-ok: apply_guardrail records its verdict into request metadata ) -> None: optional_params: Final = kwargs.get("optional_params") or {} @@ -934,8 +947,10 @@ class CustomGuardrail(CustomLogger): "metadata": scratch_metadata, } await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self) - await translation.process_output_response( - response=copy.deepcopy(result), guardrail_to_apply=self, request_data=scratch_request + if response is None: + return + await output_translation.process_output_response( + response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request ) def supports_scan_only_tool_results(self) -> bool: diff --git a/litellm/integrations/weights_biases.py b/litellm/integrations/weights_biases.py index 97a2acbac08..d1a8ec098cf 100644 --- a/litellm/integrations/weights_biases.py +++ b/litellm/integrations/weights_biases.py @@ -4,18 +4,10 @@ imported_openAIResponse = True try: import io import logging - import sys - from typing import Any, TypeVar + from typing import Any, Literal, Protocol, TypeVar from wandb.sdk.data_types import trace_tree - if sys.version_info >= (3, 8): - from typing import Literal, Protocol - else: - from typing import Literal - - from typing_extensions import Protocol - logger: Final = logging.getLogger(__name__) K = TypeVar("K", bound=str) diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 6e76bf9d49e..15380bc5d57 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -309,8 +309,8 @@ def max_retries_per_request_hit(kwargs: Mapping[str, object], num_retries_per_re metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs)) if not isinstance(metadata, Mapping): return False - attempted_retries: Final = metadata.get("attempted_retries") - return type(attempted_retries) is int and 0 < attempted_retries and num_retries_per_request <= attempted_retries + retry_count: Final = metadata.get("request_retry_count") + return type(retry_count) is int and 0 < retry_count and num_retries_per_request <= retry_count def get_or_create_metadata_bucket( diff --git a/litellm/litellm_core_utils/coroutine_checker.py b/litellm/litellm_core_utils/coroutine_checker.py index 52fc44ba8dc..7b9a650c66b 100644 --- a/litellm/litellm_core_utils/coroutine_checker.py +++ b/litellm/litellm_core_utils/coroutine_checker.py @@ -36,7 +36,7 @@ class CoroutineChecker: target = callback if not inspect.isfunction(target) and not inspect.ismethod(target): try: - call_attr: Final = getattr(target, "__call__", None) + call_attr: Final = getattr(target, "__call__", None) # noqa: B004 # value unwrap so iscoroutinefunction sees through functors if call_attr is not None: target = call_attr except Exception: diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 82708d412c9..70675966dfc 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -1,6 +1,9 @@ +import inspect import json import re import traceback +from collections.abc import Mapping +from types import MappingProxyType from typing import Any, Final, Protocol, cast import httpx @@ -202,11 +205,17 @@ def _get_response_headers(original_exception: Exception) -> httpx.Headers | None return _response_headers +def _accepted_init_kwargs(exception_class: type[Exception], candidates: Mapping[str, object]) -> Mapping[str, object]: + accepted: Final = inspect.signature(exception_class).parameters + return MappingProxyType({name: value for name, value in candidates.items() if name in accepted}) + + def extract_and_raise_litellm_exception( response: Any | None, error_str: str, model: str, custom_llm_provider: str, + body: object | None = None, ): """ Covers scenario where litellm sdk calling proxy. @@ -216,32 +225,19 @@ def extract_and_raise_litellm_exception( Relevant Issue: https://github.com/BerriAI/litellm/issues/7259 """ pattern: Final = r"litellm\.\w+Error" - - # Search for the exception in the error string match: Final = re.search(pattern, error_str) - - # Extract the exception if found - if match: - exception_name = match.group(0) - exception_name = exception_name.strip().replace("litellm.", "") - raised_exception_obj: Final = getattr(litellm, exception_name, None) - if raised_exception_obj: - # Try with response parameter first, fall back to without it - # Some exceptions (e.g., APIConnectionError) don't accept response param - try: - raise raised_exception_obj( - message=error_str, - llm_provider=custom_llm_provider, - model=model, - response=response, - ) - except TypeError: - # Exception doesn't accept response parameter - raise raised_exception_obj( - message=error_str, - llm_provider=custom_llm_provider, - model=model, - ) + if match is None: + return + exception_name: Final = match.group(0).removeprefix("litellm.") + raised_exception_obj: Final = getattr(litellm, exception_name, None) + if not raised_exception_obj: + return + raise raised_exception_obj( + message=error_str, + llm_provider=custom_llm_provider, + model=model, + **_accepted_init_kwargs(raised_exception_obj, MappingProxyType({"response": response, "body": body})), + ) class _ProviderHTTPException(Protocol): @@ -254,6 +250,23 @@ class _ProviderHTTPException(Protocol): llm_provider: str +def _litellm_proxy_response( + original_exception: _ProviderHTTPException, custom_llm_provider: str +) -> httpx.Response | None: + response: Final = getattr(original_exception, "response", None) + if custom_llm_provider != "litellm_proxy" or not isinstance(response, httpx.Response) or response.headers: + return response + headers: Final = getattr(original_exception, "headers", None) + if not isinstance(headers, Mapping) or not headers: + return response + pairs: Final = headers.multi_items() if isinstance(headers, httpx.Headers) else headers.items() + return httpx.Response( + status_code=response.status_code, + headers=[(str(k), str(v)) for k, v in pairs], + request=getattr(original_exception, "request", None), + ) + + def _map_openai_exception( *, model: str, @@ -264,6 +277,7 @@ def _map_openai_exception( exception_provider: str, extra_information: str, ) -> None: + response: Final = _litellm_proxy_response(original_exception, custom_llm_provider) # custom_llm_provider is openai, make it OpenAI message = get_error_message(error_obj=original_exception) if message is None: @@ -292,14 +306,14 @@ def _map_openai_exception( message=f"RateLimitError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, ) elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str): raise ContextWindowExceededError( message=f"ContextWindowExceededError: {exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif "invalid_request_error" in error_str and "model_not_found" in error_str: @@ -307,7 +321,7 @@ def _map_openai_exception( message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif "A timeout occurred" in error_str: @@ -326,8 +340,9 @@ def _map_openai_exception( message=f"ContentPolicyViolationError: {exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, + body=getattr(original_exception, "body", None), ) elif "invalid_encrypted_content" in error_str or "could not be verified" in error_str: helpful_message: Final = ( @@ -345,7 +360,7 @@ def _map_openai_exception( message=helpful_message, llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, body=getattr(original_exception, "body", None), ) @@ -354,7 +369,7 @@ def _map_openai_exception( message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, body=getattr(original_exception, "body", None), ) @@ -372,7 +387,7 @@ def _map_openai_exception( message=f"RateLimitError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif ( @@ -383,7 +398,7 @@ def _map_openai_exception( message=f"AuthenticationError: {exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif "Mistral API raised a streaming error" in error_str: @@ -402,15 +417,16 @@ def _map_openai_exception( message=f"{exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, + body=getattr(original_exception, "body", None), ) elif original_exception.status_code == 401: raise AuthenticationError( message=f"AuthenticationError: {exception_provider} - {message}", llm_provider=custom_llm_provider, model=model, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 404: @@ -418,7 +434,7 @@ def _map_openai_exception( message=f"NotFoundError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 408: @@ -433,7 +449,7 @@ def _map_openai_exception( message=f"{exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, body=getattr(original_exception, "body", None), ) @@ -442,7 +458,7 @@ def _map_openai_exception( message=f"RateLimitError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 500: @@ -450,7 +466,7 @@ def _map_openai_exception( message=f"InternalServerError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 502: @@ -458,7 +474,7 @@ def _map_openai_exception( message=f"BadGatewayError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 503: @@ -466,7 +482,7 @@ def _map_openai_exception( message=f"ServiceUnavailableError: {exception_provider} - {message}", model=model, llm_provider=custom_llm_provider, - response=getattr(original_exception, "response", None), + response=response, litellm_debug_info=extra_information, ) elif original_exception.status_code == 504: # gateway timeout error @@ -2423,10 +2439,11 @@ def exception_type( custom_llm_provider == "litellm_proxy" ): # handle special case where calling litellm proxy + exception str contains error message extract_and_raise_litellm_exception( - response=getattr(original_exception, "response", None), + response=_litellm_proxy_response(mappable_exception, custom_llm_provider), error_str=error_str, model=model, custom_llm_provider=custom_llm_provider, + body=getattr(original_exception, "body", None), ) if ( custom_llm_provider == "openai" diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 8fc428b38ae..baa9aab1087 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -484,11 +484,12 @@ def apply_off_peak_pricing(model_info: ModelInfo, current_time: datetime | None, def _apply_off_peak_to_base_costs( model_info: ModelInfo, current_time: datetime | None, - base_costs: tuple[float, float, float, float, float], + base_costs: tuple[float, float, float, float | None, float], ) -> tuple[float, float, float, float, float]: """Apply off-peak rates to an already-resolved set of base costs, whichever pricing path - produced them. The one-hour cache-creation rate passes through untouched, since - off_peak_pricing has no field for it, and reasoning is left to _resolve_billed_reasoning_rate. + produced them. off_peak_pricing has no field for the one-hour cache-creation rate, so a + present one passes through untouched and an absent one resolves to the applied + cache-creation rate. Reasoning is left to _resolve_billed_reasoning_rate. """ prompt, completion, cache_creation, cache_creation_above_1hr, cache_read = base_costs rates: Final = apply_off_peak_pricing( @@ -506,7 +507,7 @@ def _apply_off_peak_to_base_costs( rates.input_rate, rates.output_rate, rates.cache_creation_rate, - cache_creation_above_1hr, + rates.cache_creation_rate if cache_creation_above_1hr is None else cache_creation_above_1hr, rates.cache_read_rate, ) @@ -532,6 +533,11 @@ def _get_token_base_cost( `missing_cache_read_uses_input` resolves an absent cache-read rate to the resolved input rate instead of 0.0; an explicit 0.0 rate stays a real price either way. + An absent cache-creation rate always resolves to the resolved input rate, the way the + tiered table and custom deployment pricing already do, since a provider that publishes + no write price bills cache writes as ordinary input. An absent 1h write rate resolves + to the cache-creation rate, off-peak included. An explicit 0.0 stays a real price for both. + Returns: Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost) """ @@ -554,10 +560,9 @@ def _get_token_base_cost( output_image_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_image_token", None) if output_image_cost is not None: completion_base_cost = cast(float, output_image_cost) - cache_creation_cost = cast(float, _get_cost_per_unit(model_info, cache_creation_cost_key)) - cache_creation_cost_above_1hr = cast( - float, - _get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"), + cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None) + cache_creation_cost_above_1hr = _get_cost_per_unit( + model_info, "cache_creation_input_token_cost_above_1hr", default_value=None ) cache_read_cost = _get_cost_per_unit(model_info, cache_read_cost_key, default_value=None) @@ -639,22 +644,10 @@ def _get_token_base_cost( else f"cache_read_input_token_cost_above_{threshold_str}_tokens" ) - cache_creation_cost = cast( - float, - _get_cost_per_unit( - model_info, - cache_creation_tiered_key, - cache_creation_cost, - ), - ) + cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost) - cache_creation_cost_above_1hr = cast( - float, - _get_cost_per_unit( - model_info, - cache_creation_1hr_tiered_key, - cache_creation_cost_above_1hr, - ), + cache_creation_cost_above_1hr = _get_cost_per_unit( + model_info, cache_creation_1hr_tiered_key, cache_creation_cost_above_1hr ) cache_read_cost = _get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost) @@ -665,16 +658,16 @@ def _get_token_base_cost( except Exception: continue + input_rate_for_missing_cache_rates: Final = _off_peak_rate( + _open_off_peak_block(model_info, current_time) or MappingProxyType({}), + "input_cost_per_token", + prompt_base_cost, + ) if cache_read_cost is None: - cache_read_cost = ( - _off_peak_rate( - _open_off_peak_block(model_info, current_time) or MappingProxyType({}), - "input_cost_per_token", - prompt_base_cost, - ) - if missing_cache_read_uses_input - else 0.0 - ) + cache_read_cost = input_rate_for_missing_cache_rates if missing_cache_read_uses_input else 0.0 + resolved_cache_creation_cost: Final = ( + input_rate_for_missing_cache_rates if cache_creation_cost is None else cache_creation_cost + ) return _apply_off_peak_to_base_costs( model_info, @@ -682,7 +675,7 @@ def _get_token_base_cost( ( prompt_base_cost, completion_base_cost, - cache_creation_cost, + resolved_cache_creation_cost, cache_creation_cost_above_1hr, cache_read_cost, ), @@ -956,12 +949,16 @@ def _calculate_input_cost( ) ### AUDIO COST - if prompt_tokens_details["audio_tokens"]: + if prompt_tokens_details["audio_tokens"] and not ( + prompt_tokens_details["audio_length_seconds"] and model_info.get("input_cost_per_audio_per_second") is not None + ): audio_cost_key: Final = _get_service_tier_cost_key("input_cost_per_audio_token", service_tier) prompt_cost += calculate_cost_component(model_info, audio_cost_key, prompt_tokens_details["audio_tokens"]) ### IMAGE TOKEN COST - if prompt_tokens_details["image_tokens"]: + if prompt_tokens_details["image_tokens"] and not ( + prompt_tokens_details["image_count"] and model_info.get("input_cost_per_image") is not None + ): # For image token costs: # First check if input_cost_per_image_token is available. If not, default to generic input_cost_per_token. image_token_cost_key = "input_cost_per_image_token" @@ -970,7 +967,9 @@ def _calculate_input_cost( prompt_cost += calculate_cost_component(model_info, image_token_cost_key, prompt_tokens_details["image_tokens"]) ### VIDEO TOKEN COST - if prompt_tokens_details["video_tokens"]: + if prompt_tokens_details["video_tokens"] and not ( + prompt_tokens_details["video_length_seconds"] and model_info.get("input_cost_per_video_per_second") is not None + ): video_token_cost_key = "input_cost_per_video_token" if model_info.get(video_token_cost_key) is None: video_token_cost_key = "input_cost_per_token" diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index ece619e3883..21ae8b001dd 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1757,7 +1757,7 @@ def convert_to_anthropic_tool_invoke( anthropic_tool_invoke: Final[list[AnthropicMessagesToolUseParam | dict[str, object]]] = [] for tool in tool_calls: - if not get_attribute_or_key(tool, "type") == "function": + if get_attribute_or_key(tool, "type") != "function": continue tool_id = cast(str, get_attribute_or_key(tool, "id")) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 3b128899f45..4c61fac82bb 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -454,7 +454,7 @@ def token_counter( params: Final = _MessageCountParams(model, custom_tokenizer) num_tokens = _count_messages(params, new_messages, use_default_image_token_count, default_token_count) if count_response_tokens is False: - includes_system_message: Final = any([message.get("role", None) == "system" for message in new_messages]) + includes_system_message: Final = any(message.get("role", None) == "system" for message in new_messages) num_tokens += _count_extra(params.count_function, tools, tool_choice, includes_system_message) else: diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 9d50345d70d..2ea20143f0c 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -20,6 +20,7 @@ from itertools import chain, repeat from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable +from pydantic import TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict, assert_never from litellm._logging import verbose_proxy_logger @@ -44,6 +45,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( scoped_structured_message_indices, stream_item_field, stream_item_fingerprint, + unappliable_request_rewrite, ) from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, @@ -103,9 +105,24 @@ class ToolResultBlockTextTarget: block_idx: int -InputWriteBackTarget = ( - MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget -) +@dataclass(frozen=True, slots=True) +class SystemStringTarget: + pass + + +@dataclass(frozen=True, slots=True) +class SystemBlockTextTarget: + block_idx: int + + +@dataclass(frozen=True, slots=True) +class ToolUseInputTarget: + msg_idx: int + content_idx: int + + +MessageTextTarget = MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget +InputWriteBackTarget = SystemStringTarget | SystemBlockTextTarget | MessageTextTarget def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]: @@ -146,10 +163,17 @@ class ScannedText: target: InputWriteBackTarget +@dataclass(frozen=True, slots=True) +class ScannedToolCall: + tool_call: ChatCompletionToolCallChunk + target: ToolUseInputTarget + + @dataclass(frozen=True, slots=True) class ExtractedInput: scanned: tuple[ScannedText, ...] images: tuple[str, ...] + tool_calls: tuple[ScannedToolCall, ...] = () EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=()) @@ -161,6 +185,74 @@ class _ToolCallShape: arguments: str +def _is_client_tool_use(block: Mapping[str, object]) -> bool: + return ( + block.get("type") == "tool_use" + and isinstance(block.get("id"), str) + and isinstance(block.get("name"), str) + and isinstance(block.get("input"), dict) + ) + + +def _write_back_system_block(system: object, block_idx: int, response: str) -> None: + if not isinstance(system, list): + return + text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text") + if block_idx < len(text_blocks): + text_blocks[block_idx]["text"] = ( + response # mutable-ok: guardrails rewrite the caller's request payload in place + ) + + +def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None: + content: Final = message.get("content", None) + if content is None: + return + match target: + case MessageContentTarget(): + if isinstance(content, str): + message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place + case ContentBlockTextTarget(content_idx=content_idx): + if isinstance(content, list): + content[content_idx]["text"] = ( + response # mutable-ok: guardrails rewrite the caller's request payload in place + ) + case ToolResultStringTarget(content_idx=content_idx): + if isinstance(content, list): + content[content_idx]["content"] = ( + response # mutable-ok: guardrails rewrite the caller's request payload in place + ) + case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx): + if isinstance(content, list): + content[content_idx]["content"][block_idx]["text"] = ( + response # mutable-ok: guardrails rewrite the caller's request payload in place + ) + case _: + assert_never(target) + + +_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object]) + + +def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None: + try: + return _TOOL_USE_INPUT_ADAPTER.validate_json(arguments) + except ValidationError: + return None + + +def _write_back_tool_use( + message: _WritableMessage, target: ToolUseInputTarget, shape: _ToolCallShape, rewritten_input: Mapping[str, object] +) -> None: + content: Final = message.get("content", None) + block: Final = content[target.content_idx] if isinstance(content, list) else None + if not isinstance(block, dict): + return + block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place + if shape.name is not None and shape.name != block.get("name"): + block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place + + @dataclass(frozen=True, slots=True) class _SSEFieldRewrite: """One field of one nested section of a buffered SSE event, rewritten.""" @@ -452,9 +544,8 @@ class AnthropicMessagesHandler(BaseTranslation): skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply) scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply) - # Exclude only the trusted top-level prompt. In-sequence system entries are untrusted - # and must stay aligned with texts_to_check for positional masking. When the top-level - # prompt is included, the pre-existing count mismatch disables positional masking. + # The top-level prompt is translated on its own below so it can be hoisted in front of + # any mid-turn system entries and scanned first, aligned with that structured position. translation_source: Final = { # mutable-ok: API message payload key: value for key, value in data.items() if key != "system" } @@ -490,7 +581,12 @@ class AnthropicMessagesHandler(BaseTranslation): ] ) - # Step 1: Extract all text content and images + # Step 1: Extract all text content, images, and tool calls + top_level_system_scanned: Final = ( + () + if hoisted_system_message is None or scan_only_tool_results + else self._extract_top_level_system_text(hoisted_system_message) + ) extracted: Final = tuple( self._extract_input_text_and_images( message=message, @@ -501,17 +597,27 @@ class AnthropicMessagesHandler(BaseTranslation): ) for msg_idx, message in enumerate(messages) ) - scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned) + scanned: Final = ( + *top_level_system_scanned, + *(item for one_message in extracted for item in one_message.scanned), + ) texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str] images_to_check: Final = [ image for one_message in extracted for image in one_message.images ] # mutable-ok: GenericGuardrailAPIInputs takes list[str] + scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls) + tool_calls_to_check: Final = [ + item.tool_call for item in scanned_tool_calls + ] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk] + pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check) - # Step 2: Apply guardrail to all texts in batch - if texts_to_check: + # Step 2: Apply guardrail to all texts and tool calls in batch + if texts_to_check or tool_calls_to_check: inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check) if images_to_check: inputs["images"] = images_to_check + if tool_calls_to_check: + inputs["tool_calls"] = tool_calls_to_check if tools_to_check: inputs["tools"] = tools_to_check original_structured_messages: Final = structured_messages @@ -570,9 +676,18 @@ class AnthropicMessagesHandler(BaseTranslation): preserve_system_messages=has_midturn_system_message, ) else: + if guardrailed_texts and len(guardrailed_texts) != len(scanned): + raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name) + self._apply_guardrail_tool_calls_to_input( + messages=messages, + scanned_tool_calls=scanned_tool_calls, + pre_guardrail_tool_calls=pre_guardrail_tool_calls, + returned_tool_calls=guardrailed_inputs.get("tool_calls"), + guardrail_name=guardrail_to_apply.guardrail_name, + ) # Step 3: Map guardrail responses back to original message structure await self._apply_guardrail_responses_to_input( - messages=messages, + data=data, responses=guardrailed_texts, scanned=scanned, ) @@ -598,6 +713,19 @@ class AnthropicMessagesHandler(BaseTranslation): hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload return hoisted[0] if hoisted else None + @staticmethod + def _extract_top_level_system_text(hoisted_system_message: AllMessageValues) -> tuple[ScannedText, ...]: + content: Final = hoisted_system_message.get("content") + if isinstance(content, str): + return (ScannedText(content, SystemStringTarget()),) + if not isinstance(content, list): + return () + return tuple( + ScannedText(text_str, SystemBlockTextTarget(block_idx)) + for block_idx, block in enumerate(content) + if isinstance(block, dict) and isinstance(text_str := block.get("text"), str) + ) + @staticmethod def _openai_system_message_to_anthropic( message: Mapping[str, object], @@ -852,9 +980,25 @@ class AnthropicMessagesHandler(BaseTranslation): for content_idx, content_item in enumerate(content) if isinstance(content_item, dict) ) + tool_use_blocks: Final = ( + () + if scan_only_tool_results + else tuple( + (content_idx, content_item) + for content_idx, content_item in enumerate(content) + if isinstance(content_item, dict) and _is_client_tool_use(content_item) + ) + ) return ExtractedInput( scanned=tuple(item for block in blocks for item in block.scanned), images=tuple(image for block in blocks for image in block.images), + tool_calls=tuple( + ScannedToolCall( + tool_call=AnthropicConfig.convert_tool_use_to_openai_format(content_item, tool_call_idx), + target=ToolUseInputTarget(msg_idx, content_idx), + ) + for tool_call_idx, (content_idx, content_item) in enumerate(tool_use_blocks) + ), ) @classmethod @@ -940,43 +1084,59 @@ class AnthropicMessagesHandler(BaseTranslation): async def _apply_guardrail_responses_to_input( self, - messages: Sequence[_WritableMessage], - responses: list[str], + data: dict[str, object], # mutable-ok: API message payload + responses: Sequence[str], scanned: tuple[ScannedText, ...], ) -> None: """ - Apply guardrail responses back to input messages. + Apply guardrail responses back to the top-level system prompt and the input messages. """ + raw_messages: Final = data.get("messages") + messages: Final[Sequence[_WritableMessage]] = raw_messages if isinstance(raw_messages, list) else () for item, guardrail_response in zip(scanned, responses): - target = item.target - message = messages[target.msg_idx] - content = message.get("content", None) - if content is None: - continue - - match target: - case MessageContentTarget(): - if isinstance(content, str): - message["content"] = ( - guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place - ) - case ContentBlockTextTarget(content_idx=content_idx): - if isinstance(content, list): - content[content_idx]["text"] = ( - guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place - ) - case ToolResultStringTarget(content_idx=content_idx): - if isinstance(content, list): - content[content_idx]["content"] = ( - guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place - ) - case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx): - if isinstance(content, list): - content[content_idx]["content"][block_idx]["text"] = ( + match item.target: + case SystemStringTarget(): + if isinstance(data.get("system"), str): + data["system"] = ( guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place ) + case SystemBlockTextTarget(block_idx=block_idx): + _write_back_system_block(data.get("system"), block_idx, guardrail_response) + case ( + MessageContentTarget() + | ContentBlockTextTarget() + | ToolResultStringTarget() + | ToolResultBlockTextTarget() as message_target + ): + _write_back_message_text(messages[message_target.msg_idx], message_target, guardrail_response) case _: - assert_never(target) + assert_never(item.target) + + @staticmethod + def _apply_guardrail_tool_calls_to_input( + messages: Sequence[_WritableMessage], + scanned_tool_calls: tuple[ScannedToolCall, ...], + pre_guardrail_tool_calls: tuple[_ToolCallShape, ...], + returned_tool_calls: Sequence[object] | None, + guardrail_name: str | None, + ) -> None: + post_guardrail_tool_calls: Final = _tool_call_shapes( + returned_tool_calls + if returned_tool_calls is not None and len(returned_tool_calls) == len(pre_guardrail_tool_calls) + else tuple(item.tool_call for item in scanned_tool_calls) + ) + rewritten: Final = tuple( + (item, after, _rewritten_tool_use_input(after.arguments)) + for item, before, after in zip(scanned_tool_calls, pre_guardrail_tool_calls, post_guardrail_tool_calls) + if before != after + ) + applicable: Final = tuple( + (item, after, rewritten_input) for item, after, rewritten_input in rewritten if rewritten_input is not None + ) + if len(applicable) != len(rewritten): + raise unappliable_request_rewrite(guardrail_name) + for item, after, rewritten_input in applicable: + _write_back_tool_use(messages[item.target.msg_idx], item.target, after, rewritten_input) async def process_output_response( self, diff --git a/litellm/llms/anthropic/prompt_cache_prediction.py b/litellm/llms/anthropic/prompt_cache_prediction.py new file mode 100644 index 00000000000..e69a02bd93a --- /dev/null +++ b/litellm/llms/anthropic/prompt_cache_prediction.py @@ -0,0 +1,388 @@ +from __future__ import annotations + +import hashlib +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from itertools import accumulate +from types import MappingProxyType +from typing import Annotated, Final, Literal, Protocol, TypeAlias + +import httpx +from pydantic import BaseModel, ConfigDict, Field, JsonValue, StrictInt, TypeAdapter, ValidationError + +import litellm +from litellm.llms.anthropic.common_utils import AnthropicModelInfo, is_anthropic_oauth_key +from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import DEFAULT_ANTHROPIC_API_VERSION +from litellm.types.router import LiteLLM_Params +from litellm.types.utils import ModelResponse + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_HEADERS: Final = TypeAdapter(dict[str, str]) +_counter: Final = AnthropicCountTokensHandler() + + +_NATIVE_HEADERS: Final = frozenset( + ( + "host", + "accept", + "accept-encoding", + "connection", + "user-agent", + "content-length", + "content-type", + "x-api-key", + "anthropic-version", + ) +) + +_DEPLOYMENT_OPTIONS: Final = frozenset( + { + "model", + "api_key", + "api_base", + "custom_llm_provider", + "rpm", + "tpm", + "timeout", + "stream_timeout", + "max_retries", + "num_retries", + "max_parallel_requests", + "input_cost_per_token", + "output_cost_per_token", + "cache_read_input_token_cost", + "cache_creation_input_token_cost", + "cache_creation_input_token_cost_above_1hr", + } +) + + +class _StrictModel(BaseModel): + model_config = ConfigDict(extra="forbid", frozen=True, strict=True) + + +class _CacheControl(_StrictModel): + type: Literal["ephemeral"] + ttl: Literal["5m", "1h"] = "5m" + + +class _Text(_StrictModel): + type: Literal["text"] + text: str = Field(min_length=1, pattern=r"\S") + cache_control: _CacheControl | None = None + + +class _ToolUse(_StrictModel): + type: Literal["tool_use"] + id: str = Field(min_length=1) + name: str = Field(min_length=1) + input: Mapping[str, JsonValue] + cache_control: _CacheControl | None = None + + +class _ResultText(_StrictModel): + type: Literal["text"] + text: str + + +class _ToolResult(_StrictModel): + type: Literal["tool_result"] + tool_use_id: str = Field(min_length=1) + content: str | Annotated[tuple[_ResultText, ...], Field(strict=False)] + is_error: bool | None = None + cache_control: _CacheControl | None = None + + +_Block: TypeAlias = Annotated[_Text | _ToolUse | _ToolResult, Field(discriminator="type")] + + +class _Message(_StrictModel): + role: Literal["user", "assistant"] + content: str | Annotated[tuple[_Block, ...], Field(strict=False)] + + def blocks(self) -> tuple[_Text | _ToolUse | _ToolResult, ...]: + return (_Text(type="text", text=self.content),) if isinstance(self.content, str) else tuple(self.content) + + +class _Tool(_StrictModel): + name: str = Field(min_length=1) + description: str | None = None + input_schema: Mapping[str, JsonValue] + type: Literal["custom"] | None = None + + +class _Request(_StrictModel): + messages: tuple[_Message, ...] = Field(min_length=1, strict=False) + system: str | Annotated[tuple[_ResultText, ...], Field(strict=False)] | None = None + tools: Annotated[tuple[_Tool, ...], Field(strict=False)] | None = None + model: str | None = None + max_tokens: int | None = None + stream: bool | None = None + temperature: float | int | None = None + top_p: float | int | None = None + top_k: int | None = None + stop_sequences: Annotated[tuple[str, ...], Field(strict=False)] | None = None + metadata: Mapping[str, JsonValue] | None = None + + +@dataclass(frozen=True, slots=True) +class PromptPrefix: + prefix_body: Mapping[str, JsonValue] + fingerprint: str + fingerprints: tuple[str, ...] + ttl_seconds: int + + +def _digest(value: object) -> str: + return hashlib.sha256( + json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode() + ).hexdigest() + + +def _next_digest(previous: str, boundary: tuple[int, str, Mapping[str, JsonValue]]) -> str: + return _digest((previous, boundary)) + + +def parse_prompt(body: Mapping[str, JsonValue]) -> PromptPrefix | None: + try: + request: Final = _Request.model_validate(body) + blocks: Final = tuple(message.blocks() for message in request.messages) + except ValidationError: + return None + markers: Final = tuple( + (message_index, block_index, block.cache_control) + for message_index, message_blocks in enumerate(blocks) + for block_index, block in enumerate(message_blocks) + if block.cache_control is not None + ) + if len(markers) != 1: + return None + message_end, block_end, marker = markers[0] + normalized: Final = _JSON_OBJECT.validate_python(request.model_dump(mode="json", exclude_none=True)) + context: Final = MappingProxyType({key: normalized[key] for key in ("system", "tools") if key in normalized}) + boundaries: Final = tuple( + ( + message_index, + request.messages[message_index].role, + _JSON_OBJECT.validate_python( + block.model_dump(mode="json", exclude=MappingProxyType({"cache_control": True}), exclude_none=True) + ), + ) + for message_index, message_blocks in enumerate(blocks[: message_end + 1]) + for block_index, block in enumerate(message_blocks) + if message_index < message_end or block_index <= block_end + ) + hashes: Final = tuple( + accumulate(boundaries, _next_digest, initial=_digest((_JSON_OBJECT.validate_python(context), marker.ttl))) + )[1:] + prefix_messages: Final = tuple( + _Message( + role=request.messages[message_index].role, + content=tuple( + block + for block_index, block in enumerate(message_blocks) + if message_index < message_end or block_index <= block_end + ), + ) + for message_index, message_blocks in enumerate(blocks[: message_end + 1]) + ) + return PromptPrefix( + prefix_body=MappingProxyType( + _JSON_OBJECT.validate_python( + _Request(messages=prefix_messages, system=request.system, tools=request.tools).model_dump( + mode="json", exclude_none=True + ) + ) + ), + fingerprint=hashes[-1], + fingerprints=tuple(reversed(hashes[-20:])), + ttl_seconds=3600 if marker.ttl == "1h" else 300, + ) + + +def cache_scope( + caller_key_hash: str, + deployment_id: str, + provider_key: str, + model: str, + anthropic_version: str = DEFAULT_ANTHROPIC_API_VERSION, +) -> str: + return _digest((caller_key_hash, deployment_id, provider_key, model, anthropic_version)) + + +class _TTLUsage(BaseModel): + model_config = ConfigDict(strict=True) + ephemeral_5m_input_tokens: int = Field(default=0, ge=0) + ephemeral_1h_input_tokens: int = Field(default=0, ge=0) + + +class _CacheUsage(BaseModel): + model_config = ConfigDict(strict=True) + cached_tokens: int = Field(default=0, ge=0) + cache_creation_tokens: int = Field(default=0, ge=0) + cache_creation_token_details: _TTLUsage | None = None + + +class _Usage(BaseModel): + model_config = ConfigDict(strict=True) + prompt_tokens: int = Field(ge=0) + prompt_tokens_details: _CacheUsage + + +class _Choice(BaseModel): + finish_reason: str = Field(min_length=1) + + +class _Response(BaseModel): + model_config = ConfigDict(strict=True) + model: str + usage: _Usage + choices: tuple[_Choice, ...] = Field(min_length=1, strict=False) + + +class _CountBody(BaseModel): + messages: Sequence[Mapping[str, JsonValue]] + tools: Sequence[Mapping[str, JsonValue]] | None = None + system: str | Sequence[Mapping[str, JsonValue]] | None = None + + +class _CountResult(BaseModel): + input_tokens: Annotated[StrictInt, Field(ge=0)] + + +class TokenCounter(Protocol): + async def __call__(self, model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: ... + + +def _count_objects( + values: Sequence[Mapping[str, JsonValue]], +) -> list[dict[str, JsonValue]]: # mutable-ok: the existing provider count API requires JSON lists/dicts + return [dict(value) for value in values] # mutable-ok: serialize read-only inputs at the provider API boundary + + +async def count_prompt_tokens(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + native: Final = _CountBody.model_validate(body) + try: + result: Final = _CountResult.model_validate( + await _counter.handle_count_tokens_request( + model=model, + messages=_count_objects(native.messages), + tools=_count_objects(native.tools) if native.tools is not None else None, + system=native.system, + api_key=api_key, + timeout=15.0, + ) + ) + except Exception: # noqa: BLE001 # provider/count validation failures are unavailable estimates, not zero tokens + return None + return result.input_tokens + + +@dataclass(frozen=True, slots=True) +class NativePredictionTarget: + model: str + api_key: str + + +@dataclass(frozen=True, slots=True) +class UnsupportedPredictionTarget: + reason: Literal[ + "unsupported_deployment_configuration", + "unsupported_provider_endpoint", + "unsupported_provider", + "unsupported_provider_credentials", + ] + + +def resolve_prediction_target(params: LiteLLM_Params) -> NativePredictionTarget | UnsupportedPredictionTarget: + configured_options: Final = frozenset(params.model_dump(exclude_defaults=True, exclude_none=True)) + if configured_options - _DEPLOYMENT_OPTIONS: + return UnsupportedPredictionTarget("unsupported_deployment_configuration") + api_base: Final = AnthropicModelInfo.get_api_base(params.api_base) + if api_base not in ("https://api.anthropic.com", "https://api.anthropic.com/v1/messages"): + return UnsupportedPredictionTarget("unsupported_provider_endpoint") + try: + model, provider, _, _ = litellm.get_llm_provider( + model=params.model, custom_llm_provider=params.custom_llm_provider + ) + except Exception: # noqa: BLE001 # the shared provider resolver raises for unknown deployments + return UnsupportedPredictionTarget("unsupported_provider") + if provider != "anthropic": + return UnsupportedPredictionTarget("unsupported_provider") + api_key: Final = AnthropicModelInfo.get_api_key(params.api_key) + if api_key is None or not _supported_provider_key(api_key): + return UnsupportedPredictionTarget("unsupported_provider_credentials") + return NativePredictionTarget(model=model, api_key=api_key) + + +def _supported_provider_key(api_key: str) -> bool: + return bool(api_key) and not is_anthropic_oauth_key(api_key) + + +def supported_prediction_headers(headers: Mapping[str, str]) -> bool: + return all( + name.lower() != "anthropic-beta" + and (name.lower() != "anthropic-version" or value == DEFAULT_ANTHROPIC_API_VERSION) + for name, value in headers.items() + ) + + +@dataclass(frozen=True, slots=True) +class ObservedCachePrefix: + prefix: PromptPrefix + scope: str + cached_tokens: int + cache_creation_tokens: int + + +def parse_observed_cache( + wire: httpx.Request, response_obj: ModelResponse, caller_key_hash: str, deployment_id: str +) -> ObservedCachePrefix | None: + try: + response: Final = _Response.model_validate(response_obj, from_attributes=True) + body: Final = _JSON_OBJECT.validate_json(wire.content) + headers: Final = _HEADERS.validate_python(wire.headers) + except (ValidationError, RuntimeError, httpx.RequestNotRead): + return None + if ( + wire.url.scheme != "https" + or wire.url.host != "api.anthropic.com" + or wire.url.path != "/v1/messages" + or wire.url.query + or wire.url.port not in (None, 443) + ): + return None + if ( + frozenset(headers) - _NATIVE_HEADERS + or not supported_prediction_headers(headers) + or headers.get("anthropic-version") != DEFAULT_ANTHROPIC_API_VERSION + ): + return None + provider_key: Final = headers.get("x-api-key", "") + model: Final = body.get("model") + if not _supported_provider_key(provider_key) or not isinstance(model, str) or model != response.model: + return None + prefix: Final = parse_prompt(body) + if prefix is None: + return None + usage: Final = response.usage.prompt_tokens_details + cache_tokens: Final = usage.cached_tokens + usage.cache_creation_tokens + if cache_tokens <= 0 or cache_tokens > response.usage.prompt_tokens: + return None + split: Final = usage.cache_creation_token_details + if usage.cache_creation_tokens and split is None: + return None + if split is not None and ( + split.ephemeral_5m_input_tokens + split.ephemeral_1h_input_tokens != usage.cache_creation_tokens + or (prefix.ttl_seconds == 300 and split.ephemeral_1h_input_tokens > 0) + or (prefix.ttl_seconds == 3600 and split.ephemeral_5m_input_tokens > 0) + ): + return None + return ObservedCachePrefix( + prefix=prefix, + scope=cache_scope(caller_key_hash, deployment_id, provider_key, model), + cached_tokens=cache_tokens, + cache_creation_tokens=usage.cache_creation_tokens, + ) diff --git a/litellm/llms/azure_ai/common_utils.py b/litellm/llms/azure_ai/common_utils.py index 0665b3f64c5..53a864a880a 100644 --- a/litellm/llms/azure_ai/common_utils.py +++ b/litellm/llms/azure_ai/common_utils.py @@ -144,10 +144,13 @@ class AzureFoundryModelInfo(BaseLLMModelInfo): def get_api_key(api_key: str | None = None) -> str | None: return api_key or litellm.api_key or get_secret_str("AZURE_AI_API_KEY") + @staticmethod + def get_api_version(api_version: str | None = None) -> str | None: + return api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") + @property - def api_version(self, api_version: str | None = None) -> str | None: - api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") - return api_version + def api_version(self) -> str | None: + return AzureFoundryModelInfo.get_api_version() def get_token_counter(self) -> BaseTokenCounter | None: """ diff --git a/litellm/llms/base_llm/guardrail_translation/utils.py b/litellm/llms/base_llm/guardrail_translation/utils.py index 94a780f8148..51d43436fc9 100644 --- a/litellm/llms/base_llm/guardrail_translation/utils.py +++ b/litellm/llms/base_llm/guardrail_translation/utils.py @@ -1,8 +1,8 @@ from __future__ import annotations import json -from collections.abc import Callable, Iterator, Sequence -from typing import Final, TypeVar +from collections.abc import Callable, Iterator, Mapping, Sequence +from typing import Final, TypeVar, cast # noqa: TID251 # a rebuilt chat row has no typed constructor across roles from pydantic import BaseModel @@ -364,3 +364,67 @@ def merge_guardrailed_scoped_messages( yield from appended return list(_merged()) + + +def _content_part_text(part: object) -> str | None: + if not isinstance(part, Mapping): + return None + text: Final = part.get("text") + return text if isinstance(text, str) else None + + +def message_slot_texts(message: Mapping[str, object]) -> tuple[str, ...]: + content: Final = message.get("content") + if isinstance(content, str): + return (content,) + if isinstance(content, list): + return tuple(text for part in content if (text := _content_part_text(part)) is not None) + return () + + +def message_text_slot_count(message: AllMessageValues) -> int: + return len(message_slot_texts(message)) + + +def _part_with_text(part: object, text: str) -> object: + if not isinstance(part, Mapping): + return part + return {**part, "text": text} # mutable-ok: content parts stay JSON-plain dicts + + +def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> Sequence[object]: + remaining_texts: Final = iter(texts) + return [ # mutable-ok: message content stays a JSON list + _part_with_text(part, next(remaining_texts)) if _content_part_text(part) is not None else part + for part in content + ] + + +def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None: + """Swap one rewritten text into each text slot of a chat row, in order. + + A slot is a string ``content`` or one list part carrying a string ``text``; + images and other parts ride along untouched. Returns None unless the counts + line up exactly, so a rewrite never lands on the wrong slot. + """ + if message_text_slot_count(message) != len(texts): + return None + content: Final = message.get("content") + if not isinstance(content, (str, list)): + return message + rewritten_content: Final = texts[0] if isinstance(content, str) else _content_with_slot_texts(content, texts) + rewritten: Final = {**message, "content": rewritten_content} # mutable-ok: chat rows stay JSON-plain dicts + return cast("AllMessageValues", rewritten) # cast-ok: the same row with only its text slots swapped + + +class UnappliableRequestRewrite(Exception): + def __init__(self, guardrail_name: str) -> None: + super().__init__( + f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, " + "so the request was rejected rather than sent unrewritten" + ) + self.guardrail_name: Final = guardrail_name + + +def unappliable_request_rewrite(guardrail_name: str | None) -> UnappliableRequestRewrite: + return UnappliableRequestRewrite(guardrail_name or "unknown") diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index f52c1cec6a8..385d5898569 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -11,6 +11,7 @@ from concurrent.futures import ThreadPoolExecutor from datetime import datetime from functools import partial from threading import Lock +from types import MappingProxyType from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload import httpx @@ -96,6 +97,77 @@ def _assume_role_params( ) +_SecureTransportBool = TypedDict("_SecureTransportBool", {"aws:SecureTransport": ReadOnly[Literal["true"]]}) + + +class _SecureTransportCondition(TypedDict): + Bool: ReadOnly[_SecureTransportBool] + + +class _SessionPolicyStatement(TypedDict): + Sid: ReadOnly[str] + Effect: ReadOnly[Literal["Allow"]] + Action: ReadOnly[tuple[str, ...]] + Resource: ReadOnly[Literal["*"]] + Condition: ReadOnly[_SecureTransportCondition] + + +class WebIdentitySessionPolicy(TypedDict): + Version: ReadOnly[Literal["2012-10-17"]] + Statement: ReadOnly[tuple[_SessionPolicyStatement, ...]] + + +_WEB_IDENTITY_SESSION_POLICY_ACTIONS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "BedrockLiteLLM": ( + "bedrock:InvokeModel", + "bedrock:InvokeModelWithResponseStream", + "bedrock:CountTokens", + "bedrock:Rerank", + "bedrock:Retrieve", + "bedrock:ListKnowledgeBases", + "bedrock:InvokeAgent", + "bedrock:ApplyGuardrail", + "bedrock:GetGuardrail", + "bedrock:ListGuardrails", + ), + "BedrockAgentCoreLiteLLM": ( + "bedrock-agentcore:InvokeAgentRuntime", + "bedrock-agentcore:InvokeAgentRuntimeForUser", + "bedrock-agentcore:InvokeGateway", + ), + "ClaudePlatformLiteLLM": ( + "aws-external-anthropic:CreateInference", + "aws-external-anthropic:CreateBatchInference", + "aws-external-anthropic:CancelBatchInference", + "aws-external-anthropic:DeleteBatchInference", + "aws-external-anthropic:CountTokens", + "aws-external-anthropic:Get*", + "aws-external-anthropic:List*", + ), + "BedrockMantleLiteLLM": ("bedrock-mantle:CreateInference",), + } +) + +_SECURE_TRANSPORT_ONLY: Final = _SecureTransportCondition(Bool=_SecureTransportBool({"aws:SecureTransport": "true"})) + + +def build_web_identity_session_policy() -> WebIdentitySessionPolicy: + return WebIdentitySessionPolicy( + Version="2012-10-17", + Statement=tuple( + _SessionPolicyStatement( + Sid=sid, + Effect="Allow", + Action=actions, + Resource="*", + Condition=_SECURE_TRANSPORT_ONLY, + ) + for sid, actions in _WEB_IDENTITY_SESSION_POLICY_ACTIONS.items() + ), + ) + + class BedrockRequestTarget(BaseModel): aws_region_name: str aws_bedrock_runtime_endpoint: str | None @@ -940,60 +1012,12 @@ class BaseAWSLLM(SignsRequestsWithAWS): # auth only (static creds + IRSA take other code paths). # https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html # https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html - bedrock_session_policy: Final = { - "Version": "2012-10-17", - "Statement": [ - { - "Sid": "BedrockLiteLLM", - "Effect": "Allow", - "Action": [ - "bedrock:InvokeModel", - "bedrock:InvokeModelWithResponseStream", - "bedrock:CountTokens", - "bedrock:ApplyGuardrail", - "bedrock:GetGuardrail", - "bedrock:ListGuardrails", - ], - "Resource": "*", - "Condition": {"Bool": {"aws:SecureTransport": "true"}}, - }, - # Claude Platform on AWS (added by #27678 for the - # ``bedrock/claude_platform/`` route) lives under - # a separate IAM action namespace; without these entries - # the OIDC path 403s on every claude_platform request - # even with a fully permissive identity policy (#30200). - { - "Sid": "ClaudePlatformLiteLLM", - "Effect": "Allow", - "Action": [ - "aws-external-anthropic:CreateInference", - "aws-external-anthropic:CreateBatchInference", - "aws-external-anthropic:CancelBatchInference", - "aws-external-anthropic:DeleteBatchInference", - "aws-external-anthropic:CountTokens", - "aws-external-anthropic:Get*", - "aws-external-anthropic:List*", - ], - "Resource": "*", - "Condition": {"Bool": {"aws:SecureTransport": "true"}}, - }, - { - "Sid": "BedrockMantleLiteLLM", - "Effect": "Allow", - "Action": [ - "bedrock-mantle:CreateInference", - ], - "Resource": "*", - "Condition": {"Bool": {"aws:SecureTransport": "true"}}, - }, - ], - } assume_role_params: Final = { "RoleArn": aws_role_name, "RoleSessionName": aws_session_name, "WebIdentityToken": oidc_token, "DurationSeconds": 3600, - "Policy": json.dumps(bedrock_session_policy, separators=(",", ":")), + "Policy": json.dumps(build_web_identity_session_policy(), separators=(",", ":")), } # Add ExternalId parameter if provided diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index 53a3e634adf..57590601a3c 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -17,7 +17,7 @@ BaseAWSLLM._sign_request after the request body is finalized. import json from collections.abc import Mapping -from typing import Any, Final +from typing import Any, Final, cast # noqa: TID251 # map_openai_params returns the filtered params as a bare dict import httpx from typing_extensions import ReadOnly, TypedDict @@ -32,6 +32,7 @@ from litellm.llms.bedrock_mantle.common_utils import ( BedrockMantleAuthMixin, ) from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.responses.additional_tools import HoistedAdditionalTools, hoist_additional_tools from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( ResponseInputParam, @@ -58,8 +59,6 @@ _BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES: Final = frozenset( _BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"}) _BEDROCK_MANTLE_OPENAI_PATH_SUPPORTED_REASONING_SUMMARIES: Final = frozenset({"auto"}) -_CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools" - _CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE: Final = "agent_message" _CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE: Final = "context_compaction" _CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: Final = "local_shell_call" @@ -233,17 +232,14 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI litellm_params: GenericLiteLLMParams, headers: dict, ) -> dict: - remaining_input, hoisted_tools = self._hoist_codex_additional_tools(input) - normalized_input: Final = self._normalize_codex_input_items(remaining_input) + params: Final = cast( # cast-ok: the base signature leaves the params dict untyped + "ResponsesAPIOptionalRequestParams", response_api_optional_request_params + ) + hoisted: Final = hoist_additional_tools(input, params.get("tools")) + normalized_input: Final = self._normalize_codex_input_items(hoisted.input) request_params: Final = ( - { - **response_api_optional_request_params, - "tools": [ - *(response_api_optional_request_params.get("tools") or []), - *hoisted_tools, - ], - } - if hoisted_tools + self._params_with_hoisted_tools(params, hoisted) + if hoisted.hoisted else response_api_optional_request_params ) return super().transform_responses_api_request( @@ -254,41 +250,14 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI headers=headers, ) - @staticmethod - def _is_codex_additional_tools_item(item: Any) -> bool: - return isinstance(item, dict) and item.get("type") == _CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE - - @staticmethod - def _tools_of_additional_tools_item(item: "dict[str, Any]") -> "list[Any]": - tools: Final = item.get("tools") - return tools if isinstance(tools, list) else [] - @classmethod - def _hoist_codex_additional_tools( - cls, - input: "str | ResponseInputParam", - ) -> "tuple[str | ResponseInputParam, list[Any]]": - """Codex's "responses lite" wire mode ships tool definitions inside - `input` as {"type": "additional_tools", "role": "developer", - "tools": [...]} items. api.openai.com accepts that item type; Mantle - rejects the whole request with 400 "Invalid 'input': value did not - match any expected variant" but accepts the same tools at the top - level, so move them there and strip the items from `input`. - """ - if not isinstance(input, list): - return input, [] - additional_tools_items: Final = [item for item in input if cls._is_codex_additional_tools_item(item)] - if not additional_tools_items: - return input, [] - remaining_input: Final = [item for item in input if not cls._is_codex_additional_tools_item(item)] - hoisted_tools = [tool for item in additional_tools_items for tool in cls._tools_of_additional_tools_item(item)] - verbose_logger.debug( - "Bedrock Mantle Responses API: hoisting %d tool(s) out of %d 'additional_tools' input item(s) " - "into the top-level tools param (Mantle rejects that input item type).", - len(hoisted_tools), - len(additional_tools_items), - ) - return remaining_input, cls._filter_unsupported_tools(hoisted_tools) + def _params_with_hoisted_tools( + cls, params: Mapping[str, object], hoisted: HoistedAdditionalTools + ) -> dict[str, object]: + supported_tools: Final = cls._filter_unsupported_tools(list(hoisted.tools)) + if supported_tools: + return {**params, "tools": supported_tools} + return {key: value for key, value in params.items() if key != "tools"} @staticmethod def _agent_message_text(item: "Mapping[str, object]") -> str: diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index c92af7de145..79985569c5f 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -115,6 +115,19 @@ def _parse_setup(session_configuration_request: str) -> BidiGenerateContentSetup return envelope.get("setup", empty_setup) +def _grounding_metadata_from_frame(frame: Mapping[str, object]) -> tuple[Mapping[str, object], ...]: + """Read ``serverContent.groundingMetadata`` off the frame that carries the turn's usage. + + Live reports grounding in the server frames rather than in ``usageMetadata``, and it emits both + on the same frame, so the per-query charge is countable at the point usage is built. + """ + server_content: Final = frame.get("serverContent") + if not isinstance(server_content, Mapping): + return () + metadata: Final = server_content.get("groundingMetadata") + return (metadata,) if isinstance(metadata, Mapping) else () + + # Google bills Live transcription at an estimated 25 audio tokens/sec of input and # 175 text tokens/min of output (ai.google.dev/gemini-api/docs/pricing). GEMINI_LIVE_TRANSCRIBE_AUDIO_TOKENS_PER_SECOND: Final = 25 @@ -323,7 +336,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) elif key == "input_audio_transcription" and value is not None: optional_params["inputAudioTranscription"] = {} - elif key == "turn_detection": + elif key == "turn_detection" and value is not None: value_typed = cast(OpenAIRealtimeTurnDetection, value) if ( isinstance(value_typed, dict) @@ -1049,6 +1062,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): {**cast(dict, message), "usageMetadata": resolved_usage_metadata}, ), ) + grounding_metadata: Final = _grounding_metadata_from_frame(message) + if grounding_metadata: + VertexGeminiConfig._set_grounding_usage_counters( # pyright: ignore[reportPrivateUsage] # shared with the chat path; no public alias exists yet + _chat_completion_usage, grounding_metadata + ) else: _chat_completion_usage = get_empty_usage() diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 58ff03e6a0d..01e14f2248d 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -42,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( stream_item_field, stream_item_fingerprint, stream_item_items, + unappliable_request_rewrite, ) from litellm.main import stream_chunk_builder from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam @@ -196,6 +197,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): else: # Step 3: Map guardrail responses back to original message structure if guardrailed_texts and texts_to_check: + if len(guardrailed_texts) != len(text_task_mappings): + raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name) await self._apply_guardrail_responses_to_input_texts( messages=messages, responses=guardrailed_texts, @@ -210,6 +213,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation): task_mappings=tool_call_task_mappings, ) + elif ( + not images_to_check + and not guardrail_to_apply.records_own_guardrail_information + and (not_run_reason := self._not_run_reason(messages)) is not None + ): + guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=not_run_reason, + request_data=data, + guardrail_status="not_run", + ) + verbose_proxy_logger.debug( "OpenAI Chat Completions: Processed input messages: %s", data.get("messages"), @@ -217,6 +231,28 @@ class OpenAIChatCompletionsHandler(BaseTranslation): return data + def _not_run_reason( + self, + messages: Sequence[dict[str, Any]], # mutable-ok: raw request messages consumed by _extract_inputs + ) -> str | None: + """Why nothing was scanned, or None when the only unscoped content is images, which this handler never scans.""" + texts: Final[list[str]] = [] # mutable-ok: filled by _extract_inputs + images: Final[list[str]] = [] # mutable-ok: filled by _extract_inputs + tool_calls: Final[list[ChatCompletionToolParam]] = [] # mutable-ok: filled by _extract_inputs + for msg_idx, message in enumerate(messages): + self._extract_inputs( + message=message, + msg_idx=msg_idx, + texts_to_check=texts, + images_to_check=images, + tool_calls_to_check=tool_calls, + text_task_mappings=[], # mutable-ok: required by _extract_inputs, unused here + tool_call_task_mappings=[], # mutable-ok: required by _extract_inputs, unused here + ) + if texts or tool_calls: + return "no scannable content after message scoping" + return None if images else "no scannable content" + def extract_request_tool_names(self, data: dict) -> list[str]: """Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name).""" names: Final[list[str]] = [] diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 2fe11d9f7bd..27ff55f120c 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -56,6 +56,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( stream_item_field, stream_item_fingerprint, stream_item_items, + unappliable_request_rewrite, ) from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools from litellm.responses.litellm_completion_transformation.transformation import ( @@ -495,13 +496,13 @@ class OpenAIResponsesHandler(BaseTranslation): data["instructions"] = written_back.instructions # rebind-ok: data is an out-param elif isinstance(input_data, str): guardrailed_texts: Final = guardrailed_inputs.get("texts") or () + if len(guardrailed_texts) > 1: + raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name) data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param else: rewritten_texts: Final = guardrailed_inputs.get("texts") or () if len(rewritten_texts) != len(extracted.task_mappings): - from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite - - raise UnappliableRequestRewrite(guardrail_to_apply.guardrail_name or "unknown") + raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name) await self._apply_guardrail_responses_to_input( messages=input_data, responses=rewritten_texts, diff --git a/litellm/llms/openai/responses/guardrail_translation/tool_merge.py b/litellm/llms/openai/responses/guardrail_translation/tool_merge.py index b596adfad6f..ff67c6220e1 100644 --- a/litellm/llms/openai/responses/guardrail_translation/tool_merge.py +++ b/litellm/llms/openai/responses/guardrail_translation/tool_merge.py @@ -6,8 +6,10 @@ from typing import Final, TypeAlias from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_logger +from litellm.responses.litellm_completion_transformation.custom_tools import custom_tool_grammar_suffix from litellm.responses.litellm_completion_transformation.transformation import ( NAMESPACE_DESCRIPTION_SEPARATOR, + NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS, LiteLLMCompletionResponsesConfig, ) @@ -34,8 +36,8 @@ def _validated_tools(values: Iterable[object]) -> tuple[Tool, ...]: return tuple(tool for tool in validated if tool is not None) -def _is_function(tool: Tool) -> bool: - return tool.get("type") == "function" +def _has_chat_tool(member: Tool) -> bool: + return member.get("type") in NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS def _chat_tool_key(tool: Tool) -> str: @@ -67,18 +69,19 @@ def _function_fields(tool: Tool) -> Tool: return function if function is not None else MappingProxyType({}) -def _without_namespace_prefix(key: str, value: object, prefix: str) -> object: - if key != "description" or not isinstance(value, str) or not value.startswith(prefix): +def _member_description(key: str, value: object, prefix: str, suffix: str) -> object: + if key != "description" or not isinstance(value, str): return value - return value[len(prefix) :] + return value.replace(prefix, "", 1).replace(suffix, "", 1) def _rebuilt_member(member: Tool, flattened: Tool, guardrailed: Tool, namespace_description: str) -> Tool: flattened_function: Final = _function_fields(flattened) prefix: Final = f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}" if namespace_description else "" + suffix: Final = custom_tool_grammar_suffix(member.get("format")) if member.get("type") == "custom" else "" changed_function: Final = MappingProxyType( { - key: _without_namespace_prefix(key, value, prefix) + key: _member_description(key, value, prefix, suffix) for key, value in _function_fields(guardrailed).items() if flattened_function.get(key) != value } @@ -93,8 +96,8 @@ def _rebuilt_member(member: Tool, flattened: Tool, guardrailed: Tool, namespace_ return {**member, **changed_extras, **changed_function} # mutable-ok: json.dumps rejects MappingProxyType -def _rebuilt_function_members( - function_members: Sequence[Tool], +def _rebuilt_flattened_members( + flattened_members: Sequence[Tool], flattened_group: Sequence[Tool], group_keys: Sequence[IndexedKey], guardrailed_by_key: Mapping[IndexedKey, Tool], @@ -106,7 +109,7 @@ def _rebuilt_function_members( else member if guardrailed_by_key[key] == flattened else _rebuilt_member(member, flattened, guardrailed_by_key[key], namespace_description) - for member, flattened, key in zip(function_members, flattened_group, group_keys) + for member, flattened, key in zip(flattened_members, flattened_group, group_keys) ) @@ -118,9 +121,9 @@ def _rebuilt_namespace( guardrailed_by_key: Mapping[IndexedKey, Tool], ) -> tuple[Tool, ...]: namespace_description: Final = str(original.get("description") or "") - rebuilt_functions: Final = iter( - _rebuilt_function_members( - tuple(member for member in members if _is_function(member)), + rebuilt_flattened: Final = iter( + _rebuilt_flattened_members( + tuple(member for member in members if _has_chat_tool(member)), flattened_group, group_keys, guardrailed_by_key, @@ -129,7 +132,7 @@ def _rebuilt_namespace( ) rebuilt_members: Final = tuple( rebuilt - for rebuilt in (next(rebuilt_functions) if _is_function(member) else member for member in members) + for rebuilt in (next(rebuilt_flattened) if _has_chat_tool(member) else member for member in members) if rebuilt is not None ) if not rebuilt_members: @@ -149,7 +152,7 @@ def _merged_original( if guardrailed_group == tuple(flattened_group): return (original,) members: Final = _namespace_members(original) if original.get("type") == "namespace" else () - if members and sum(map(_is_function, members)) == len(flattened_group): + if members and sum(map(_has_chat_tool, members)) == len(flattened_group): return _rebuilt_namespace(original, members, flattened_group, group_keys, guardrailed_by_key) if not guardrailed_group: return () diff --git a/litellm/llms/vertex_ai/batches/transformation.py b/litellm/llms/vertex_ai/batches/transformation.py index e63c80dd3cf..f5f1ab2068a 100644 --- a/litellm/llms/vertex_ai/batches/transformation.py +++ b/litellm/llms/vertex_ai/batches/transformation.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Any, Final from urllib.parse import unquote @@ -8,7 +9,36 @@ from litellm.llms.vertex_ai.common_utils import ( ) from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest from litellm.types.llms.vertex_ai import * -from litellm.types.utils import LiteLLMBatch +from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper + + +def vertex_prompt_tokens_details( + usage_metadata: Mapping[str, object], +) -> PromptTokensDetailsWrapper | None: + raw_details: Final = usage_metadata.get("promptTokensDetails") + if not isinstance(raw_details, list): + return None + + def _normalize(detail: object) -> tuple[str, int] | None: + if not isinstance(detail, Mapping): + return None + modality: Final = detail.get("modality") + token_count: Final = detail.get("tokenCount") + if not isinstance(modality, str) or not isinstance(token_count, int): + return None + return modality.upper(), token_count + + parsed_details: Final = tuple(_normalize(detail) for detail in raw_details) + normalized: Final = tuple(detail for detail in parsed_details if detail is not None) + if len(normalized) != len(parsed_details): + return None + + return PromptTokensDetailsWrapper( + text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")), + audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"), + image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"), + video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"), + ) class VertexAIBatchTransformation: diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index e7fd9a0d08b..d669acecfd9 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -298,8 +298,6 @@ def transform_openai_input_gemini_embed_content( _IMAGE_MIME_TYPES: Final = frozenset({"image/png", "image/jpeg"}) -_VIDEO_TOKENS_PER_SECOND: Final = 258.0 -_AUDIO_TOKENS_PER_SECOND: Final = 32.0 _usage_metadata_adapter: Final = TypeAdapter(UsageMetadata) @@ -339,11 +337,12 @@ def _is_image_element( return False -def _count_input_images( +def _is_image_only_input( input: GeminiEmbeddingInput, resolved_files: Mapping[str, Mapping[str, str]], -) -> int: - return sum(1 for element in _flatten_input(input) if _is_image_element(element, resolved_files)) +) -> bool: + elements: Final = _flatten_input(input) + return bool(elements) and all(_is_image_element(element, resolved_files) for element in elements) def _tokens_for_modality(details: Sequence[PromptTokensDetails], modality: str) -> int: @@ -372,30 +371,29 @@ def _usage_from_embed_content_response( total_tokens: Final = usage_metadata.get("totalTokenCount") or prompt_tokens details: Final[Sequence[PromptTokensDetails]] = usage_metadata.get("promptTokensDetails") or () + if not details: + return Usage( + prompt_tokens=prompt_tokens, + total_tokens=total_tokens, + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=0, + image_tokens=prompt_tokens if _is_image_only_input(input, resolved_files) else 0, + ), + ) + text_tokens: Final = _tokens_for_modality(details, "TEXT") audio_tokens: Final = _tokens_for_modality(details, "AUDIO") + image_tokens: Final = _tokens_for_modality(details, "IMAGE") video_tokens: Final = _tokens_for_modality(details, "VIDEO") - image_count: Final = _count_input_images(input, resolved_files) - - video_length_seconds: Final = video_tokens / _VIDEO_TOKENS_PER_SECOND if video_tokens > 0 else 0.0 - audio_length_seconds: Final = audio_tokens / _AUDIO_TOKENS_PER_SECOND if audio_tokens > 0 else 0.0 - - # generic_cost_per_token rewrites text_tokens to the full prompt minus - # other modalities when both text_tokens and image_count are zero. For - # video, that misallocates video tokens to text; a 1-token floor sidesteps - # the rewrite and keeps billing on input_cost_per_video_per_second. - needs_video_text_floor: Final = video_length_seconds > 0 and text_tokens == 0 and image_count == 0 - resolved_text_tokens: Final = 1 if needs_video_text_floor else text_tokens return Usage( prompt_tokens=prompt_tokens, total_tokens=total_tokens, prompt_tokens_details=PromptTokensDetailsWrapper( - text_tokens=resolved_text_tokens, + text_tokens=text_tokens, audio_tokens=audio_tokens, - image_count=image_count, - video_length_seconds=video_length_seconds, - audio_length_seconds=audio_length_seconds, + image_tokens=image_tokens, + video_tokens=video_tokens, ), ) @@ -415,8 +413,7 @@ def process_embed_content_response( model_response: EmbeddingResponse to populate model: Model name response_json: Raw JSON response from embedContent endpoint - resolved_files: Mapping of file references (files/abc) to {mime_type, uri}, - used to bill resolved image references at the per-image rate + resolved_files: Mapping of file references to resolved metadata Returns: EmbeddingResponse with single embedding diff --git a/litellm/main.py b/litellm/main.py index 17edafcdfca..f6f4ec1bf63 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2592,7 +2592,9 @@ def _complete_custom_openai( copilot_headers.update(extra_headers) extra_headers = copilot_headers - if extra_headers is not None: + use_base_llm_http_handler: Final = get_secret_bool("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER") + + if extra_headers is not None and not use_base_llm_http_handler: optional_params["extra_headers"] = extra_headers if litellm.enable_preview_features and metadata is not None: # [PREVIEW] allow metadata to be passed to OPENAI @@ -2609,8 +2611,6 @@ def _complete_custom_openai( optional_params[k] = v ## COMPLETION CALL - use_base_llm_http_handler: Final = get_secret_bool("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER") - try: if use_base_llm_http_handler: response = base_llm_http_handler.completion( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 685e63dcc56..9f91cf82f41 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -11216,12 +11216,15 @@ "babbage-002": { "deprecation_date": "2026-09-28", "input_cost_per_token": 4e-07, + "input_cost_per_token_batches": 2e-07, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 4e-07 + "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "bedrock/*/1-month-commitment/cohere.command-light-text-v14": { "input_cost_per_second": 0.001902, @@ -13286,7 +13289,9 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 0.0001, + "source": "https://developers.openai.com/api/docs/pricing" }, "claude-haiku-4-5-20251001": { "deprecation_date": "2026-10-15", @@ -13334,7 +13339,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-3-7-sonnet-20250219": { "cache_creation_input_token_cost": 3.75e-06, @@ -13493,7 +13499,8 @@ "supports_native_structured_output": true, "supports_tool_choice": true, "supports_vision": true, - "prompt_cache_min_tokens": 1024 + "prompt_cache_min_tokens": 1024, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-5-20250929": { "deprecation_date": "2026-09-29", @@ -13567,7 +13574,7 @@ }, "supports_output_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-6": { "deprecation_date": "2027-02-17", @@ -13603,7 +13610,8 @@ "prompt_cache_min_tokens": 1024, "provider_specific_entry": { "us": 1.1 - } + }, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, @@ -13780,7 +13788,8 @@ "supports_tool_choice": true, "supports_vision": true, "supports_output_config": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6": { "deprecation_date": "2027-02-05", @@ -13817,7 +13826,8 @@ "supports_output_config": true, "supports_max_reasoning_effort": true, "supports_speed": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6-20260205": { "deprecation_date": "2027-02-05", @@ -13892,7 +13902,8 @@ }, "supports_output_config": true, "supports_speed": true, - "prompt_cache_min_tokens": 2048 + "prompt_cache_min_tokens": 2048, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-7-20260416": { "deprecation_date": "2027-04-16", @@ -13970,7 +13981,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-fable-5-1": { "deprecation_date": "2027-09-01", @@ -14011,7 +14022,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://platform.claude.com/docs/en/models/fable-5-1/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5": { "deprecation_date": "2027-07-24", @@ -14052,7 +14063,7 @@ "supports_output_config": true, "supports_speed": true, "prompt_cache_min_tokens": 512, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-8": { "deprecation_date": "2027-05-28", @@ -14092,7 +14103,8 @@ }, "supports_output_config": true, "supports_speed": true, - "prompt_cache_min_tokens": 1024 + "prompt_cache_min_tokens": 1024, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-20250514": { "deprecation_date": "2026-06-15", @@ -19361,12 +19373,15 @@ "davinci-002": { "deprecation_date": "2026-09-28", "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 2e-06 + "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "deepgram/base": { "input_cost_per_second": 0.00020833, @@ -22413,15 +22428,18 @@ "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.45e-07, - "input_cost_per_token": 1.74e-06, + "cache_read_input_token_cost": 6e-07, + "cache_read_input_token_cost_priority": 6e-07, + "input_cost_per_token": 1.2e-06, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 3.48e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22430,14 +22448,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 4.4e-08, + "cache_read_input_token_cost_priority": 5.5e-08, "input_cost_per_token": 1.32e-06, + "input_cost_per_token_priority": 1.65e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 3.96e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 4.95e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22532,14 +22553,17 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p2": { "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost_priority": 1.75e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22548,14 +22572,17 @@ }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { "cache_read_input_token_cost": 1.5e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.8e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 7.2e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22634,14 +22661,17 @@ }, "fireworks_ai/accounts/fireworks/models/kimi-k2p6": { "cache_read_input_token_cost": 1.6e-07, + "cache_read_input_token_cost_priority": 2.2e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22650,14 +22680,17 @@ }, "fireworks_ai/accounts/fireworks/models/kimi-k2p7-code": { "cache_read_input_token_cost": 1.9e-07, + "cache_read_input_token_cost_priority": 2.85e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.425e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22783,14 +22816,17 @@ }, "fireworks_ai/accounts/fireworks/models/minimax-m2p7": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 6e-07, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 196608, "max_output_tokens": 196608, "max_tokens": 196608, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22799,14 +22835,17 @@ }, "fireworks_ai/accounts/fireworks/models/minimax-m3": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 9e-08, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 512000, "max_output_tokens": 512000, "max_tokens": 512000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.8e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22882,15 +22921,18 @@ "supports_vision": false }, "fireworks_ai/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.45e-07, - "input_cost_per_token": 1.74e-06, + "cache_read_input_token_cost": 6e-07, + "cache_read_input_token_cost_priority": 6e-07, + "input_cost_per_token": 1.2e-06, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 3.48e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22946,14 +22988,17 @@ }, "fireworks_ai/glm-5p2": { "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost_priority": 1.75e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22962,14 +23007,17 @@ }, "fireworks_ai/gpt-oss-120b": { "cache_read_input_token_cost": 1.5e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.8e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 7.2e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23008,14 +23056,17 @@ }, "fireworks_ai/kimi-k2p6": { "cache_read_input_token_cost": 1.6e-07, + "cache_read_input_token_cost_priority": 2.2e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23040,14 +23091,17 @@ }, "fireworks_ai/kimi-k2p7-code": { "cache_read_input_token_cost": 1.9e-07, + "cache_read_input_token_cost_priority": 2.85e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.425e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23086,14 +23140,17 @@ }, "fireworks_ai/minimax-m2p7": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 6e-07, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 196608, "max_output_tokens": 196608, "max_tokens": 196608, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23102,14 +23159,17 @@ }, "fireworks_ai/minimax-m3": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 9e-08, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 512000, "max_output_tokens": 512000, "max_tokens": 512000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.8e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23125,7 +23185,7 @@ "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23374,26 +23434,28 @@ "ft:babbage-002": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1.6e-06, - "input_cost_per_token_batches": 2e-07, + "input_cost_per_token_batches": 8e-07, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.6e-06, - "output_cost_per_token_batches": 2e-07 + "output_cost_per_token_batches": 9e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "ft:davinci-002": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1.2e-05, - "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_batches": 6e-06, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.2e-05, - "output_cost_per_token_batches": 1e-06 + "output_cost_per_token_batches": 6e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "ft:gpt-3.5-turbo": { "deprecation_date": "2026-10-23", @@ -23406,6 +23468,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_batches": 3e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_system_messages": true, "supports_tool_choice": true }, @@ -23462,14 +23525,15 @@ "ft:gpt-4o-2024-08-06": { "cache_read_input_token_cost": 1.875e-06, "input_cost_per_token": 3.75e-06, - "input_cost_per_token_batches": 1.875e-06, + "input_cost_per_token_batches": 2.225e-06, "litellm_provider": "openai", "max_input_tokens": 128000, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-05, - "output_cost_per_token_batches": 7.5e-06, + "output_cost_per_token_batches": 1.25e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -23507,6 +23571,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "output_cost_per_token_batches": 6e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -23526,6 +23591,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23544,6 +23610,7 @@ "mode": "chat", "output_cost_per_token": 3.2e-06, "output_cost_per_token_batches": 1.6e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23563,6 +23630,7 @@ "mode": "chat", "output_cost_per_token": 8e-07, "output_cost_per_token_batches": 4e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23582,6 +23650,7 @@ "mode": "chat", "output_cost_per_token": 1.6e-05, "output_cost_per_token_batches": 8e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -23592,15 +23661,18 @@ "gemini-2.0-flash": { "cache_read_input_token_cost": 2.5e-08, "deprecation_date": "2026-06-01", - "input_cost_per_audio_token": 7e-07, - "input_cost_per_token": 1e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_character": 3.75e-08, + "input_cost_per_token": 1.5e-07, + "input_cost_per_token_batches": 7.5e-08, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 4e-07, - "source": "https://ai.google.dev/pricing#2_0flash", + "output_cost_per_token": 6e-07, + "output_cost_per_token_batches": 3e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image", @@ -23669,13 +23741,16 @@ "cache_read_input_token_cost": 1.875e-08, "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 7.5e-08, + "input_cost_per_character": 1.875e-08, "input_cost_per_token": 7.5e-08, + "input_cost_per_token_batches": 3.75e-08, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", + "output_cost_per_token_batches": 1.5e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image", @@ -23736,8 +23811,11 @@ } }, "gemini-2.5-flash": { + "cache_read_input_audio_token_cost": 1e-07, "deprecation_date": "2026-10-20", "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 3e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, "litellm_provider": "vertex_ai-language-models", @@ -23747,7 +23825,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23780,6 +23858,12 @@ "search_context_size_high": 0.035 }, "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, "supports_image_size": false }, "gemini-2.5-flash-image": { @@ -23787,6 +23871,9 @@ "cache_read_input_token_cost": 3e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 32768, "max_output_tokens": 32768, @@ -23796,8 +23883,10 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, "rpm": 100000, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23828,10 +23917,19 @@ "supports_image_size": false }, "gemini-3-pro-image": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 1e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 3.6e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -23840,8 +23938,12 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3-pro-image", + "output_cost_per_token_flex": 6e-06, + "output_cost_per_token_priority": 2.16e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23909,9 +24011,13 @@ "web_search_billing_unit": "per_query" }, "gemini-3.1-flash-image": { + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 2.5e-08, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -23920,7 +24026,9 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23987,9 +24095,11 @@ }, "gemini-3.1-flash-lite-image": { "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_flex": 1.25e-08, "input_cost_per_image": 0.00028, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 4096, @@ -23999,6 +24109,7 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_token": 1.5e-06, "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -24073,6 +24184,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "gemini-3.1-flash-lite": { + "cache_read_input_audio_token_cost": 5e-08, "deprecation_date": "2027-05-07", "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, @@ -24092,7 +24204,7 @@ "output_cost_per_token_batches": 7.5e-07, "output_cost_per_token_flex": 7.5e-07, "output_cost_per_token_priority": 2.7e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24134,7 +24246,7 @@ "deprecation_date": "2027-07-21", "cache_read_input_token_cost": 3e-08, "cache_read_input_token_cost_flex": 1.5e-08, - "cache_read_input_token_cost_priority": 5e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, "input_cost_per_token_flex": 1.5e-07, @@ -24149,7 +24261,7 @@ "output_cost_per_token_batches": 1.25e-06, "output_cost_per_token_flex": 1.25e-06, "output_cost_per_token_priority": 4.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24188,6 +24300,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "deep-research-pro-preview-12-2025": { + "cache_read_input_token_cost": 2e-07, "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -24200,7 +24313,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24222,8 +24335,11 @@ "supports_web_search": true }, "gemini-2.5-flash-lite": { + "cache_read_input_audio_token_cost": 3e-08, "deprecation_date": "2026-10-20", "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_flex": 1e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_audio_token": 3e-07, "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-language-models", @@ -24233,7 +24349,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24266,6 +24382,12 @@ "search_context_size_high": 0.035 }, "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_batches": 5e-08, + "input_cost_per_token_flex": 5e-08, + "input_cost_per_token_priority": 1.8e-07, + "output_cost_per_token_batches": 2e-07, + "output_cost_per_token_flex": 2e-07, + "output_cost_per_token_priority": 7.2e-07, "supports_image_size": false }, "gemini-2.5-flash-lite-preview-09-2025": { @@ -24417,7 +24539,7 @@ "mode": "realtime", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_token": 2e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/vertex_ai/live" ], @@ -24448,7 +24570,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "gemini_native_audio": true + "gemini_native_audio": true, + "input_cost_per_image_token": 3e-06 }, "gemini/gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -24548,6 +24671,9 @@ "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 4.5e-07, + "cache_read_input_token_cost_flex": 1.25e-07, + "cache_read_input_token_cost_priority": 2.25e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, "litellm_provider": "vertex_ai-language-models", @@ -24557,7 +24683,7 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions" @@ -24587,7 +24713,15 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "google_maps_grounding_cost_per_query": 0.025 + "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_above_200k_tokens_priority": 4.5e-06, + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "input_cost_per_token_priority": 2.25e-06, + "output_cost_per_token_above_200k_tokens_priority": 2.7e-05, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "output_cost_per_token_priority": 1.8e-05 }, "gemini-3-pro-preview": { "deprecation_date": "2026-03-26", @@ -24662,7 +24796,7 @@ "output_cost_per_token_above_200k_tokens": 1.8e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_image": 0.00012, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24696,13 +24830,16 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_low": 0.014, "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_flex": 1e-06, + "output_cost_per_token_flex": 6e-06 }, "gemini-3.1-pro-preview-customtools": { "prompt_cache_min_tokens": 4096, @@ -24813,7 +24950,9 @@ "web_search_billing_unit": "per_query" }, "vertex_ai/gemini-3-flash-preview": { + "cache_read_input_audio_token_cost": 1e-07, "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, "input_cost_per_token": 5e-07, "input_cost_per_audio_token": 1e-06, "litellm_provider": "vertex_ai", @@ -24822,7 +24961,7 @@ "max_tokens": 65535, "mode": "chat", "output_cost_per_token": 3e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24859,7 +24998,11 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06 }, "vertex_ai/gemini-3.5-flash": { "prompt_cache_min_tokens": 4096, @@ -24875,7 +25018,7 @@ "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24938,7 +25081,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24995,7 +25138,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25052,7 +25195,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25109,7 +25252,7 @@ "output_cost_per_token_above_200k_tokens": 1.8e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_image": 0.00012, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25143,13 +25286,16 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_low": 0.014, "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_flex": 1e-06, + "output_cost_per_token_flex": 6e-06 }, "vertex_ai/gemini-3.1-pro-preview-customtools": { "prompt_cache_min_tokens": 4096, @@ -25214,13 +25360,15 @@ "cache_read_input_token_cost": 1.25e-07, "input_cost_per_audio_token": 7e-07, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 65535, "max_tokens": 65535, "mode": "chat", + "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text" ], @@ -25427,7 +25575,7 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/computer-use", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image" @@ -25453,10 +25601,14 @@ "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models" }, "gemini-embedding-2-preview": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25467,25 +25619,33 @@ "uses_embed_content": true }, "gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, "vertex_ai/gemini-embedding-2-preview": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25497,17 +25657,21 @@ "uses_embed_content": true }, "vertex_ai/gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, @@ -25539,10 +25703,14 @@ }, "gemini/gemini-embedding-2-preview": { "deprecation_date": "2026-08-10", - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "gemini", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25555,10 +25723,14 @@ "tpm": 10000000 }, "gemini/gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "gemini", "max_input_tokens": 8192, "max_tokens": 8192, @@ -27141,7 +27313,9 @@ "google_maps_grounding_cost_per_query": 0.014 }, "gemini-3-flash-preview": { + "cache_read_input_audio_token_cost": 1e-07, "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 5e-07, "litellm_provider": "vertex_ai-language-models", @@ -27151,7 +27325,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 3e-06, "output_cost_per_token": 3e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27189,7 +27363,11 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06 }, "gemini-omni-flash-preview": { "input_cost_per_audio_token": 1.5e-06, @@ -27202,7 +27380,7 @@ "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, "output_cost_per_video_token": 1.75e-05, - "source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/gemini/omni-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions" ], @@ -27235,7 +27413,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27298,7 +27476,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27355,7 +27533,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27412,7 +27590,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -28870,6 +29048,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_system_messages": true, @@ -28878,12 +29057,15 @@ "gpt-3.5-turbo-0125": { "deprecation_date": "2026-10-23", "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -28893,12 +29075,15 @@ "gpt-3.5-turbo-1106": { "deprecation_date": "2026-09-28", "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -28926,7 +29111,8 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 2e-06 + "output_cost_per_token": 2e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-3.5-turbo-instruct-0914": { "input_cost_per_token": 1.5e-06, @@ -28981,12 +29167,15 @@ "gpt-4-0613": { "deprecation_date": "2026-10-23", "input_cost_per_token": 3e-05, + "input_cost_per_token_batches": 1.5e-05, "litellm_provider": "openai", "max_input_tokens": 8192, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_system_messages": true, @@ -29027,12 +29216,15 @@ "gpt-4-turbo-2024-04-09": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "openai", "max_input_tokens": 128000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 3e-05, + "output_cost_per_token_batches": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29076,6 +29268,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29118,6 +29311,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29160,6 +29354,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29202,6 +29397,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29240,6 +29436,7 @@ "output_cost_per_token": 4e-07, "output_cost_per_token_batches": 2e-07, "output_cost_per_token_priority": 8e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29277,6 +29474,7 @@ "output_cost_per_token": 4e-07, "output_cost_per_token_priority": 8e-07, "output_cost_per_token_batches": 2e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29313,6 +29511,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_batches": 5e-06, "output_cost_per_token_priority": 1.7e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29335,6 +29534,7 @@ "output_cost_per_token": 1.5e-05, "output_cost_per_token_batches": 7.5e-06, "output_cost_per_token_priority": 2.625e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29357,6 +29557,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_priority": 1.7e-05, "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29380,6 +29581,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_priority": 1.7e-05, "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29454,6 +29656,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29490,6 +29693,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions" ], @@ -29524,6 +29728,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29561,6 +29766,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29598,6 +29804,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29659,7 +29866,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": false, - "deprecation_date": "2027-01-20" + "deprecation_date": "2027-01-20", + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini": { "cache_read_input_token_cost": 7.5e-08, @@ -29687,7 +29895,8 @@ "search_context_size_high": 0.025, "search_context_size_low": 0.025, "search_context_size_medium": 0.025 - } + }, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini-2024-07-18": { "cache_read_input_token_cost": 7.5e-08, @@ -29708,6 +29917,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29856,15 +30066,18 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 5e-05, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini-tts": { - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -29996,7 +30209,9 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 0.0001, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-1.5": { "cache_read_input_token_cost": 1.25e-06, @@ -30006,7 +30221,10 @@ "mode": "image_generation", "output_cost_per_token": 1e-05, "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations" ], @@ -30021,7 +30239,10 @@ "mode": "image_generation", "output_cost_per_token": 1e-05, "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations" ], @@ -30034,7 +30255,9 @@ "litellm_provider": "openai", "mode": "image_generation", "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -30481,6 +30704,7 @@ "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -30489,6 +30713,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, "search_context_cost_per_query": { @@ -30496,6 +30721,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -30525,6 +30751,7 @@ }, "gpt-5.1": { "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_priority": 2.5e-06, @@ -30565,11 +30792,17 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, "supports_minimal_reasoning_effort": false }, "gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_priority": 2.5e-06, @@ -30610,6 +30843,11 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, "supports_minimal_reasoning_effort": false }, @@ -30661,6 +30899,7 @@ }, "gpt-5.2": { "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_flex": 8.75e-08, "cache_read_input_token_cost_priority": 3.5e-07, "input_cost_per_token": 1.75e-06, "input_cost_per_token_priority": 3.5e-06, @@ -30702,11 +30941,17 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 8.75e-07, + "input_cost_per_token_flex": 8.75e-07, + "output_cost_per_token_batches": 7e-06, + "output_cost_per_token_flex": 7e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, "gpt-5.2-2025-12-11": { "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_flex": 8.75e-08, "cache_read_input_token_cost_priority": 3.5e-07, "input_cost_per_token": 1.75e-06, "input_cost_per_token_priority": 3.5e-06, @@ -30748,6 +30993,11 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 8.75e-07, + "input_cost_per_token_flex": 8.75e-07, + "output_cost_per_token_batches": 7e-06, + "output_cost_per_token_flex": 7e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -30841,17 +31091,20 @@ }, "gpt-5.2-pro": { "input_cost_per_token": 2.1e-05, + "input_cost_per_token_batches": 1.05e-05, "litellm_provider": "openai", "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "output_cost_per_token": 0.000168, + "output_cost_per_token_batches": 8.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -30880,17 +31133,20 @@ }, "gpt-5.2-pro-2025-12-11": { "input_cost_per_token": 2.1e-05, + "input_cost_per_token_batches": 1.05e-05, "litellm_provider": "openai", "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "output_cost_per_token": 0.000168, + "output_cost_per_token_batches": 8.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -30956,6 +31212,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31092,6 +31349,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31160,6 +31418,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31227,6 +31486,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31290,7 +31550,7 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "source": "https://developers.openai.com/api/docs/models/gpt-5.6-cyber", + "source": "https://developers.openai.com/api/docs/pricing", "supports_computer_use": true, "supports_parallel_function_calling": true }, @@ -31460,7 +31720,7 @@ "reasoning_effort_levels": [ "medium" ], - "source": "https://developers.openai.com/api/docs/models/chat-latest", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -31539,7 +31799,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 5e-06, "output_cost_per_token_above_272k_tokens_flex": 2.25e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.5-2026-04-23": { "cache_read_input_token_cost": 5e-07, @@ -31596,7 +31857,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 5e-06, "output_cost_per_token_above_272k_tokens_flex": 2.25e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.5-pro": { "input_cost_per_token": 3e-05, @@ -31619,6 +31881,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -31667,6 +31930,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -31744,7 +32008,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 2.5e-06, "output_cost_per_token_above_272k_tokens_flex": 1.125e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-2026-03-05": { "cache_read_input_token_cost": 2.5e-07, @@ -31796,7 +32061,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 2.5e-06, "output_cost_per_token_above_272k_tokens_flex": 1.125e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-pro": { "input_cost_per_token": 3e-05, @@ -31845,7 +32111,8 @@ "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 3e-05, - "output_cost_per_token_above_272k_tokens_flex": 0.000135 + "output_cost_per_token_above_272k_tokens_flex": 0.000135, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-pro-2026-03-05": { "input_cost_per_token": 3e-05, @@ -31894,7 +32161,8 @@ "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 3e-05, - "output_cost_per_token_above_272k_tokens_flex": 0.000135 + "output_cost_per_token_above_272k_tokens_flex": 0.000135, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, @@ -31945,6 +32213,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -31997,6 +32266,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32046,6 +32316,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32095,6 +32366,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32113,6 +32385,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -32155,6 +32428,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -32187,6 +32461,7 @@ "cache_read_input_token_cost_priority": 2.5e-07, "deprecation_date": "2026-12-11", "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -32195,6 +32470,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, "search_context_cost_per_query": { @@ -32202,6 +32478,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32526,6 +32803,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses" ], @@ -32556,6 +32834,7 @@ "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, "input_cost_per_token_flex": 1.25e-07, "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "openai", @@ -32564,6 +32843,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, "output_cost_per_token_flex": 1e-06, "output_cost_per_token_priority": 3.6e-06, "search_context_cost_per_query": { @@ -32571,6 +32851,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32604,6 +32885,7 @@ "cache_read_input_token_cost_priority": 4.5e-08, "deprecation_date": "2026-12-11", "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, "input_cost_per_token_flex": 1.25e-07, "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "openai", @@ -32612,6 +32894,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, "output_cost_per_token_flex": 1e-06, "output_cost_per_token_priority": 3.6e-06, "search_context_cost_per_query": { @@ -32619,6 +32902,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32650,6 +32934,7 @@ "cache_read_input_token_cost": 5e-09, "cache_read_input_token_cost_flex": 2.5e-09, "input_cost_per_token": 5e-08, + "input_cost_per_token_batches": 2.5e-08, "input_cost_per_token_flex": 2.5e-08, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -32658,12 +32943,14 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, "output_cost_per_token_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32696,6 +32983,7 @@ "cache_read_input_token_cost_flex": 2.5e-09, "deprecation_date": "2026-12-11", "input_cost_per_token": 5e-08, + "input_cost_per_token_batches": 2.5e-08, "input_cost_per_token_priority": 2.5e-06, "input_cost_per_token_flex": 2.5e-08, "litellm_provider": "openai", @@ -32704,12 +32992,14 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, "output_cost_per_token_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32742,9 +33032,11 @@ "deprecation_date": "2026-10-23", "input_cost_per_image_token": 1e-05, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "openai", "mode": "image_generation", "output_cost_per_image_token": 4e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32755,9 +33047,11 @@ "deprecation_date": "2026-12-01", "input_cost_per_image_token": 2.5e-06, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "openai", "mode": "image_generation", "output_cost_per_image_token": 8e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32778,6 +33072,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32811,6 +33106,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32844,6 +33140,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 2.4e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32879,6 +33176,7 @@ "output_cost_per_token": 2.4e-05, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32914,6 +33212,7 @@ "output_cost_per_token": 2.4e-06, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32939,6 +33238,7 @@ "cache_read_input_token_cost": 6e-08, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1e-05, + "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, "litellm_provider": "openai", "max_input_tokens": 32000, @@ -32947,6 +33247,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32981,6 +33282,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -37700,12 +38002,15 @@ "cache_read_input_token_cost": 7.5e-06, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.5e-05, + "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -37720,12 +38025,15 @@ "cache_read_input_token_cost": 7.5e-06, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.5e-05, + "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -37747,6 +38055,7 @@ "mode": "responses", "output_cost_per_token": 0.0006, "output_cost_per_token_batches": 0.0003, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -37780,6 +38089,7 @@ "mode": "responses", "output_cost_per_token": 0.0006, "output_cost_per_token_batches": 0.0003, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -37807,6 +38117,7 @@ "cache_read_input_token_cost_flex": 2.5e-07, "cache_read_input_token_cost_priority": 8.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 3.5e-06, "litellm_provider": "openai", @@ -37815,6 +38126,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8e-06, + "output_cost_per_token_batches": 4e-06, "output_cost_per_token_flex": 4e-06, "output_cost_per_token_priority": 1.4e-05, "search_context_cost_per_query": { @@ -37822,6 +38134,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/chat/completions", @@ -37851,6 +38164,7 @@ "cache_read_input_token_cost_priority": 8.75e-07, "deprecation_date": "2026-12-11", "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 3.5e-06, "litellm_provider": "openai", @@ -37859,6 +38173,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8e-06, + "output_cost_per_token_batches": 4e-06, "output_cost_per_token_flex": 4e-06, "output_cost_per_token_priority": 1.4e-05, "search_context_cost_per_query": { @@ -37866,6 +38181,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/chat/completions", @@ -37975,12 +38291,15 @@ "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -37993,12 +38312,15 @@ "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -38022,6 +38344,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -38059,6 +38382,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -38082,10 +38406,11 @@ }, "o4-mini": { "cache_read_input_token_cost": 2.75e-07, - "cache_read_input_token_cost_flex": 1.375e-07, + "cache_read_input_token_cost_flex": 1.38e-07, "cache_read_input_token_cost_priority": 5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "input_cost_per_token_flex": 5.5e-07, "input_cost_per_token_priority": 2e-06, "litellm_provider": "openai", @@ -38094,6 +38419,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, "output_cost_per_token_flex": 2.2e-06, "output_cost_per_token_priority": 8e-06, "search_context_cost_per_query": { @@ -38101,6 +38427,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -38113,10 +38440,11 @@ }, "o4-mini-2025-04-16": { "cache_read_input_token_cost": 2.75e-07, - "cache_read_input_token_cost_flex": 1.375e-07, + "cache_read_input_token_cost_flex": 1.38e-07, "cache_read_input_token_cost_priority": 5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "input_cost_per_token_flex": 5.5e-07, "input_cost_per_token_priority": 2e-06, "litellm_provider": "openai", @@ -38125,6 +38453,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, "output_cost_per_token_flex": 2.2e-06, "output_cost_per_token_priority": 8e-06, "search_context_cost_per_query": { @@ -38132,6 +38461,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -43294,7 +43624,8 @@ "mode": "embedding", "output_cost_per_token": 0.0, "output_cost_per_token_batches": 0.0, - "output_vector_size": 3072 + "output_vector_size": 3072, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-3-small": { "input_cost_per_token": 2e-08, @@ -43305,7 +43636,8 @@ "mode": "embedding", "output_cost_per_token": 0.0, "output_cost_per_token_batches": 0.0, - "output_vector_size": 1536 + "output_vector_size": 1536, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-ada-002": { "input_cost_per_token": 1e-07, @@ -43314,7 +43646,8 @@ "max_tokens": 8191, "mode": "embedding", "output_cost_per_token": 0.0, - "output_vector_size": 1536 + "output_vector_size": 1536, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-ada-002-v2": { "input_cost_per_token": 1e-07, @@ -43485,7 +43818,7 @@ "input_cost_per_token": 1.2e-06, "output_cost_per_token": 1.2e-06, "max_input_tokens": 131072, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-7B-Instruct-Turbo": { "litellm_provider": "together_ai", @@ -43497,7 +43830,7 @@ "input_cost_per_token": 3e-07, "output_cost_per_token": 3e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { "deprecation_date": "2026-07-10", @@ -43544,7 +43877,7 @@ "max_input_tokens": 256000, "mode": "chat", "output_cost_per_token": 2e-06, - "source": "https://www.together.ai/models/qwen3-coder-480b-a35b-instruct", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43606,7 +43939,7 @@ }, "mode": "chat", "output_cost_per_token": 1.7e-06, - "source": "https://www.together.ai/models/deepseek-v3-1", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43630,7 +43963,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.04e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43664,6 +43997,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 5.9e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43686,6 +44020,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43697,6 +44032,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 1.8e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43713,7 +44049,7 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 2e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { "deprecation_date": "2026-04-02", @@ -43725,7 +44061,7 @@ "input_cost_per_token": 1e-07, "output_cost_per_token": 3e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { "deprecation_date": "2026-04-16", @@ -43733,6 +44069,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 6e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43759,7 +44096,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://www.together.ai/models/gpt-oss-120b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43773,7 +44110,7 @@ "max_input_tokens": 131072, "mode": "chat", "output_cost_per_token": 2e-07, - "source": "https://www.together.ai/models/gpt-oss-20b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43793,7 +44130,7 @@ "max_input_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.1e-06, - "source": "https://www.together.ai/models/glm-4-5-air", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43809,7 +44146,7 @@ }, "mode": "chat", "output_cost_per_token": 2.2e-06, - "source": "https://www.together.ai/models/glm-4-6", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43826,7 +44163,7 @@ }, "mode": "chat", "output_cost_per_token": 2e-06, - "source": "https://www.together.ai/models/glm-4-7", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43874,7 +44211,7 @@ }, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://www.together.ai/models/qwen3-next-80b-a3b-instruct", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43890,7 +44227,7 @@ }, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://www.together.ai/models/qwen3-next-80b-a3b-thinking", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43904,7 +44241,7 @@ "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 3.6e-06, - "source": "https://www.together.ai/models/qwen3-5-397b-a17b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -43919,7 +44256,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -43944,7 +44281,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 2.5e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43959,7 +44296,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 3e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_reasoning": true }, "together_ai/Qwen/Qwen3.7-Max": { @@ -43970,7 +44307,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 7.5e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/Qwen/Qwen3.7-Plus": { @@ -43980,7 +44317,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 1.28e-06, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3.8-2.4T-A95B": { "cache_read_input_token_cost": 2.5e-07, @@ -43990,7 +44327,7 @@ "max_tokens": 1010000, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/arize-ai/qwen-2-1.5b-instruct": { @@ -44000,7 +44337,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 1e-07, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { "cache_read_input_token_cost": 3e-08, @@ -44010,7 +44347,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.8e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44042,7 +44379,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 3.96e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44067,7 +44404,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 9.7e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -44103,7 +44440,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/moonshotai/Kimi-K2.7-Code": { @@ -44115,7 +44452,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44136,7 +44473,7 @@ "high", "max" ], - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44154,7 +44491,7 @@ "max_tokens": 512288, "mode": "chat", "output_cost_per_token": 3.6e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44180,7 +44517,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 4.05e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44196,7 +44533,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/zai-org/GLM-5.2": { @@ -44208,7 +44545,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44225,7 +44562,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44242,7 +44579,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44255,6 +44592,7 @@ "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", "mode": "audio_speech", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ] @@ -44263,6 +44601,7 @@ "input_cost_per_character": 3e-05, "litellm_provider": "openai", "mode": "audio_speech", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ] @@ -47605,6 +47944,9 @@ "cache_read_input_token_cost": 3e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 32768, "max_output_tokens": 32768, @@ -47614,8 +47956,10 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, "rpm": 100000, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-generation#edit-an-image", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47647,10 +47991,19 @@ "supports_image_size": false }, "vertex_ai/gemini-3-pro-image": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 1e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 3.6e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -47659,9 +48012,13 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, + "output_cost_per_token_flex": 6e-06, + "output_cost_per_token_priority": 2.16e-05, "supports_reasoning": false, - "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -47680,9 +48037,13 @@ "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, "vertex_ai/gemini-3.1-flash-image": { + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 2.5e-08, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -47691,8 +48052,10 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_reasoning": false, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/gemini-3.1-flash-image-preview": { "input_cost_per_image": 0.00056, @@ -47710,9 +48073,11 @@ }, "vertex_ai/gemini-3.1-flash-lite-image": { "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_flex": 1.25e-08, "input_cost_per_image": 0.00028, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 4096, @@ -47722,6 +48087,7 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_token": 1.5e-06, "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -47796,6 +48162,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "vertex_ai/gemini-3.1-flash-lite": { + "cache_read_input_audio_token_cost": 5e-08, "deprecation_date": "2027-05-07", "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, @@ -47816,7 +48183,7 @@ "output_cost_per_token_flex": 7.5e-07, "output_cost_per_token_priority": 2.7e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47858,7 +48225,7 @@ "deprecation_date": "2027-07-21", "cache_read_input_token_cost": 3e-08, "cache_read_input_token_cost_flex": 1.5e-08, - "cache_read_input_token_cost_priority": 5e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, "input_cost_per_token_flex": 1.5e-07, @@ -47874,7 +48241,7 @@ "output_cost_per_token_flex": 1.25e-06, "output_cost_per_token_priority": 4.5e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47913,6 +48280,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "vertex_ai/deep-research-pro-preview-12-2025": { + "cache_read_input_token_cost": 2e-07, "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -47925,7 +48293,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", @@ -49588,7 +49956,8 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "source": "https://developers.openai.com/api/docs/pricing" }, "xai/grok-3": { "cache_read_input_token_cost": 2e-07, @@ -52685,10 +53054,11 @@ "max_tokens": 40960, "max_input_tokens": 40960, "max_output_tokens": 40960, - "input_cost_per_token": 0.0, + "input_cost_per_token": 2e-07, "output_cost_per_token": 0.0, "litellm_provider": "fireworks_ai", - "mode": "rerank" + "mode": "rerank", + "source": "https://api.fireworks.ai/v1/serverless/models" }, "fireworks_ai/accounts/fireworks/models/qwen3-vl-235b-a22b-instruct": { "max_tokens": 262144, @@ -52753,7 +53123,7 @@ "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -54652,12 +55022,13 @@ }, "gpt-4o-mini-tts-2025-03-20": { "deprecation_date": "2026-07-23", - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -54670,12 +55041,13 @@ ] }, "gpt-4o-mini-tts-2025-12-15": { - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -54690,24 +55062,28 @@ "gpt-4o-mini-transcribe-2025-03-20": { "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1.25e-06, + "input_cost_per_second": 5e-05, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions" ] }, "gpt-4o-mini-transcribe-2025-12-15": { "input_cost_per_audio_token": 1.25e-06, + "input_cost_per_second": 5e-05, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -54726,6 +55102,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -54753,6 +55130,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -54780,6 +55158,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -54831,13 +55210,14 @@ "supports_parallel_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "deprecation_date": "2027-01-20" + "deprecation_date": "2027-01-20", + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-realtime-whisper": { - "input_cost_per_second": 0.0002833333333333333, + "input_cost_per_second": 0.000283333333333, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-realtime-whisper", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime", "/v1/realtime/transcription_sessions" @@ -54855,7 +55235,7 @@ "litellm_provider": "openai", "mode": "video_generation", "output_cost_per_video_per_second": 0.1, - "source": "https://platform.openai.com/docs/api-reference/videos", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "image" @@ -54869,7 +55249,7 @@ "litellm_provider": "openai", "mode": "video_generation", "output_cost_per_video_per_second": 0.3, - "source": "https://platform.openai.com/docs/api-reference/videos", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "image" @@ -54894,11 +55274,15 @@ "chatgpt-image-latest": { "cache_read_input_token_cost": 1.25e-06, "deprecation_date": "2026-12-01", - "input_cost_per_image_token": 1e-05, + "input_cost_per_image_token": 8e-06, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "openai", "mode": "image_generation", - "output_cost_per_image_token": 4e-05, + "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -57373,7 +57757,7 @@ "input_cost_per_second": 7.5e-05, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-transcribe", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions", "/v1/realtime/transcription_sessions" @@ -57388,10 +57772,10 @@ "supports_audio_input": true }, "gpt-live-transcribe": { - "input_cost_per_second": 0.0002833333333333333, + "input_cost_per_second": 0.000283333333333, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-live-transcribe", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime", "/v1/realtime/transcription_sessions" @@ -57406,10 +57790,10 @@ "supports_audio_input": true }, "gpt-live-1": { - "input_cost_per_second": 0.0008333333333333334, + "input_cost_per_second": 0.000833333333333, "litellm_provider": "openai", "mode": "realtime", - "source": "https://developers.openai.com/api/docs/models/gpt-live-1", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "audio" @@ -57423,13 +57807,13 @@ "supports_function_calling": true }, "gpt-realtime-translate": { - "input_cost_per_second": 0.0005666666666666667, + "input_cost_per_second": 0.000566666666667, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "max_tokens": 2000, "mode": "realtime", - "source": "https://developers.openai.com/api/docs/models/gpt-realtime-translate", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "audio" ], @@ -57457,7 +57841,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, - "source": "https://platform.claude.com/docs/en/about-claude/models/overview", + "source": "https://platform.claude.com/docs/en/about-claude/pricing", "supports_adaptive_thinking": true, "thinking_always_on": true, "supports_mid_conversation_system": true, @@ -57518,7 +57902,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://platform.claude.com/docs/en/models/mythos-5-1/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-mythos-preview": { "cache_creation_input_token_cost": 1.25e-05, @@ -57874,6 +58258,7 @@ }, "vertex_ai/gemini-3.5-live-translate-preview": { "input_cost_per_audio_token": 3.5e-06, + "input_cost_per_second": 8.83333333333e-05, "input_cost_per_token": 3.5e-06, "litellm_provider": "vertex_ai", "mode": "realtime", @@ -57976,14 +58361,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -57992,14 +58380,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4p1-flash": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://fireworks.ai/models/deepseek-ai/deepseek-v4p1-flash", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -58015,26 +58406,29 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, "supports_vision": true }, "fireworks_ai/accounts/fireworks/models/kimi-k3": { "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_priority": 3.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_priority": 1.875e-05, "reasoning_effort_levels": [ "low", "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58043,14 +58437,17 @@ }, "fireworks_ai/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58059,14 +58456,17 @@ }, "fireworks_ai/deepseek-v4p1-flash": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://fireworks.ai/models/deepseek-ai/deepseek-v4p1-flash", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -58082,7 +58482,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, "supports_vision": true @@ -58121,19 +58521,22 @@ }, "fireworks_ai/kimi-k3": { "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_priority": 3.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_priority": 1.875e-05, "reasoning_effort_levels": [ "low", "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58184,12 +58587,15 @@ }, "fireworks_ai/qwen3p8-max": { "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 9e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58205,7 +58611,7 @@ "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58237,7 +58643,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 2.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58253,7 +58659,7 @@ "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58285,7 +58691,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 2.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58294,12 +58700,15 @@ }, "fireworks_ai/accounts/fireworks/models/qwen3p8-max": { "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 9e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58315,7 +58724,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58352,7 +58761,7 @@ "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -61050,14 +61459,17 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p3": { "cache_read_input_token_cost": 2.6e-07, + "cache_read_input_token_cost_priority": 3.25e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -61066,13 +61478,16 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p3-flash": { "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_priority": 3.75e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.875e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 5e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -61100,7 +61515,7 @@ "max_output_tokens": 40960, "max_tokens": 40960, "mode": "embedding", - "source": "https://docs.fireworks.ai/serverless/pricing" + "source": "https://api.fireworks.ai/v1/serverless/models" }, "zai/glm-5.2": { "cache_creation_input_token_cost": 0, @@ -61124,7 +61539,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 4.7e-07, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/moonshotai/Kimi-K2.6": { "deprecation_date": "2026-08-19", @@ -61134,7 +61549,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/moonshotai/Kimi-K2.5-fp4": { "input_cost_per_token": 5e-07, @@ -61142,7 +61557,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/MiniMaxAI/MiniMax-M2.7": { "input_cost_per_token": 3e-07, @@ -61151,7 +61566,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 196608, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5": { "deprecation_date": "2026-06-22", @@ -61160,7 +61575,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 202752, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5.1": { "deprecation_date": "2026-07-10", @@ -61170,7 +61585,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 202752, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-0528": { "input_cost_per_token": 3e-06, @@ -61178,7 +61593,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 163840, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-Next-FP8": { "deprecation_date": "2026-05-14", @@ -61187,7 +61602,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-32B-Instruct": { "deprecation_date": "2026-02-25", @@ -61196,7 +61611,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-8B-Instruct": { "deprecation_date": "2026-04-16", @@ -61205,7 +61620,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Ministral-3-14B-Instruct-2512": { "input_cost_per_token": 2e-07, @@ -61213,7 +61628,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": { "input_cost_per_token": 6e-08, @@ -61221,7 +61636,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 131072, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-7B-Instruct-v0.3": { "input_cost_per_token": 2e-07, @@ -61229,7 +61644,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 32768, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/QwQ-32B": { "deprecation_date": "2025-11-13", @@ -61238,7 +61653,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 131072, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "cerebras/gemma-4-31b": { "input_cost_per_token": 9.9e-07, @@ -65365,5 +65780,263 @@ "supports_tool_choice": false, "supports_response_schema": true, "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": { + "cache_read_input_token_cost": 3.9e-07, + "input_cost_per_token": 2.1e-06, + "litellm_provider": "fireworks_ai", + "mode": "chat", + "output_cost_per_token": 6.6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models" + }, + "together_ai/arcee-ai/trinity-mini": { + "input_cost_per_token": 4.5e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.5e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/deepseek-coder-33b-instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Llama-70B": { + "input_cost_per_token": 2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-14B": { + "input_cost_per_token": 1.6e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-V4.1-Flash": { + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "vertex_ai/gemini-2.5-flash-native-audio": { + "input_cost_per_audio_token": 3e-06, + "input_cost_per_token": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "realtime", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_token": 2e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-2.5-flash-preview-tts": { + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "litellm_provider": "vertex_ai", + "mode": "audio_speech", + "output_cost_per_audio_token": 1e-05, + "output_cost_per_token": 1e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.1-flash-live-preview": { + "input_cost_per_audio_token": 3e-06, + "input_cost_per_second": 8.33333333333e-05, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "vertex_ai", + "mode": "realtime", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_token": 4.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.1-flash-tts-preview": { + "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "audio_speech", + "output_cost_per_audio_token": 2e-05, + "output_cost_per_token": 2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.5-transcribe": { + "input_cost_per_audio_token": 2e-06, + "input_cost_per_second": 5e-05, + "litellm_provider": "vertex_ai", + "mode": "audio_transcription", + "output_cost_per_token": 1.2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.5-transcribe-live": { + "input_cost_per_audio_token": 3.5e-06, + "input_cost_per_second": 8.33333333333e-05, + "litellm_provider": "vertex_ai", + "mode": "audio_transcription", + "output_cost_per_token": 2.1e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-omni-1.1-flash": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 9e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-robotics-er-2": { + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemma-4-26b-a4b-it": { + "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "together_ai/google/gemma-2-27b-it": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "gpt-5.5-cyber": { + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_token": 1.25e-05, + "litellm_provider": "openai", + "mode": "chat", + "output_cost_per_token": 7.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supports_reasoning": true + }, + "gpt-rosalind-research": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 5e-06, + "litellm_provider": "openai", + "mode": "chat", + "output_cost_per_token": 2.5e-05, + "source": "https://developers.openai.com/api/docs/pricing" + }, + "together_ai/meta-llama/Llama-3-8b-chat-hf": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.1-405B-Instruct": { + "input_cost_per_token": 3.5e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.2-1B-Instruct": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.2-3B-Instruct": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Meta-Llama-3-70B-Instruct-Turbo": { + "input_cost_per_token": 8.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Meta-Llama-3-8B-Instruct": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/NousResearch/Nous-Hermes-2-Mixtral-8x7B-DPO": { + "input_cost_per_token": 6e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/nvidia/Llama-3.1-Nemotron-70B-Instruct-HF": { + "input_cost_per_token": 8.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-1.5B-Instruct": { + "input_cost_per_token": 2e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-72B-Instruct": { + "input_cost_per_token": 9e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 9e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-VL-72B-Instruct": { + "input_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-14B-Instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-72B-Instruct": { + "input_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-Coder-32B-Instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-VL-72B-Instruct": { + "input_cost_per_token": 1.95e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://api.together.ai/v1/models" } } diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 7a110eff080..fb2f014e3d8 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -11503,6 +11503,12 @@ "description": "Enable content moderation to check for harmful content (harassment, hate speech, etc.).", "title": "Content Moderation Check" }, + "contextual_grounding_from_messages": { + "default": false, + "description": "ApplyGuardrail: when True, post-call scans of a request with no grounding_source / query content parts send the system and developer messages as the grounding source and the latest user message as the query, so the guardrail's contextual grounding policy can score the response. Bedrock bills contextual grounding units for these scans and rejects queries, sources and responses over its contextual grounding length limits, so leave this off for guardrails without a contextual grounding policy. Default False: plain messages are never sent as grounding context.", + "title": "Contextual Grounding From Messages", + "type": "boolean" + }, "credentials": { "anyOf": [ { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9d6fff0fa37..7d2f1b33823 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4,11 +4,12 @@ import os from collections.abc import Callable, Mapping from datetime import datetime from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple import httpx from pydantic import ( BaseModel, + BeforeValidator, ConfigDict, Field, Json, @@ -47,6 +48,7 @@ from litellm.types.proxy.carried_budget_state import ( ) from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry from litellm.types.router import RouterErrors, UpdateRouterConfig +from litellm.types.router_weights import validate_router_settings_dict from litellm.types.secret_managers.main import KeyManagementSystem from litellm.types.utils import ( CallTypes, @@ -656,6 +658,7 @@ class LiteLLMRoutes(enum.Enum): [ # user "/user/new", + "/management/v1/users/bulk", "/user/update", "/user/bulk_update", "/user/delete", @@ -890,6 +893,7 @@ class LiteLLMRoutes(enum.Enum): "/auto_router/validate_complexity_router_config", # Per-session auto-router read - the endpoint scopes the row to the caller's own key hash "/auto_router/session", + "/cost/predict-cache", # Agent registry - reads are role-scoped and writes are proxy-admin-gated # inside agent_endpoints/endpoints.py *agent_management_routes, @@ -1983,8 +1987,14 @@ class OrgMember(MemberBase): from litellm.models.team import TeamBase as TeamBase # noqa: E402 +RouterSettingsDict = Annotated[ + dict[str, object], + BeforeValidator(validate_router_settings_dict, json_schema_input_type=UpdateRouterConfig), +] + class NewTeamRequest(TeamBase): + router_settings: RouterSettingsDict | None = None model_aliases: dict | None = None tags: list | None = None guardrails: list[str] | None = None @@ -2082,7 +2092,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None enforced_batch_output_expires_after: dict | None = None enforced_file_expires_after: dict | None = None - router_settings: dict | None = None + router_settings: RouterSettingsDict | None = None access_group_ids: list[str] | None = None budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows default_team_member_models: list[str] | None = None # default allowed_models seeded onto new team members @@ -4421,6 +4431,9 @@ class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable): access_group_mcp_server_ids: list[str] | None = None access_group_agent_ids: list[str] | None = None access_group_details: tuple[TeamAccessGroupModelGrant, ...] | None = None + # Parent org's model ceiling, reported only to callers who can manage the team. + # None = no org or not a manager; [] or ["all-proxy-models"] = no ceiling. + organization_models: list[str] | None = None class TeamInfoResponseObject(TypedDict): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 0117edef627..6d61ad4d3e8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -895,6 +895,7 @@ async def common_checks( request_query_params=_safe_get_request_query_params(request=request), llm_router=llm_router, request=request, + team_id=valid_token.team_id if valid_token is not None else None, ) skip_all_budget_checks: Final = skip_budget_checks or ( @@ -4471,7 +4472,7 @@ async def stamp_matched_model_access_groups( async def can_key_call_model( model: str | list[str], - llm_model_list: list | None, + llm_model_list: Sequence[object] | None, valid_token: UserAPIKeyAuth, llm_router: litellm.Router | None, ) -> Literal[True]: @@ -4518,7 +4519,7 @@ async def can_key_call_model( async def can_key_call_resolved_model( model: str, - llm_model_list: list | None, + llm_model_list: Sequence[object] | None, valid_token: UserAPIKeyAuth, llm_router: litellm.Router | None, ) -> None: diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index be65c3b39ec..dc304a156cf 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -33,7 +33,7 @@ from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_me from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_ENDPOINT_MARKER, ) -from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS +from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment from litellm.types.utils import CustomPricingLiteLLMParams @@ -1736,7 +1736,7 @@ def _append_model_candidates(candidates: list[str], value: Any) -> None: candidates.extend(model for model in model_names if model) -def _dedupe_model_candidates(candidates: list[str]) -> list[str]: +def _dedupe_model_candidates(candidates: Collection[str]) -> list[str]: deduped: Final[list[str]] = [] for model in candidates: if model not in deduped: @@ -1845,13 +1845,42 @@ def _resolve_model_id_with_router(model_id: str | None, llm_router: Router | Non return model_id +def get_cache_prediction_deployments( + *, current_deployment_id: str, candidate_deployment_id: str, llm_router: Router, team_id: str | None +) -> tuple[Deployment, Deployment] | None: + current: Final = llm_router.get_deployment(current_deployment_id) + candidate: Final = llm_router.get_deployment(candidate_deployment_id) + if current is None or candidate is None: + return None + if any(deployment.model_info.team_id not in (None, team_id) for deployment in (current, candidate)): + return None + return current, candidate + + +def _cache_prediction_model_candidates( + request_data: Mapping[str, object], llm_router: Router | None, team_id: str | None +) -> tuple[str, ...]: + current_id: Final = request_data.get("current_deployment_id") + candidate_id: Final = request_data.get("candidate_deployment_id") + if llm_router is None or not isinstance(current_id, str) or not isinstance(candidate_id, str): + return () + deployments: Final = get_cache_prediction_deployments( + current_deployment_id=current_id, candidate_deployment_id=candidate_id, llm_router=llm_router, team_id=team_id + ) + return tuple(deployment.model_name for deployment in deployments) if deployments is not None else () + + def _extract_model_candidates_from_request( request_data: dict, route: str, request_headers: Mapping[str, object] | None = None, request_query_params: Mapping[str, object] | None = None, llm_router: Router | None = None, + team_id: str | None = None, ) -> list[str]: + if route == "/cost/predict-cache": + prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload + return _dedupe_model_candidates(prediction_models) candidates: Final[list[str]] = [] uses_model_routing_sources: Final = _route_uses_model_routing_sources(route=route) uses_header_or_query_model_sources: Final = _route_matches_any_marker( @@ -1945,6 +1974,7 @@ def get_model_from_request( request_query_params: Mapping[str, object] | None = None, llm_router: Router | None = None, request: Request | None = None, + team_id: str | None = None, ) -> str | list[str] | None: """Resolve the model(s) a request targets, for model-access and budget checks. @@ -1967,6 +1997,7 @@ def get_model_from_request( request_headers=request_headers, request_query_params=request_query_params, llm_router=llm_router, + team_id=team_id, ) model = _format_model_candidates(candidates) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index c0a76a4fc20..b7064802878 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -249,6 +249,7 @@ async def authenticate_user( if os.getenv("DATABASE_URL") is not None: response = await generate_key_helper_fn( + llm_router=None, request_type="key", **{ "user_role": LitellmUserRoles.PROXY_ADMIN, @@ -324,6 +325,7 @@ async def authenticate_user( await _rehash_password_if_needed(_user_row.user_id, password, _password) if os.getenv("DATABASE_URL") is not None: response = await generate_key_helper_fn( + llm_router=None, request_type="key", **{ "user_role": user_role, diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 4c64d4d7e23..166a0500cee 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -24,6 +24,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES: Final = frozenset( [ # user "/user/new", + "/management/v1/users/bulk", "/user/delete", "/management/v1/users/bulk_delete", "/user/bulk_update", @@ -760,6 +761,7 @@ class RouteChecks: _ADMIN_VIEWER_BLOCKED_WRITE_ROUTES = frozenset( [ "/user/new", + "/management/v1/users/bulk", "/user/delete", "/management/v1/users/bulk_delete", "/user/bulk_update", diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 22826f48b52..9491f77ecfc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -191,6 +191,7 @@ def _get_model_from_request_context( route: str, request: Request | None, llm_router: Any | None = None, + team_id: str | None = None, ) -> str | list[str] | None: return get_model_from_request( request_data=request_data, @@ -199,6 +200,7 @@ def _get_model_from_request_context( request_query_params=_safe_get_request_query_params(request=request), llm_router=llm_router, request=request, + team_id=team_id, ) @@ -217,7 +219,7 @@ async def _normalize_claude_model( return if request is not None and request.scope.get(_CLAUDE_MODEL_NORMALIZED) is True: return - requested: Final = _get_model_from_request_context(request_data, route, request, llm_router) + requested: Final = _get_model_from_request_context(request_data, route, request, llm_router, valid_token.team_id) if not isinstance(requested, str) or requested != request_data.get("model"): return if not requested.startswith("claude-router-") and not requested.lower().endswith("[1m]"): @@ -876,6 +878,7 @@ async def _auto_register_jwt_mapping( # the NOT NULL @id constraint. Every successful key-creation caller (e.g. # /key/generate) passes table_name="key" explicitly. key_data: Final = await generate_key_helper_fn( + llm_router=None, request_type="key", table_name="key", team_id=team_id, @@ -1652,6 +1655,7 @@ async def _user_api_key_auth_builder( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) skip_budget_checks = False if model is not None and llm_router is not None: @@ -1692,6 +1696,7 @@ async def _user_api_key_auth_builder( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) ), ) @@ -2091,6 +2096,7 @@ async def _user_api_key_auth_builder( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) skip_budget_checks = False if model is not None and llm_router is not None: @@ -2209,6 +2215,7 @@ async def _user_api_key_auth_builder( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) current_models = _get_model_names_for_budget_checks(model=current_model) @@ -2239,6 +2246,7 @@ async def _user_api_key_auth_builder( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) current_models = _get_model_names_for_budget_checks(model=current_model) @@ -2734,6 +2742,7 @@ async def _run_centralized_common_checks( route=route, request=request, llm_router=llm_router, + team_id=user_api_key_auth_obj.team_id, ) # Pin the metadata variable name (litellm_metadata vs metadata) before @@ -2850,12 +2859,14 @@ def _should_skip_budget_checks( route: str, request: Request | None, llm_router: Any | None, + team_id: str | None = None, ) -> bool: model: Final = _get_model_from_request_context( request_data=request_data, route=route, request=request, llm_router=llm_router, + team_id=team_id, ) if model is not None and llm_router is not None: return _is_model_cost_zero(model=model, llm_router=llm_router) @@ -3301,6 +3312,7 @@ async def _enforce_key_and_fallback_model_access( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) if model is not None: @@ -3408,6 +3420,7 @@ async def _run_post_custom_auth_checks( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) current_models = _get_model_names_for_budget_checks(model=current_model) @@ -3449,6 +3462,7 @@ async def _run_post_custom_auth_checks( route=route, request=request, llm_router=llm_router, + team_id=valid_token.team_id, ) current_models = _get_model_names_for_budget_checks(model=current_model) diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index cb867cf9e61..a5a2675ed6e 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -580,8 +580,8 @@ What the command changed is recorded in `~/.litellm/claude_configure_state.json` `lite configure claude`, `lite login --config-claude`, `lite up` and `lite autoroute up` also install a status line (`~/.litellm/statusline.py`, registered as `statusLine` in `~/.claude/settings.json` unless you already run one) that shows which model the auto-router actually served the last turn and, once the proxy has recorded the session, what the session cost against the router's savings baseline: ``` -claude-auto · Routed to: claude-haiku-4-5 -63% vs Claude Opus 5 -LiteLLM ████████░░░░░░░░░░░░░░░░ $0.14 +Routed to: claude-haiku-4-5 -63% vs Claude Opus 5 +claude-auto ████████░░░░░░░░░░░░░░░░ $0.14 Claude Opus 5 ████████████████████████ $0.38 ``` diff --git a/litellm/proxy/client/cli/commands/pi.py b/litellm/proxy/client/cli/commands/pi.py index 9810e81ae36..5c749959638 100644 --- a/litellm/proxy/client/cli/commands/pi.py +++ b/litellm/proxy/client/cli/commands/pi.py @@ -10,7 +10,7 @@ import os import tempfile from collections.abc import Callable, Mapping from dataclasses import dataclass -from enum import StrEnum +from enum import Enum from pathlib import Path from types import MappingProxyType from typing import Annotated, Final @@ -25,7 +25,7 @@ LITELLM_PROXY_API_KEY_ENV: Final = "LITELLM_PROXY_API_KEY" _REJECTED_STATUSES: Final = frozenset((401, 403)) -class ListingFailure(StrEnum): +class ListingFailure(str, Enum): """Why a proxy could not be listed, decided once where the HTTP outcome is classified. `unreachable` means no response at all; the other kinds prove the proxy answered, so callers diff --git a/litellm/proxy/client/cli/commands/statusline_script.py b/litellm/proxy/client/cli/commands/statusline_script.py index 5493586f627..47be3888a58 100644 --- a/litellm/proxy/client/cli/commands/statusline_script.py +++ b/litellm/proxy/client/cli/commands/statusline_script.py @@ -28,6 +28,7 @@ import os import sys import tempfile import time +import unicodedata import urllib.error import urllib.request from collections.abc import Callable, Mapping @@ -42,7 +43,6 @@ FETCH_TIMEOUT_SECONDS: Final = 3 BAR_WIDTH: Final = 24 BAR_FULL: Final = "\u2588" BAR_EMPTY: Final = "\u2591" -SEPARATOR: Final = " \u00b7 " TRANSCRIPT_SCAN_LIMIT_BYTES: Final = 4 * 1024 * 1024 CLAUDE_BASE_URL_ENV_KEYS: Final = ("ANTHROPIC_BASE_URL",) CLAUDE_API_KEY_ENV_KEYS: Final = ("ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_KEY") @@ -50,7 +50,6 @@ CODEX_BASE_URL_ENV_KEYS: Final = ("OPENAI_BASE_URL",) CODEX_API_KEY_ENV_KEYS: Final = ("OPENAI_API_KEY",) CODEX_STOP_EVENT: Final = "Stop" SYNTHETIC_MODEL: Final = "" -LITELLM_LABEL: Final = "LiteLLM" RESET: Final = "\033[0m" BOLD: Final = "\033[1m" DIM: Final = "\033[90m" @@ -302,31 +301,37 @@ def _bar(fraction: float, color: str, width: int, use_color: bool) -> str: return f"{color}{BAR_FULL * filled}{DIM}{BAR_EMPTY * (width - filled)}{RESET}" +def _display_width(label: str) -> int: + return sum( + 2 if unicodedata.east_asian_width(character) in ("W", "F") else 1 + for character in label + if unicodedata.category(character) not in ("Mn", "Me") + ) + + def render(model: str, session: Session | None, config_dir: Path, use_color: bool, bar_width: int = BAR_WIDTH) -> str: def paint(code: str, text: str) -> str: return f"{code}{text}{RESET}" if use_color else text routed: Final = paint(BOLD, f"Routed to: {model}") - if session is None: + if session is None or session.baseline_model is None or session.baseline_spend <= 0: return routed - header: Final = f"{session.router_name}{SEPARATOR}{routed}" - if session.baseline_model is None or session.baseline_spend <= 0: - return header reference: Final = baseline_label(session.baseline_model, config_dir) pct: Final = (session.baseline_spend - session.spend) / session.baseline_spend * 100 delta: Final = paint(LITELLM_COLOR, f"{'-' if pct >= 0 else '+'}{abs(round(pct))}% vs {reference}") peak: Final = max(session.spend, session.baseline_spend) - label_width: Final = max(len(LITELLM_LABEL), len(reference)) + label_width: Final = max(_display_width(session.router_name), _display_width(reference)) rows: Final = ( - (LITELLM_LABEL, session.spend, LITELLM_COLOR), + (session.router_name, session.spend, LITELLM_COLOR), (reference, session.baseline_spend, BASELINE_COLOR), ) lines: Final = ( - f"{paint(DIM, label.ljust(label_width))} {_bar(amount / peak, color, bar_width, use_color)} " + f"{paint(DIM, label + ' ' * (label_width - _display_width(label)))} " + f"{_bar(amount / peak, color, bar_width, use_color)} " f"{paint(DIM, f'${amount:.2f}')}" for label, amount, color in rows ) - return "\n".join((f"{header} {delta}", *lines)) + return "\n".join((f"{routed} {delta}", *lines)) def color_enabled(env: Mapping[str, str]) -> bool: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c73a888ba58..4a4daa68cce 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -14,6 +14,7 @@ import httpx import orjson from fastapi import HTTPException, Request, status from fastapi.responses import JSONResponse, Response, StreamingResponse +from pydantic import ValidationError from starlette.types import Receive, Scope, Send import litellm @@ -76,6 +77,7 @@ from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_di from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.types.guardrails import GuardrailEventHooks from litellm.types.router import RouterRateLimitError +from litellm.types.router_weights import validate_router_weights _LateResponseT = TypeVar("_LateResponseT", bound=Response) _LlmCallT = TypeVar("_LlmCallT") @@ -1571,6 +1573,9 @@ class ProxyBaseLLMRequestProcessing: ) -> dict: exclude_values: Final = {"", None, "None"} hidden_params = hidden_params or {} + resolved_call_id: Final = ( + call_id or hidden_params.get("litellm_call_id") or (request_data or {}).get("litellm_call_id") + ) timing_values: Final = _timing_values( hidden_params=hidden_params, logging_obj=litellm_logging_obj, @@ -1598,7 +1603,7 @@ class ProxyBaseLLMRequestProcessing: classifier_cost: Final = _classifier_cost_from_request_data(request_data) headers: Final = { - "x-litellm-call-id": call_id, + "x-litellm-call-id": resolved_call_id, "x-litellm-model-id": model_id, "x-litellm-model-name": model_name, "x-litellm-cache-key": cache_key, @@ -1936,6 +1941,13 @@ class ProxyBaseLLMRequestProcessing: # This avoids expensive Router instantiation on each request if router_settings is not None: self.data["router_settings_override"] = router_settings + try: + self.data["_router_weights"] = validate_router_weights(router_settings.get("weights")) + except ValidationError: + self.data["_router_weights"] = None + verbose_proxy_logger.warning( + "Ignoring invalid saved router weights; update team/key router_settings" + ) alias_target: Final = await _resolve_per_request_model_group_alias( requested_model=self.data.get("model"), router_settings=router_settings, diff --git a/litellm/proxy/common_utils/openai_error_payload.py b/litellm/proxy/common_utils/openai_error_payload.py index 89f735ee8b6..fe23ab2c4b6 100644 --- a/litellm/proxy/common_utils/openai_error_payload.py +++ b/litellm/proxy/common_utils/openai_error_payload.py @@ -8,6 +8,8 @@ from typing import Final from fastapi import status +from litellm.constants import STRINGIFIED_NONE + _OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType( { status.HTTP_401_UNAUTHORIZED: "authentication_error", @@ -35,7 +37,7 @@ def openai_error_type(exc: object, status_code: int) -> str: """OpenAI types ``error.type`` as a required string, so an exception carrying none falls back to the type its status code stands for.""" carried: Final = attribute_of(exc, "type") - if isinstance(carried, str): + if isinstance(carried, str) and carried != STRINGIFIED_NONE: return carried mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code) if mapped is not None: @@ -49,4 +51,4 @@ def openai_error_param(exc: object) -> str | None: """OpenAI types ``error.param`` as nullable, so an exception carrying none serializes as JSON ``null``.""" carried: Final = attribute_of(exc, "param") - return carried if isinstance(carried, str) else None + return carried if isinstance(carried, str) and carried != STRINGIFIED_NONE else None diff --git a/litellm/proxy/common_utils/prompt_cache_pricing.py b/litellm/proxy/common_utils/prompt_cache_pricing.py new file mode 100644 index 00000000000..ff070853b46 --- /dev/null +++ b/litellm/proxy/common_utils/prompt_cache_pricing.py @@ -0,0 +1,91 @@ +from collections.abc import Mapping +from math import isfinite +from typing import Final + +from pydantic import TypeAdapter + +import litellm +from litellm.cost_calculator import ( + _select_model_name_for_cost_calc, # pyright: ignore[reportPrivateUsage] # shares completion_cost's deployment tariff selection + completion_cost, # pyright: ignore[reportUnknownVariableType] # legacy optional parameters are untyped +) +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.types.management_endpoints.prompt_cache_prediction import CacheTokenBuckets +from litellm.types.utils import CacheCreationTokenDetails, ModelResponse, PromptTokensDetailsWrapper, Usage + +_PRICE_ENTRY: Final = TypeAdapter(Mapping[str, object]) + + +def _valid_price(value: object) -> bool: + return isinstance(value, (int, float)) and not isinstance(value, bool) and isfinite(value) and value >= 0 + + +def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets) -> bool: + required: Final = ( + ("input_cost_per_token", True), + ("cache_read_input_token_cost", tokens.cache_read_input_tokens > 0), + ("cache_creation_input_token_cost", tokens.cache_creation_5m_input_tokens > 0), + ("cache_creation_input_token_cost_above_1hr", tokens.cache_creation_1h_input_tokens > 0), + ) + if any(needed and not _valid_price(prices.get(key)) for key, needed in required): + return False + return all( + _valid_price(value) + for key, value in prices.items() + if value is not None and any(needed and key.startswith(f"{base}_above_") for base, needed in required) + ) + + +def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> float | None: + try: + selected_model: Final = _select_model_name_for_cost_calc( + model=model, + completion_response=None, + custom_pricing=True, + custom_llm_provider="anthropic", + router_model_id=deployment_id, + ) + if selected_model is None: + return None + model_info: Final = litellm.get_model_info(model=selected_model, custom_llm_provider="anthropic") + registry: Final = _PRICE_ENTRY.validate_python(litellm.model_cost) # pyright: ignore[reportUnknownMemberType] # legacy registry is validated at this boundary + price_entry: Final = registry.get(model_info["key"]) + if price_entry is None: + return None + prices: Final = _PRICE_ENTRY.validate_python(price_entry) + if not _has_required_prices(prices, tokens): + return None + usage: Final = Usage( + prompt_tokens=tokens.total_tokens, + completion_tokens=0, + total_tokens=tokens.total_tokens, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=tokens.cache_read_input_tokens, + cache_creation_tokens=tokens.cache_creation_5m_input_tokens + tokens.cache_creation_1h_input_tokens, + cache_creation_token_details=CacheCreationTokenDetails( + ephemeral_5m_input_tokens=tokens.cache_creation_5m_input_tokens, + ephemeral_1h_input_tokens=tokens.cache_creation_1h_input_tokens, + ), + ), + ) + logging_obj: Final = Logging( + model=model, + messages=[], # mutable-ok: Logging requires a list + stream=False, + call_type="completion", + start_time=None, + litellm_call_id="prompt-cache-prediction", + function_id="prompt-cache-prediction", + ) + completion_cost( + completion_response=ModelResponse(model=model, usage=usage), + model=model, + custom_llm_provider="anthropic", + custom_pricing=True, + router_model_id=deployment_id, + litellm_logging_obj=logging_obj, + ) + cost: Final = logging_obj.cost_breakdown.get("input_cost") if logging_obj.cost_breakdown is not None else None + return cost if cost is not None and _valid_price(cost) else None + except Exception: # noqa: BLE001 # the shared pricing owners raise plain Exception for unpriceable models + return None diff --git a/litellm/proxy/compliance_checks.py b/litellm/proxy/compliance_checks.py index ff311911742..9d2f2dc7c69 100644 --- a/litellm/proxy/compliance_checks.py +++ b/litellm/proxy/compliance_checks.py @@ -26,7 +26,7 @@ class ComplianceChecker: def __init__(self, data: ComplianceCheckRequest): self.data = data - self.guardrails = data.guardrail_information or [] + self.guardrails = tuple(g for g in data.guardrail_information or () if g.get("guardrail_status") != "not_run") def _get_guardrails_by_mode(self, mode: str) -> list[dict]: """ diff --git a/litellm/proxy/db/health_check_latest.py b/litellm/proxy/db/health_check_latest.py index 35bc838379c..21438f095bb 100644 --- a/litellm/proxy/db/health_check_latest.py +++ b/litellm/proxy/db/health_check_latest.py @@ -74,10 +74,14 @@ class LatestHealthCheckRow(BaseModel): _ROWS_ADAPTER: Final = TypeAdapter(tuple[LatestHealthCheckRow, ...]) +async def query_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: + rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) + return _ROWS_ADAPTER.validate_python(rows) + + async def fetch_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: try: - rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) - return _ROWS_ADAPTER.validate_python(rows) + return await query_latest_health_checks(prisma_client) except Exception as query_err: # noqa: BLE001 # health decorates other reads; a driver error must not fail them verbose_proxy_logger.error("Error getting all latest health checks: %s", query_err) return () diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 6a7ac4361b9..2c407d91a48 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -244,6 +244,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): prompt_attack_threshold: float | None = 0.5, pii_confidence_threshold: float | None = 0.5, chunk_budget_chars: int = BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS, + contextual_grounding_from_messages: bool = False, streaming_buffer_until_moderated: bool | None = None, streaming_sampling_rate: int | None = None, streaming_end_of_stream_only: bool | None = None, @@ -265,6 +266,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self.guardrailVersion = guardrailVersion self.guardrail_provider = "bedrock" self.chunk_budget_chars = chunk_budget_chars + self.contextual_grounding_from_messages = contextual_grounding_from_messages self.experimental_use_latest_role_message_only = bool(kwargs.get("experimental_use_latest_role_message_only")) # Resource-less, detect-only InvokeGuardrailChecks mode. Present `checks` @@ -459,8 +461,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """ Flatten a message into text blocks, preserving any contextual-grounding qualifier carried by the content-block ``type`` (grounding_source / query). - Untagged text keeps ``qualifier=None`` so the payload is unchanged for - callers that do not use grounding. + Untagged text keeps ``qualifier=None``; the OUTPUT scan decides whether to + derive grounding qualifiers from it. """ content: Final = message.get("content") if content is None: @@ -493,6 +495,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): result carrying externally-influenced content can supply fake evidence for the contextual-grounding check to grade the response against. ``query`` is accepted from any role (it is the user's question). + + With ``contextual_grounding_from_messages`` on, a request with no tagged blocks + falls back to the plain messages: system / developer text is the grounding + source and the latest user message is the query. """ grounding: Final[list[QualifiedTextBlock]] = [] for message in messages or []: @@ -504,7 +510,33 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): and role in _GROUNDING_SOURCE_TRUSTED_ROLES ): grounding.append(block) - return grounding + if grounding or not self.contextual_grounding_from_messages: + return grounding + return self._derive_grounding_blocks_from_plain_messages(messages) + + def _derive_grounding_blocks_from_plain_messages( + self, messages: list[AllMessageValues] | None + ) -> list[QualifiedTextBlock]: + if not messages: + return [] + latest_user_index: Final = self._find_latest_message_index(messages, target_role="user") + if latest_user_index is None: + return [] + sources: Final = tuple( + QualifiedTextBlock(text=block.text, qualifier="grounding_source") + for message in messages + if message.get("role") in _GROUNDING_SOURCE_TRUSTED_ROLES + for block in self.get_content_items_for_message(message=message) or [] + if block.text + ) + queries: Final = tuple( + QualifiedTextBlock(text=block.text, qualifier="query") + for block in self.get_content_items_for_message(message=messages[latest_user_index]) or [] + if block.text + ) + if not sources or not queries: + return [] + return [*sources, *queries] def supports_scan_only_tool_results(self) -> bool: return self.experimental_use_latest_role_message_only is not True @@ -3210,6 +3242,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): bedrock_response = await self.make_bedrock_api_request( source="OUTPUT", response=synthetic_response, + messages=request_data.get("messages"), request_data=request_data, logging_event_type=_log_hook, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index d8296003ae9..3d1a173635e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -7,7 +7,7 @@ import fnmatch import os -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional import httpx @@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.llms.openai import ChatCompletionToolParam +from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( GenericGuardrailAPIMetadata, GenericGuardrailAPIRequest, @@ -150,6 +150,26 @@ def _extract_inbound_headers( return None +def _structured_rows_to_write_back( + original_rows: Sequence[AllMessageValues] | None, + shown_rows: Sequence[AllMessageValues] | None, + returned_rows: Sequence[AllMessageValues], +) -> tuple[AllMessageValues, ...] | None: + """The request model drops row keys its message types do not declare, so a + row the server echoes back verbatim is restored to the original row object. + A server that echoes every row back unchanged has not rewritten anything + per row, so its answer is read from texts, as it was before rows could be + returned at all.""" + if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows): + return tuple(returned_rows) + if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)): + return None + return tuple( + original if returned == shown else returned + for original, shown, returned in zip(original_rows, shown_rows, returned_rows) + ) + + class GenericGuardrailAPI(CustomGuardrail): """ Generic Guardrail API integration for LiteLLM. @@ -322,6 +342,8 @@ class GenericGuardrailAPI(CustomGuardrail): texts: list, images: list[str] | None, tools: list[ChatCompletionToolParam] | None, + structured_messages: Sequence[AllMessageValues] | None, + shown_messages: Sequence[AllMessageValues] | None, guardrail_response: GenericGuardrailAPIResponse, ) -> GenericGuardrailAPIInputs: # Action is NONE or no modifications needed @@ -336,6 +358,13 @@ class GenericGuardrailAPI(CustomGuardrail): return_inputs["tools"] = guardrail_response.tools elif tools: return_inputs["tools"] = tools + rows_to_write_back: Final = ( + _structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages) + if guardrail_response.structured_messages + else None + ) + if rows_to_write_back is not None: + return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list if guardrail_response.stream_holdback_chars is not None: return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars return return_inputs @@ -473,6 +502,8 @@ class GenericGuardrailAPI(CustomGuardrail): texts=texts, images=images, tools=tools, + structured_messages=structured_messages, + shown_messages=guardrail_request.structured_messages, guardrail_response=guardrail_response, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py index 3e5d8fb311d..e54e07b6a1b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py +++ b/litellm/proxy/guardrails/guardrail_hooks/javelin/javelin.py @@ -44,7 +44,7 @@ class JavelinGuardrail(CustomGuardrail): application: str | None = None, **kwargs, ): - f""" + """ Initialize the JavelinGuardrail class. This calls: {api_base}/{api_version}/guardrail/{guardrail_name}/apply diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py index a911c78ddc3..95db4bd4f77 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/__init__.py @@ -15,7 +15,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" # We check the raw guardrail dict because LitellmParams normalizes None → False, # making it impossible to distinguish "not set" from "explicitly false" via litellm_params. _raw_default_on: Final = cast(dict[str, Any], guardrail).get("litellm_params", {}).get("default_on") - _default_on: Final = False if _raw_default_on is False else True + _default_on: Final = _raw_default_on is not False _callback: Final = MCPEndUserPermissionGuardrail( guardrail_name=guardrail.get("guardrail_name", ""), diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index fde40111d49..a7c93e63b32 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -1,11 +1,13 @@ -from collections.abc import AsyncGenerator, Mapping, Sequence +import time +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence from enum import Enum, auto -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal import httpx from fastapi import HTTPException if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel import json @@ -23,6 +25,7 @@ from litellm.litellm_core_utils.core_helpers import ( get_or_create_metadata_bucket, ) from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, get_async_httpx_client, httpxSpecialProvider, ) @@ -52,6 +55,7 @@ from litellm.types.utils import ( CallTypes, CallTypesLiteral, Choices, + GenericGuardrailAPIInputs, GuardrailStatus, ModelResponse, ModelResponseStream, @@ -118,8 +122,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): Supports: - Pre-call sanitization (sanitizeUserPrompt) - Post-call sanitization (sanitizeModelResponse) + - logging_only: scans the completed response after it reaches the client and + records the verdict in spend logs without blocking """ + use_native_lifecycle_hooks: ClassVar[bool] = True + @classmethod def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: return [ @@ -128,6 +136,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): GuardrailEventHooks.post_call, GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.during_mcp_call, + GuardrailEventHooks.logging_only, ] def __init__( @@ -138,6 +147,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): credentials: VERTEX_CREDENTIALS_TYPES | None = None, api_endpoint: str | None = None, sanitize_error_detail: "bool | None" = True, + async_handler: AsyncHTTPHandler | None = None, + access_token_provider: Callable[[], Awaitable[tuple[str, str]]] | None = None, **kwargs, ): # Set supported event hooks if not already provided @@ -154,7 +165,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): VertexBase.__init__(self) # Then set our attributes (this ensures project_id is not overwritten) - self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + self.async_handler = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + self.access_token_provider = access_token_provider self.template_id = template_id self.project_id = project_id self.location = location or "us-central1" @@ -278,11 +292,14 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): If file_bytes and file_type are provided, file prompt sanitization is performed. """ # Get access token using VertexBase auth - access_token, resolved_project_id = await self._ensure_access_token_async( - credentials=self.credentials, - project_id=self.project_id, - custom_llm_provider="vertex_ai", - ) + if self.access_token_provider is not None: + access_token, resolved_project_id = await self.access_token_provider() + else: + access_token, resolved_project_id = await self._ensure_access_token_async( + credentials=self.credentials, + project_id=self.project_id, + custom_llm_provider="vertex_ai", + ) # Use resolved project ID if not explicitly set if not self.project_id and resolved_project_id: @@ -1096,6 +1113,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): add_guardrail_to_applied_guardrails_header, ) + if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: + async for chunk in response: + yield chunk + return + all_chunks: Final[Sequence[object]] = tuple([chunk async for chunk in response]) if not all_chunks or self._is_terminal_error_stream(all_chunks): @@ -1213,6 +1235,60 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): for chunk in all_chunks: yield chunk + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: "LiteLLMLoggingObj | None" = None, + ) -> GenericGuardrailAPIInputs: + content: Final = "\n".join(text for text in inputs.get("texts") or () if text) + if not content: + return inputs + + source: Final[Literal["user_prompt", "model_response"]] = ( + "user_prompt" if input_type == "request" else "model_response" + ) + start_time: Final = time.time() + try: + armor_response: Final = await self.make_model_armor_request( + content=content, source=source, request_data=request_data + ) + except (ModelArmorAPIError, httpx.HTTPError) as e: + error_end_time: Final = time.time() + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=str(e), + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + guardrail_provider="model_armor", + start_time=start_time, + end_time=error_end_time, + duration=error_end_time - start_time, + ) + return inputs + + flagged: Final = self._should_block_content(armor_response, allow_sanitization=False) + end_time: Final = time.time() + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=self._build_logging_response(armor_response), + request_data=request_data, + guardrail_status="guardrail_flagged" if flagged else "success", + guardrail_provider="model_armor", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + if flagged and not self._event_hook_is_event_type(GuardrailEventHooks.logging_only): + raise HTTPException( + status_code=400, + detail=self._build_block_error_detail( + "Response blocked by Model Armor" if input_type == "response" else "Content blocked by Model Armor", + armor_response, + ), + ) + return inputs + @staticmethod def get_config_model() -> type["GuardrailConfigModel"] | None: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 3bc0dfabefc..9002e2aea07 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -1600,8 +1600,8 @@ class PanwPrismaAirsHandler(CustomGuardrail): Args: texts: Flattened text entries from the framework. - messages: Original request messages (request_data["messages"]), - NOT structured_messages (which may have injected system content). + messages: The structured messages the framework flattened into ``texts``, + hoisted top-level system prompt included, so positions line up. Returns a set of scannable indices, or None on count mismatch or no user/developer message (safety fallback to existing role-filter behavior). @@ -1788,15 +1788,10 @@ class PanwPrismaAirsHandler(CustomGuardrail): structured_messages: Final = inputs.get("structured_messages") if structured_messages: # For Anthropic /v1/messages: default to latest-user-only scanning. - # Uses request_data["messages"] (original format), NOT structured_messages - # (which has injected system content from adapter translation). if self._use_latest_user_only(request_data, logging_obj): - original_messages: Final = request_data.get("messages") - if original_messages: - scannable_indices = self._get_latest_user_text_indices(texts, original_messages) + scannable_indices = self._get_latest_user_text_indices(texts, structured_messages) # Fall through to existing role filtering if: # - not Anthropic, OR flag explicitly False, OR - # - no original messages, OR # - latest-user extraction returned None (no user / count mismatch) if scannable_indices is None: scannable_indices = self._get_scannable_text_indices(texts, structured_messages) diff --git a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py index 0954fe1698a..7e43566f224 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py +++ b/litellm/proxy/guardrails/guardrail_hooks/prompt_security/prompt_security.py @@ -2,6 +2,7 @@ import asyncio import base64 import os from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Optional import httpx @@ -14,11 +15,13 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, ) +from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts, message_with_slot_texts from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: @@ -27,12 +30,37 @@ if TYPE_CHECKING: _SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0 +_SANITIZE_FILE_QUEUED_STATUSES: Final = frozenset({"created", "in progress"}) +_PROTECT_ROLES: Final = frozenset({"system", "user", "assistant"}) class PromptSecurityGuardrailMissingSecrets(Exception): pass +def _inputs_with_structured_messages( + inputs: GenericGuardrailAPIInputs, rewritten_messages: Sequence[AllMessageValues] | None +) -> GenericGuardrailAPIInputs: + if rewritten_messages is None: + return inputs + patched: Final[GenericGuardrailAPIInputs] = { + **inputs, + "structured_messages": list(rewritten_messages), # mutable-ok: the TypedDict field is declared as a list + } + return patched + + +def _inputs_with_modifications( + inputs: GenericGuardrailAPIInputs, + modified_texts: list[str], + rewritten_messages: Sequence[AllMessageValues] | None, +) -> GenericGuardrailAPIInputs: + if not modified_texts: + return _inputs_with_structured_messages(inputs, rewritten_messages) + with_texts: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": modified_texts} + return _inputs_with_structured_messages(with_texts, rewritten_messages) + + class _ProtectVerdict(TypedDict, total=False): """One side (``prompt`` or ``response``) of an ``/api/protect`` verdict.""" @@ -275,14 +303,39 @@ class PromptSecurityGuardrail(CustomGuardrail): detail="Blocked by Prompt Security, Violations: " + ", ".join(violations), ) elif action == "modify": - # Extract modified texts from modified_messages modified_messages: Final = result.get("modified_messages", []) - modified_texts: Final = self._extract_texts_from_messages(modified_messages) - if modified_texts: - inputs["texts"] = modified_texts + return _inputs_with_modifications( + inputs, + self._extract_texts_from_messages(modified_messages), + self._structured_messages_with_modifications(structured_messages, modified_messages), + ) return inputs + def _is_sent_to_protect(self, message: Mapping[str, object]) -> bool: + return self.check_tool_results or message.get("role") in _PROTECT_ROLES + + def _structured_messages_with_modifications( + self, + structured_messages: Sequence[AllMessageValues], + modified_messages: Sequence[Mapping[str, object]], + ) -> tuple[AllMessageValues, ...] | None: + sent_indices: Final = tuple( + index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message) + ) + if not sent_indices or len(sent_indices) != len(modified_messages): + return None + rewritten: Final = tuple( + message_with_slot_texts(structured_messages[index], self._extract_texts_from_messages((modified,))) + for index, modified in zip(sent_indices, modified_messages) + ) + replacements: Final = MappingProxyType( + {index: message for index, message in zip(sent_indices, rewritten) if message is not None} + ) + if len(replacements) != len(sent_indices): + return None + return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages)) + async def _apply_guardrail_on_response( self, inputs: GenericGuardrailAPIInputs, @@ -346,19 +399,7 @@ class PromptSecurityGuardrail(CustomGuardrail): return inputs def _extract_texts_from_messages(self, messages: Sequence[Mapping[str, object]]) -> list[str]: - """Extract text content from messages.""" - texts: Final = [] - for message in messages: - content = message.get("content") - if isinstance(content, str): - texts.append(content) - elif isinstance(content, list): - for item in content: - if isinstance(item, dict) and item.get("type") == "text": - text = item.get("text") - if text: - texts.append(text) - return texts + return [text for message in messages for text in message_slot_texts(message)] async def _process_standalone_images(self, images: list[str], user_api_key_alias: str | None) -> None: """Process standalone images from inputs (data URLs).""" @@ -512,16 +553,18 @@ class PromptSecurityGuardrail(CustomGuardrail): "metadata": result.get("metadata", {}), "violations": result.get("metadata", {}).get("violations", []), } - elif status == "in progress": - verbose_proxy_logger.debug( - "Prompt Security Guardrail: File sanitization in progress (attempt %d/%d)", - attempt + 1, - self.max_poll_attempts, - ) - continue - else: + + if status not in _SANITIZE_FILE_QUEUED_STATUSES: raise HTTPException(status_code=500, detail=f"Unexpected sanitization status: {status}") + verbose_proxy_logger.debug( + "Prompt Security Guardrail: File sanitization status=%s for jobId=%s (attempt %d/%d)", + status, + job_id, + attempt + 1, + self.max_poll_attempts, + ) + raise HTTPException(status_code=408, detail="File sanitization timeout") def _raise_if_file_blocked(self, sanitization_result: _SanitizeResult, resource_name: str) -> None: @@ -678,14 +721,13 @@ class PromptSecurityGuardrail(CustomGuardrail): This allows checking tool results for indirect prompt injection when enabled. """ - supported_roles: Final = ["system", "user", "assistant"] filtered_messages: Final = [] transformed_count = 0 filtered_count = 0 for message in messages: role = message.get("role", "") - if role in supported_roles: + if role in _PROTECT_ROLES: filtered_messages.append(message) else: if self.check_tool_results: diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 16369abbfb0..7858adeb55d 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -23,6 +23,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail): prompt_attack_threshold=litellm_params.prompt_attack_threshold, pii_confidence_threshold=litellm_params.pii_confidence_threshold, chunk_budget_chars=litellm_params.chunk_budget_chars, + contextual_grounding_from_messages=litellm_params.contextual_grounding_from_messages, default_on=litellm_params.default_on, disable_exception_on_block=litellm_params.disable_exception_on_block, mask_request_content=litellm_params.mask_request_content, diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 6259efb6654..556b6a4e919 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -42,7 +42,7 @@ if TYPE_CHECKING: router: Final = APIRouter() _EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({}) -_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"passed": 0, "flagged": 1, "blocked": 2}) +_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"not_run": 0, "passed": 1, "flagged": 2, "blocked": 3}) _T = TypeVar("_T") @@ -325,7 +325,7 @@ class UsageDetailResponse(BaseModel): class UsageLogEntry(BaseModel): id: str timestamp: str - action: str # blocked | passed | flagged + action: str # blocked | passed | flagged | not_run score: float | None latency_ms: float | None model: str | None diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index 797323794d2..7e11b69108b 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -193,10 +193,12 @@ async def _upsert_rows_with_retry( def guardrail_status_to_action(status: str | None) -> str: - """Map StandardLogging guardrail_status to blocked/passed/flagged.""" + """Map StandardLogging guardrail_status to blocked/passed/flagged/not_run.""" if not status: return "passed" s: Final = (status or "").lower() + if s == "not_run": + return "not_run" if "intervened" in s or "block" in s: return "blocked" if "flagged" in s or "fail" in s or "error" in s: @@ -354,37 +356,49 @@ async def process_spend_logs_guardrail_usage( "flagged_count": 0, } ) - index_rows: Final[list[dict[str, object]]] = [] + index_rows_by_key: Final[dict[tuple[str, str], dict[str, object]]] = {} for payload in logs_to_process: request_id = payload.get("request_id") start_time = _parse_payload_start_time(payload) - if not request_id or start_time is None: + if not isinstance(request_id, str) or not request_id or start_time is None: continue date_key = _date_str(start_time) - for entry in _parse_guardrail_info_from_payload(payload): - guardrail_id = entry.get("guardrail_id") or entry.get("guardrail_name") or "" - if not guardrail_id: + entries = _parse_guardrail_info_from_payload(payload) + ids_by_name = MappingProxyType( + { + e["guardrail_name"]: e["guardrail_id"] + for e in entries + if e.get("guardrail_id") and isinstance(e.get("guardrail_name"), str) and e["guardrail_name"] + } + ) + for entry in entries: + raw_name = entry.get("guardrail_name") + guardrail_name = raw_name if isinstance(raw_name, str) else "" + guardrail_id = entry.get("guardrail_id") or ids_by_name.get(guardrail_name) or guardrail_name + if not isinstance(guardrail_id, str) or not guardrail_id: continue - key = _MetricsKey(guardrail_id, date_key) - daily_guardrail[key]["requests_evaluated"] += 1 action = guardrail_status_to_action(entry.get("guardrail_status")) - if action == "passed": - daily_guardrail[key]["passed_count"] += 1 - elif action == "blocked": - daily_guardrail[key]["blocked_count"] += 1 - else: - daily_guardrail[key]["flagged_count"] += 1 + if action != "not_run": + key = _MetricsKey(guardrail_id, date_key) + daily_guardrail[key]["requests_evaluated"] += 1 + if action == "passed": + daily_guardrail[key]["passed_count"] += 1 + elif action == "blocked": + daily_guardrail[key]["blocked_count"] += 1 + else: + daily_guardrail[key]["flagged_count"] += 1 policy_id = entry.get("policy_id") - index_rows.append( - { + prior = index_rows_by_key.get((request_id, guardrail_id)) + if prior is None or (prior["policy_id"] is None and policy_id is not None): + index_rows_by_key[(request_id, guardrail_id)] = { "request_id": request_id, "guardrail_id": guardrail_id, "policy_id": policy_id, "start_time": start_time, } - ) + index_rows: Final = tuple(index_rows_by_key.values()) async with pending.lock: pending_metrics: Final = pending.metrics diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index db6ec754c6e..64fd59bbe44 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -45,7 +45,10 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.model_checks import get_key_models from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler -from litellm.proxy.db.health_check_latest import LatestHealthCheckRow +from litellm.proxy.db.health_check_latest import ( + LatestHealthCheckRow, + query_latest_health_checks, +) from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers from litellm.proxy.health_check import ( ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, @@ -876,7 +879,7 @@ async def _save_background_health_checks_to_db( ) # Step 3: Get latest health checks for all models in one query to compare status - latest_checks: Final = await prisma_client.get_all_latest_health_checks() + latest_checks: Final = await query_latest_health_checks(prisma_client) latest_checks_map: Final = {} for check in latest_checks: # Use model_id as primary key, fallback to model_name diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 8714dd5f3d2..f3542098f95 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -9,6 +9,7 @@ from .max_budget_per_session_limiter import _PROXY_MaxBudgetPerSessionHandler from .max_iterations_limiter import _PROXY_MaxIterationsHandler from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from .prompt_cache_prediction import PromptCacheObserver from .responses_id_security import ResponsesIDSecurity from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler @@ -25,6 +26,7 @@ PROXY_HOOKS: Final = { "max_iterations_limiter": _PROXY_MaxIterationsHandler, "max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler, "sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler, + "prompt_cache_prediction": PromptCacheObserver, } ## FEATURE FLAG HOOKS ## diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c398abff099..a34dc99e472 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -9,10 +9,12 @@ import binascii import logging import os import uuid -from collections.abc import Awaitable, Callable, Mapping, Sequence, Set +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence, Set +from contextlib import asynccontextmanager from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, @@ -23,6 +25,7 @@ from typing import ( TypedDict, ) +from pydantic import TypeAdapter from typing_extensions import NotRequired, ReadOnly from litellm import DualCache @@ -84,6 +87,9 @@ else: InternalUsageCache = Any +_REQUEST_RATE_LIMIT_DATA: Final = TypeAdapter(Mapping[str, object]) + + BATCH_RATE_LIMITER_SCRIPT: Final = """ local results = {} local now = tonumber(ARGV[1]) @@ -2673,12 +2679,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Returns list of descriptors for API key, user, team, team member, end user, model-specific, agent, and agent-session limits. """ - from litellm.proxy.auth.auth_utils import ( - get_team_model_rpm_limit, - get_team_model_tpm_limit, - ) - - descriptors: Final = [] + descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: existing descriptor helpers append in place # API Key rate limits if user_api_key_dict.api_key and ( @@ -2803,34 +2804,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptors=descriptors, ) - if ( - get_team_model_rpm_limit(user_api_key_dict) is not None - or get_team_model_tpm_limit(user_api_key_dict) is not None - ): - _tpm_limit_for_team_model: Final = get_team_model_tpm_limit(user_api_key_dict) or {} - _rpm_limit_for_team_model: Final = get_team_model_rpm_limit(user_api_key_dict) or {} - should_check_rate_limit = False - if requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model: - should_check_rate_limit = True - - if should_check_rate_limit: - model_specific_tpm_limit = None - model_specific_rpm_limit = None - if requested_model in _tpm_limit_for_team_model: - model_specific_tpm_limit = _tpm_limit_for_team_model[requested_model] - if requested_model in _rpm_limit_for_team_model: - model_specific_rpm_limit = _rpm_limit_for_team_model[requested_model] - descriptors.append( - RateLimitDescriptor( - key="model_per_team", - value=f"{user_api_key_dict.team_id}:{requested_model}", - rate_limit={ - "requests_per_unit": model_specific_rpm_limit, - "tokens_per_unit": model_specific_tpm_limit, - "window_size": self.window_size, - }, - ) - ) + self._add_team_model_rate_limit_descriptor_from_metadata( + user_api_key_dict=user_api_key_dict, + requested_model=requested_model if isinstance(requested_model, str) else None, + descriptors=descriptors, + ) # Agent-level and session-level rate limits resolved_agent_id: Final = self._get_resolved_agent_id(user_api_key_dict, data) @@ -3416,6 +3394,108 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): requested_model, ) + async def _build_request_rate_limit_descriptors( + self, + user_api_key_dict: UserAPIKeyAuth, + data: Mapping[str, object], + call_type: str | None, + ) -> list[RateLimitDescriptor]: # mutable-ok: the shared generation reservation helpers require a list + metadata: Final = _REQUEST_RATE_LIMIT_DATA.validate_python( + user_api_key_dict.metadata or MappingProxyType({}) # pyright: ignore[reportUnknownMemberType] # validates the legacy auth metadata boundary + ) + rpm_value: Final = metadata.get("rpm_limit_type") + tpm_value: Final = metadata.get("tpm_limit_type") + rpm_limit_type: Final = rpm_value if isinstance(rpm_value, str) else None + tpm_limit_type: Final = tpm_value if isinstance(tpm_value, str) else None + model_value: Final = data.get("model") + requested_model: Final = model_value if isinstance(model_value, str) else None + model_has_failures: Final = ( + await self._check_model_has_recent_failures( + model=requested_model, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + if requested_model and self._is_dynamic_rate_limiting_enabled(rpm_limit_type, tpm_limit_type) + else False + ) + descriptors: Final = self._create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys + user_api_key_dict=user_api_key_dict, + data=dict(data), # mutable-ok: legacy descriptor helpers accept a request dictionary + rpm_limit_type=rpm_limit_type, + tpm_limit_type=tpm_limit_type, + model_has_failures=model_has_failures, + call_type=call_type, + ) + self._add_project_model_rate_limit_descriptor_from_metadata( + user_api_key_dict=user_api_key_dict, + requested_model=requested_model, + descriptors=descriptors, + ) + self.add_project_io_token_rate_limit_descriptors_from_metadata( + user_api_key_dict=user_api_key_dict, + requested_model=requested_model, + descriptors=descriptors, + ) + return [ # mutable-ok: the shared generation reservation helpers require a list + *descriptors, + *self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model), + ] + + async def _release_request_capacity_when_admitted( + self, + admission: asyncio.Task[RateLimitResponse], + acquisition: ParallelSlotAcquisition, + user_api_key_dict: UserAPIKeyAuth, + ) -> None: + response: Final = await admission + if response["overall_code"] == "OK": + await self._release_parallel_request_slots(acquisition, user_api_key_dict.parent_otel_span) + + @asynccontextmanager + async def request_capacity( + self, + user_api_key_dict: UserAPIKeyAuth, + model: str, + *, + request_data: Mapping[str, object] | None = None, + ) -> AsyncGenerator[None, None]: + """Charge one non-generation provider request to RPM and hold its concurrency slot.""" + data: Final = MappingProxyType({**(request_data or MappingProxyType({})), "model": model}) + descriptors: Final = await self._build_request_rate_limit_descriptors(user_api_key_dict, data, None) + acquisition: Final = ParallelSlotAcquisition( + slot_id=uuid.uuid4().hex, + counter_keys=[ # mutable-ok: the shared slot-release contract requires a list + self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") + for d in descriptors + if d["rate_limit"] is not None and d["rate_limit"].get("max_parallel_requests") is not None + ], + ) + admission: Final = asyncio.create_task( + self.should_rate_limit( + descriptors=descriptors, + parent_otel_span=user_api_key_dict.parent_otel_span, + skip_tpm_check=True, + parallel_slot_id=acquisition["slot_id"], + ) + ) + try: + response: Final = await asyncio.shield(admission) + if response["overall_code"] == "OVER_LIMIT": + self._handle_rate_limit_error(response, descriptors, model) + yield + finally: + cleanup: Final = asyncio.create_task( + self._release_request_capacity_when_admitted(admission, acquisition, user_api_key_dict) + ) + cancellation: asyncio.CancelledError | None = None # rebind-ok: retain cancellation until cleanup finishes + while not cleanup.done(): + try: + await asyncio.shield(cleanup) + except asyncio.CancelledError as exc: + cancellation = exc # rebind-ok: retain the latest cancellation without interrupting slot release + cleanup.result() + if cancellation is not None: + raise cancellation + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -3444,59 +3524,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): call_type=call_type, ) - # Get rate limit types from metadata - metadata: Final = user_api_key_dict.metadata or {} - rpm_limit_type: Final = metadata.get("rpm_limit_type") - tpm_limit_type: Final = metadata.get("tpm_limit_type") - - # For dynamic mode, check if the model has recent failures - model_has_failures = False - requested_model: Final = data.get("model", None) - - if ( - self._is_dynamic_rate_limiting_enabled( - rpm_limit_type=rpm_limit_type, - tpm_limit_type=tpm_limit_type, - ) - and requested_model - ): - model_has_failures = await self._check_model_has_recent_failures( - model=requested_model, - parent_otel_span=user_api_key_dict.parent_otel_span, - ) - - # Create rate limit descriptors - descriptors: Final = self._create_rate_limit_descriptors( + request_data: Final = _REQUEST_RATE_LIMIT_DATA.validate_python(data) + model_value: Final = request_data.get("model") + requested_model: Final = model_value if isinstance(model_value, str) else None + descriptors: Final = await self._build_request_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, - data=data, - rpm_limit_type=rpm_limit_type, - tpm_limit_type=tpm_limit_type, - model_has_failures=model_has_failures, + data=request_data, call_type=call_type, ) - # Add team model rate limits from team_metadata - self._add_team_model_rate_limit_descriptor_from_metadata( - user_api_key_dict=user_api_key_dict, - requested_model=requested_model, - descriptors=descriptors, - ) - - # Project Level Rate Limits - self._add_project_model_rate_limit_descriptor_from_metadata( - user_api_key_dict=user_api_key_dict, - requested_model=requested_model, - descriptors=descriptors, - ) - self.add_project_io_token_rate_limit_descriptors_from_metadata( - user_api_key_dict=user_api_key_dict, - requested_model=requested_model, - descriptors=descriptors, - ) - - # Org Level Rate Limits - descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) - # Only check rate limits if we have descriptors with actual limits if descriptors: # First pass: RPM and max_parallel_requests sliding-window check. diff --git a/litellm/proxy/hooks/prompt_cache_prediction.py b/litellm/proxy/hooks/prompt_cache_prediction.py new file mode 100644 index 00000000000..65c456c5666 --- /dev/null +++ b/litellm/proxy/hooks/prompt_cache_prediction.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +import asyncio +import time +from collections.abc import Callable, Mapping +from datetime import datetime +from typing import TYPE_CHECKING, Final, Literal + +import httpx +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError + +from litellm.caching.dual_cache import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, parse_observed_cache +from litellm.types.utils import ModelResponse + +if TYPE_CHECKING: + from litellm.proxy.utils import InternalUsageCache + +_RETENTION_SECONDS: Final = 86_400 + + +class CacheObservation(BaseModel): + model_config = ConfigDict(extra="forbid", frozen=True, strict=True) + + fingerprint: str = Field(pattern=r"^[0-9a-f]{64}$") + cached_tokens: int = Field(gt=0) + observed_at: float = Field(ge=0, allow_inf_nan=False) + expires_at: float = Field(ge=0, allow_inf_nan=False) + + +_CACHE_ENTRY: Final[TypeAdapter[CacheObservation | str | None]] = TypeAdapter(CacheObservation | str | None) + + +def _cache_key(scope: str, fingerprint: str) -> str: + return f"prompt-cache-observation:{scope}:{fingerprint}" + + +async def lookup( + cache: DualCache, scope: str, prefix: PromptPrefix, now: float | None = None +) -> CacheObservation | None: + checked_at: Final = time.time() if now is None else now + exact: Final = await _read_exact(cache, scope, prefix.fingerprint) + if exact is not None and exact.expires_at > checked_at: + return exact + older: Final = await asyncio.gather( + *(_read_exact(cache, scope, fingerprint) for fingerprint in prefix.fingerprints[1:]) + ) + observations: Final = tuple(observation for observation in (exact, *older) if observation is not None) + return next( + (observation for observation in observations if observation.expires_at > checked_at), + next(iter(observations), None), + ) + + +async def _read_exact(cache: DualCache, scope: str, fingerprint: str) -> CacheObservation | None: + try: + value: Final = _CACHE_ENTRY.validate_python(await cache.async_get_cache(_cache_key(scope, fingerprint), ttl=1)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # validate the legacy cache's untyped result at the I/O boundary + if value is None: + return None + observation: Final = CacheObservation.model_validate_json(value) if isinstance(value, str) else value + except ValidationError: + return None + return observation if observation.fingerprint == fingerprint else None + + +class _Metadata(BaseModel): + model_config = ConfigDict(strict=True) + user_api_key_hash: str = Field(min_length=1) + + +class _Logged(BaseModel): + model_config = ConfigDict(strict=True) + status: Literal["success"] + model_id: str = Field(min_length=1) + metadata: _Metadata + + +class _Event(BaseModel): + model_config = ConfigDict(strict=True, arbitrary_types_allowed=True) + call_type: Literal["anthropic_messages"] + custom_llm_provider: Literal["anthropic"] + cache_hit: bool | None = None + httpx_response: httpx.Response + first_api_call_start_time: datetime + standard_logging_object: _Logged + stream: bool = False + prompt_cache_response_complete: bool = False + + +class PromptCacheObserver(CustomLogger): + def __init__(self, internal_usage_cache: InternalUsageCache, clock: Callable[[], float] = time.time) -> None: + super().__init__() # pyright: ignore[reportUnknownMemberType] # base callback constructor accepts untyped kwargs + self.cache = internal_usage_cache.dual_cache + self.clock = clock + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + if not isinstance(response_obj, ModelResponse): + return + try: + event: Final = _Event.model_validate(kwargs) + wire: Final = event.httpx_response.request + except (ValidationError, RuntimeError, httpx.RequestNotRead): + return + if ( + event.cache_hit + or event.httpx_response.status_code != 200 + or (event.stream and not event.prompt_cache_response_complete) + ): + return + observed: Final = parse_observed_cache( + wire, + response_obj, + event.standard_logging_object.metadata.user_api_key_hash, + event.standard_logging_object.model_id, + ) + if observed is None: + return + prefix: Final = observed.prefix + scope: Final = observed.scope + cache_tokens: Final = observed.cached_tokens + now: Final = self.clock() + started: Final = event.first_api_call_start_time.timestamp() + if started > now: + return + if observed.cache_creation_tokens == 0: + previous: Final = await _read_exact(self.cache, scope, prefix.fingerprint) + if previous is None or previous.fingerprint != prefix.fingerprint or previous.cached_tokens != cache_tokens: + return + observation: Final = CacheObservation( + fingerprint=prefix.fingerprint, + cached_tokens=cache_tokens, + observed_at=now, + expires_at=started + prefix.ttl_seconds, + ) + key: Final = _cache_key(scope, prefix.fingerprint) + payload: Final = observation.model_dump_json() + await self.cache.async_set_cache(key, payload, ttl=_RETENTION_SECONDS) # pyright: ignore[reportUnknownMemberType] # legacy cache accepts a serialized validated observation + if self.cache.redis_cache is not None: + await self.cache.async_set_cache(key, payload, local_only=True, ttl=1) # pyright: ignore[reportUnknownMemberType] # keep the local copy short-lived while Redis retains stale evidence diff --git a/litellm/proxy/list_api/common.py b/litellm/proxy/list_api/common.py index 7ef2827f30e..daa6414fd94 100644 --- a/litellm/proxy/list_api/common.py +++ b/litellm/proxy/list_api/common.py @@ -1,5 +1,6 @@ """Contract machinery shared by every LiteLLM-defined list route, on any surface.""" +from collections.abc import Sequence from typing import Final from urllib.parse import urlencode @@ -7,6 +8,7 @@ from fastapi import Request from fastapi.dependencies.utils import get_flat_params from fastapi.params import ParamTypes from fastapi.responses import JSONResponse +from typing_extensions import ReadOnly, TypedDict from litellm.types.proxy.management_endpoints.management_v1 import ( ListLinks, @@ -56,6 +58,40 @@ def escape_like(value: str) -> str: return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") +class ValidationErrorDetail(TypedDict): + """The keys of a pydantic/FastAPI validation error a problem document needs.""" + + type: ReadOnly[str] + loc: ReadOnly[tuple[int | str, ...]] + msg: ReadOnly[str] + + +def _is_length_error_of_rejected_items(error: ValidationErrorDetail, errors: Sequence[ValidationErrorDetail]) -> bool: + """pydantic counts only items that validated, so a bad item also trips the parent's min_length.""" + return error["type"] == "too_short" and any( + len(other["loc"]) > len(error["loc"]) and other["loc"][: len(error["loc"])] == error["loc"] for other in errors + ) + + +def request_validation_problem(raw_errors: Sequence[ValidationErrorDetail]) -> ProblemDetail: + """A body that fails validation (an unknown field included) is 422; a bad query parameter is 400.""" + errors: Final = tuple(error for error in raw_errors if not _is_length_error_of_rejected_items(error, raw_errors)) + detail: Final = "; ".join(f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in errors) + if any(error["loc"] and error["loc"][0] == "body" for error in errors): + return ProblemDetail( + type=f"{PROBLEM_TYPE_BASE}invalid-request-body", + title="Invalid request body", + status=422, + detail=detail or "The request body is invalid.", + ) + return ProblemDetail( + type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter", + title="Invalid query parameter", + status=400, + detail=detail or "The request query parameters are invalid.", + ) + + def unknown_query_param_problem(unknown: tuple[str, ...], allowed: tuple[str, ...]) -> ProblemDetail: return ProblemDetail( type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter", diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 59971e54e46..563db811edc 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -221,6 +221,8 @@ LITELLM_TRACE_CONTROL_METADATA_FIELDS: Final = frozenset( ) _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = ( + "weights", + "_router_weights", "proxy_server_request", "standard_logging_object", "secret_fields", @@ -334,7 +336,7 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg # and read by spend logs as fact; a client value has no legitimate meaning and no # key or team setting keeps it, so the strip is never gated. _ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset( - {"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY} + {"attempted_fallbacks", "original_model_group", "request_retry_count", CLIENT_OUTPUT_CEILING_METADATA_KEY} ) _ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override" diff --git a/litellm/proxy/management_endpoints/cost_tracking_settings.py b/litellm/proxy/management_endpoints/cost_tracking_settings.py index dc0da63555f..cb376f286ec 100644 --- a/litellm/proxy/management_endpoints/cost_tracking_settings.py +++ b/litellm/proxy/management_endpoints/cost_tracking_settings.py @@ -28,6 +28,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.prompt_cache_prediction import router as prompt_cache_prediction_router from litellm.types.utils import ( CostBreakdown, CostPerToken, @@ -39,6 +40,7 @@ from litellm.types.utils import ( ) router: Final = APIRouter() +router.include_router(prompt_cache_prediction_router) @dataclass(frozen=True, slots=True) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index e3efda507f6..ba7a3309a90 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -567,7 +567,7 @@ async def new_user( teams = check_if_default_team_set() organization_ids: Final = cast(list[str] | None, data_json.pop("organizations", None)) - response: Final = await generate_key_helper_fn(request_type="user", **data_json) + response: Final = await generate_key_helper_fn(request_type="user", **data_json, llm_router=None) # Admin UI Logic # Add User to Team and Organization # if team_id passed add this user to the team diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 95ccb7bbe0b..62019f462a2 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -96,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import ( 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_helpers.access_group_key_sync import ( sync_key_access_group_membership, sync_key_regeneration_access_group_membership, @@ -148,6 +149,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import ( BulkUpdateKeyRequest, BulkUpdateKeyResponse, BulkUpdateTeamKeysRequest, + CustomKeyPolicyRequest, FailedKeyUpdate, KeySearchWhere, SuccessfulKeyUpdate, @@ -201,6 +203,10 @@ class _KeyUpdateResult(TypedDict): data: ReadOnly[Mapping[str, object]] +class _StoredKeyRouterSettings(BaseModel): + router_settings: Mapping[str, object] | None = None + + class _KeyRowWhere(TypedDict): token: ReadOnly[str] @@ -280,6 +286,7 @@ def _config_table(prisma_client: PrismaClient) -> _ConfigTableActions: class _CustomKeyHooksModule(Protocol): user_custom_key_generate: Callable[..., Awaitable[Mapping[str, object]]] | None user_custom_key_update: Callable[..., Awaitable[Mapping[str, object]]] | None + user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None def _custom_key_generate_hook( @@ -294,6 +301,161 @@ def _custom_key_update_hook( return hooks.user_custom_key_update +def _custom_key_policy_hook( + hooks: _CustomKeyHooksModule, +) -> Callable[..., Awaitable[Mapping[str, object]]] | None: + return hooks.user_custom_key_policy + + +async def _enforce_custom_key_update_policy( + hook: Callable[..., Awaitable[Mapping[str, object]]] | None, + data: UpdateKeyRequest, +) -> None: + if hook is None: + return + if not inspect.iscoroutinefunction(hook): + raise ValueError("user_custom_key_update must be a coroutine") + result: Final = await hook(data) + if not result.get("decision", True): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=result.get("message", "Authentication Failed - Custom Auth Rule"), + ) + + +async def _enforce_custom_key_policy( + hook: Callable[..., Awaitable[Mapping[str, object]]] | None, + build_policy_request: Callable[[], CustomKeyPolicyRequest], +) -> None: + if hook is None: + return + if not inspect.iscoroutinefunction(hook): + raise ValueError("user_custom_key_policy must be a coroutine") + result: Final = await hook(build_policy_request()) + if not result.get("decision", True): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=result.get("message", "Authentication Failed - Custom Auth Rule"), + ) + + +_KEY_UPDATE_JSON_STRING_COLUMNS: Final = frozenset({"router_settings", "budget_limits"}) + +_KEY_METADATA_REQUEST_FIELDS: Final = frozenset( + (*LiteLLM_ManagementEndpoint_MetadataFields_Premium, *LiteLLM_ManagementEndpoint_MetadataFields) +) + + +def _decode_json_string_column(column: str, value: object) -> object: + if column in _KEY_UPDATE_JSON_STRING_COLUMNS and isinstance(value, str): + return json.loads(value) + return value + + +def _verification_token_from_row(row: Mapping[str, object]) -> LiteLLM_VerificationToken: + org_id: Final = row["organization_id"] if "organization_id" in row else row.get("org_id") + return LiteLLM_VerificationToken.model_validate(MappingProxyType({**row, "org_id": org_id})) + + +def _effective_key_after_update( + existing_key_row: LiteLLM_VerificationToken, + non_default_values: Mapping[str, object], +) -> LiteLLM_VerificationToken: + overlay: Final = MappingProxyType( + {column: _decode_json_string_column(column, value) for column, value in non_default_values.items()} + ) + return _verification_token_from_row( + MappingProxyType({**existing_key_row.model_dump(), **overlay, "object_permission": None}) + ) + + +def _update_policy_request( + operation: Literal["update", "regenerate"], + existing_key_row: LiteLLM_VerificationToken, + non_default_values: Mapping[str, object], + request: UpdateKeyRequest | RegenerateKeyRequest, +) -> CustomKeyPolicyRequest: + return CustomKeyPolicyRequest( + operation=operation, + existing_key=_verification_token_from_row(existing_key_row.model_dump()), + effective_key=_effective_key_after_update( + existing_key_row=existing_key_row, non_default_values=non_default_values + ), + request=request, + ) + + +def _generate_budget_windows( + budget_limits: Sequence[BudgetLimitEntry] | None, +) -> tuple[Mapping[str, object], ...] | None: + if not budget_limits: + return None + return tuple( + MappingProxyType( + { + **window.model_dump(), + "reset_at": get_budget_reset_time(budget_duration=window.budget_duration).isoformat(), + } + ) + for window in budget_limits + ) + + +def _effective_key_for_generate(data: GenerateKeyRequest, now: datetime) -> LiteLLM_VerificationToken: + requested: Final = data.model_dump(exclude_unset=True, exclude_none=True) + metadata_fields: Final = MappingProxyType( + {field: value for field, value in requested.items() if field in _KEY_METADATA_REQUEST_FIELDS} + ) + column_fields: Final = MappingProxyType( + {field: value for field, value in requested.items() if field not in _KEY_METADATA_REQUEST_FIELDS} + ) + metadata: Final = data.metadata or MappingProxyType({}) + folded_metadata: Final = {**metadata, **metadata_fields} # mutable-ok: encrypt_callback_vars needs a dict + columns: Final = handle_key_type(data, {**column_fields}) # mutable-ok: handle_key_type mutates in place + expires: Final = ( + now + timedelta(seconds=duration_in_seconds(duration=data.duration)) if data.duration is not None else None + ) + budget_reset_at: Final = ( + get_budget_reset_time(budget_duration=data.budget_duration) if data.budget_duration is not None else None + ) + key_rotation_at: Final = ( + now + timedelta(seconds=duration_in_seconds(duration=data.rotation_interval)) + if data.auto_rotate and data.rotation_interval + else None + ) + return _verification_token_from_row( + MappingProxyType( + { + **columns, + "metadata": encrypt_callback_vars(folded_metadata), + "expires": expires, + "budget_reset_at": budget_reset_at, + "key_rotation_at": key_rotation_at, + "budget_limits": _generate_budget_windows(data.budget_limits), + "object_permission": None, + } + ) + ) + + +_EMPTY_DURATION_MEANS_UNCHANGED: Final = frozenset({"duration", "budget_duration"}) + + +def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest) -> UpdateKeyRequest | None: + changed_fields: Final = MappingProxyType( + { + field: value + for field, value in data.model_dump(exclude_unset=True).items() + if field in UpdateKeyRequest.model_fields + and field != "key" + and not (field in _EMPTY_DURATION_MEANS_UNCHANGED and value == "") + } + ) + if not changed_fields: + return None + return UpdateKeyRequest(key=key, **changed_fields) + + class _LegacyDumpable(Protocol): def dict(self) -> Mapping[str, object]: ... @@ -987,6 +1149,7 @@ async def _common_key_generation_helper( litellm_changed_by: str | None, team_table: LiteLLM_TeamTableCachedObj | None, ) -> GenerateKeyResponse: + from litellm.proxy import proxy_server from litellm.proxy.proxy_server import ( litellm_proxy_admin_name, llm_router, @@ -1135,6 +1298,16 @@ async def _common_key_generation_helper( "litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e ) + await _enforce_custom_key_policy( + hook=_custom_key_policy_hook(proxy_server), + build_policy_request=lambda: CustomKeyPolicyRequest( + operation="generate", + existing_key=None, + effective_key=_effective_key_for_generate(data=data, now=datetime.now(timezone.utc)), + request=data, + ), + ) + # TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable _budget_id = data.budget_id if prisma_client is not None and data.soft_budget is not None: @@ -1330,7 +1503,7 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) - response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key") + response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router) response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response @@ -2234,7 +2407,26 @@ async def _update_key_row_with_soft_budget( async def prepare_key_update_data( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, + *, + prisma_client: PrismaClient | None = None, + llm_router: Router | None = None, ): + if data.router_settings is not None or ( + "router_settings" not in data.model_fields_set + and "team_id" in data.model_fields_set + and data.team_id != existing_key_row.team_id + ): + effective_settings: Final = ( + data.router_settings + if data.router_settings is not None + else _StoredKeyRouterSettings.model_validate(existing_key_row, from_attributes=True).router_settings + ) + await validate_router_settings_weights( + effective_settings, + team_id=data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id, + prisma_client=prisma_client, + llm_router=llm_router, + ) data_json: Final[dict] = data.model_dump(exclude_unset=True) data_json.pop("key", None) data_json.pop("new_key", None) @@ -2301,12 +2493,6 @@ async def prepare_key_update_data( # sentinel for Json? columns, so store the JSON literal null non_default_values["budget_limits"] = json.dumps(None) - if "object_permission" in non_default_values: - non_default_values = await _handle_update_object_permission( - data_json=non_default_values, - existing_key_row=existing_key_row, - ) - _metadata: Final = existing_key_row.metadata or {} # validate model_max_budget @@ -2327,13 +2513,12 @@ async def prepare_key_update_data( async def _handle_update_object_permission( data_json: dict, existing_key_row: LiteLLM_VerificationToken, + prisma_client: PrismaClient, ) -> dict: - """ - Handle the update of object permission. - """ - from litellm.proxy.proxy_server import prisma_client + """Persist the requested object permission row and swap it for its id, only after the key policy allowed the write.""" + if "object_permission" not in data_json: + return data_json - # Use the common helper to handle the object permission update object_permission_id: Final = await handle_update_object_permission_common( data_json=data_json, existing_object_permission_id=existing_key_row.object_permission_id, @@ -2467,6 +2652,7 @@ async def _process_single_key_update( llm_router: Router | None, user_custom_key_update: Callable | None = None, existing_key_row: LiteLLM_VerificationToken | None = None, + user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None = None, ) -> dict[str, object]: """ Process a single key update with all validations and checks. @@ -2575,7 +2761,19 @@ async def _process_single_key_update( ) # Prepare update data - non_default_values = await prepare_key_update_data(data=update_key_request, existing_key_row=existing_key_row) + non_default_values = await prepare_key_update_data( + data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router + ) + + await _enforce_custom_key_policy( + hook=user_custom_key_policy, + build_policy_request=lambda: _update_policy_request( + operation="update", + existing_key_row=existing_key_row, + non_default_values=non_default_values, + request=update_key_request, + ), + ) # Update key in database if prisma_client is None: @@ -2584,7 +2782,12 @@ async def _process_single_key_update( detail={"error": "Database not connected"}, ) - _data: Final = {**non_default_values, "token": update_key_request.key} + update_values: Final = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) + _data: Final = {**update_values, "token": update_key_request.key} response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict "Mapping[str, object] | None", await prisma_client.update_data(token=update_key_request.key, data=_data), @@ -3077,23 +3280,13 @@ async def update_key_fn( user_api_key_cache=user_api_key_cache, ) - # Custom key update hook - custom_key_update_hook: Final[Callable[..., Awaitable[Mapping[str, object]]] | None] = _custom_key_update_hook( - proxy_server - ) - if custom_key_update_hook is not None: - if inspect.iscoroutinefunction(custom_key_update_hook): - result: Final = await custom_key_update_hook(data) - else: - raise ValueError("user_custom_key_update must be a coroutine") - decision: Final = result.get("decision", True) - message: Final = result.get("message", "Authentication Failed - Custom Auth Rule") - if not decision: - raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message) + await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data) # Enforce upperbound key params on update (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) - non_default_values: Final = await prepare_key_update_data(data=data, existing_key_row=existing_key_row) + non_default_values: Final = await prepare_key_update_data( + data=data, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router + ) # Only validate key_alias format if it's actually being changed new_key_alias: Final = non_default_values.get("key_alias", None) @@ -3114,21 +3307,36 @@ async def update_key_fn( existing_key_alias=existing_key_row.key_alias, ) + await _enforce_custom_key_policy( + hook=_custom_key_policy_hook(proxy_server), + build_policy_request=lambda: _update_policy_request( + operation="update", + existing_key_row=existing_key_row, + non_default_values=non_default_values, + request=data, + ), + ) + if prisma_client is None: raise Exception("Not connected to DB!") + update_values: Final = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name response: Final = ( await _update_key_row_with_soft_budget( prisma_client=prisma_client, key=key, data=data, - non_default_values=non_default_values, + non_default_values=update_values, existing_key_row=existing_key_row, changed_by=changed_by, ) if "soft_budget" in data.model_fields_set - else await prisma_client.update_data(token=key, data=MappingProxyType({**non_default_values, "token": key})) + else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) ) # Delete - key from cache, since it's been updated! @@ -3263,6 +3471,7 @@ async def bulk_update_keys( ) custom_key_update_hook: Final = _custom_key_update_hook(proxy_server) + custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server) if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: raise HTTPException( @@ -3310,6 +3519,7 @@ async def bulk_update_keys( proxy_logging_obj=proxy_logging_obj, llm_router=llm_router, user_custom_key_update=custom_key_update_hook, + user_custom_key_policy=custom_key_policy_hook, ) successful_updates.append( @@ -3427,6 +3637,7 @@ async def bulk_update_team_keys( ) custom_key_update_hook: Final = _custom_key_update_hook(proxy_server) + custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server) if prisma_client is None: raise HTTPException( @@ -3557,6 +3768,7 @@ async def bulk_update_team_keys( proxy_logging_obj=proxy_logging_obj, llm_router=llm_router, user_custom_key_update=custom_key_update_hook, + user_custom_key_policy=custom_key_policy_hook, existing_key_row=existing_by_token[db_token], ) @@ -4082,6 +4294,40 @@ def _check_model_access_group(models: list[str] | None, llm_router: Router | Non return True +_NO_METADATA: Final[Mapping[str, object]] = MappingProxyType({}) + + +def metadata_json_with_limits( + metadata: Mapping[str, object] | None, + *, + model_rpm_limit: Mapping[str, object] | None, + model_tpm_limit: Mapping[str, object] | None, + mcp_rpm_limit: Mapping[str, int] | None, + tag_rpm_limit: Mapping[str, int] | None, + guardrails: Sequence[str] | None, + policies: Sequence[str] | None, + prompts: Sequence[str] | None, +) -> str: + """Serialize the stored metadata blob with the per-model, MCP, tag, guardrail, policy and prompt settings folded in.""" + limits: Final = tuple( + (name, value) + for name, value in ( + ("model_rpm_limit", model_rpm_limit), + ("model_tpm_limit", model_tpm_limit), + ("mcp_rpm_limit", mcp_rpm_limit), + ("tag_rpm_limit", tag_rpm_limit), + ("guardrails", guardrails), + ("policies", policies), + ("prompts", prompts), + ) + if value is not None + ) + if metadata is None and not limits: + return json.dumps(None) + merged: Final = {**(metadata or _NO_METADATA), **dict(limits)} # mutable-ok: encrypt_callback_vars takes a dict + return json.dumps(encrypt_callback_vars(merged)) + + async def generate_key_helper_fn( request_type: Literal["user", "key"], # identifies if this request is from /user/new or /key/generate duration: str | None = None, @@ -4137,15 +4383,24 @@ async def generate_key_helper_fn( object_permission: LiteLLM_ObjectPermissionBase | None = None, auto_rotate: bool | None = None, rotation_interval: str | None = None, - router_settings: dict | None = None, + router_settings: dict[str, object] | None = None, access_group_ids: list[str] | None = None, budget_limits: list | None = None, # multiple concurrent budget windows + *, + llm_router: Router | None = None, ): from litellm.proxy.proxy_server import premium_user, prisma_client if prisma_client is None: raise Exception("Connect Proxy to database to generate keys - https://docs.litellm.ai/docs/proxy/virtual_keys ") + await validate_router_settings_weights( + router_settings, + team_id=team_id, + prisma_client=prisma_client, + llm_router=llm_router, + ) + if token is None: if key is not None: token = key @@ -4184,31 +4439,16 @@ async def generate_key_helper_fn( permissions_json: Final = json.dumps(permissions) router_settings_json: Final = safe_dumps(router_settings) if router_settings is not None else safe_dumps({}) - # Add model_rpm_limit and model_tpm_limit to metadata - if model_rpm_limit is not None: - metadata = metadata or {} - metadata["model_rpm_limit"] = model_rpm_limit - if model_tpm_limit is not None: - metadata = metadata or {} - metadata["model_tpm_limit"] = model_tpm_limit - if mcp_rpm_limit is not None: - metadata = metadata or {} - metadata["mcp_rpm_limit"] = mcp_rpm_limit - if tag_rpm_limit is not None: - metadata = metadata or {} - metadata["tag_rpm_limit"] = tag_rpm_limit - if guardrails is not None: - metadata = metadata or {} - metadata["guardrails"] = guardrails - if policies is not None: - metadata = metadata or {} - metadata["policies"] = policies - if prompts is not None: - metadata = metadata or {} - metadata["prompts"] = prompts - - metadata = encrypt_callback_vars(metadata) - metadata_json: Final = json.dumps(metadata) + metadata_json: Final = metadata_json_with_limits( + metadata, + model_rpm_limit=model_rpm_limit, + model_tpm_limit=model_tpm_limit, + mcp_rpm_limit=mcp_rpm_limit, + tag_rpm_limit=tag_rpm_limit, + guardrails=guardrails, + policies=policies, + prompts=prompts, + ) validate_model_max_budget(model_max_budget) model_max_budget_json: Final = json.dumps(model_max_budget) budget_fallbacks_json: Final = json.dumps(budget_fallbacks or {}) @@ -5070,6 +5310,7 @@ async def _insert_deprecated_key( async def _execute_virtual_key_regeneration( *, prisma_client: PrismaClient, + llm_router: Router | None = None, key_in_db: LiteLLM_VerificationToken, hashed_api_key: str, key: str, @@ -5080,6 +5321,7 @@ async def _execute_virtual_key_regeneration( proxy_logging_obj: ProxyLogging, ) -> GenerateKeyResponse: """Generate new token, update DB, invalidate cache, and return response.""" + from litellm.proxy import proxy_server from litellm.proxy.proxy_server import hash_token # Mirror the /key/update ownership rebind guard. See helper docstring. @@ -5127,15 +5369,34 @@ async def _execute_virtual_key_regeneration( non_default_values = {} if data is not None: + update_request: Final = _regenerate_request_as_update_request(key=hashed_api_key, data=data) + if update_request is not None: + await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=update_request) # Enforce upperbound key params on regenerate (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) - non_default_values = await prepare_key_update_data(data=data, existing_key_row=key_in_db) + non_default_values = await prepare_key_update_data( + data=data, existing_key_row=key_in_db, prisma_client=prisma_client, llm_router=llm_router + ) # Only validate key_alias format if it's actually being changed new_key_alias: Final = non_default_values.get("key_alias") if new_key_alias != key_in_db.key_alias: _validate_key_alias_format(key_alias=new_key_alias) verbose_proxy_logger.debug("non_default_values: %s", non_default_values) - update_data.update(non_default_values) + await _enforce_custom_key_policy( + hook=_custom_key_policy_hook(proxy_server), + build_policy_request=lambda: _update_policy_request( + operation="regenerate", + existing_key_row=key_in_db, + non_default_values=non_default_values, + request=data if data is not None else RegenerateKeyRequest(), + ), + ) + update_values: Final = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=key_in_db, + prisma_client=prisma_client, + ) + update_data.update(update_values) jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) # Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash, @@ -5145,6 +5406,13 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) + await _persist_deleted_verification_tokens( + keys=[key_in_db], + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=litellm_changed_by, + ) + # If grace period set, insert deprecated key so old key remains valid await _insert_deprecated_key( prisma_client=prisma_client, @@ -5268,6 +5536,7 @@ async def regenerate_key_fn( try: from litellm.proxy.proxy_server import ( hash_token, + llm_router, master_key, premium_user, prisma_client, @@ -5443,19 +5712,9 @@ async def regenerate_key_fn( if litellm_changed_by is not None and not isinstance(litellm_changed_by, str): litellm_changed_by = None - # Save the old key record to deleted table before regeneration. - # This preserves key_alias and team_id metadata for historical spend records. - # If this fails, abort the regeneration to avoid permanently losing the - # old hash→metadata mapping. - await _persist_deleted_verification_tokens( - keys=[_key_in_db], - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=litellm_changed_by, - ) - return await _execute_virtual_key_regeneration( prisma_client=prisma_client, + llm_router=llm_router, key_in_db=_key_in_db, hashed_api_key=hashed_api_key, key=key, diff --git a/litellm/proxy/management_endpoints/management_v1/users.py b/litellm/proxy/management_endpoints/management_v1/users.py index fa2a5cf3c3d..afe4482c9da 100644 --- a/litellm/proxy/management_endpoints/management_v1/users.py +++ b/litellm/proxy/management_endpoints/management_v1/users.py @@ -1,4 +1,4 @@ -"""`POST /management/v1/users/bulk_delete`.""" +"""`POST /management/v1/users/bulk` and `POST /management/v1/users/bulk_delete`.""" from typing import Annotated, Final @@ -9,19 +9,105 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem, reject_unknown_query_params from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX +from litellm.proxy.management_helpers.bulk_user_creation import bulk_create_users from litellm.proxy.management_helpers.bulk_user_deletion import bulk_delete_users from litellm.proxy.management_helpers.utils import ( - management_endpoint_wrapper, # pyright: ignore[reportUnknownVariableType] # legacy decorator is untyped + management_endpoint_wrapper, # pyright: ignore[reportUnknownVariableType] # legacy untyped decorator ) from litellm.types.proxy.management_endpoints.internal_user_endpoints import ( BulkDeleteUserRequest, BulkDeleteUsersResponse, + BulkNewUserRequest, + BulkNewUserResponse, ) from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail router: Final = APIRouter(prefix=MANAGEMENT_V1_PREFIX) +@router.post( + "/users/bulk", + tags=["Internal User management"], # mutable-ok: fastapi types tags as list[str | Enum] + dependencies=(Depends(user_api_key_auth),), + response_model=BulkNewUserResponse, +) +@management_endpoint_wrapper +async def bulk_create_users_route( + data: BulkNewUserRequest, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> BulkNewUserResponse: + """ + Create up to 500 internal users in one request, optionally adding each one to teams. + + Every entry in `users` takes the same fields as `/user/new`, with two differences: `auto_create_key` + defaults to `false` (opt in per user to also get a virtual key back) and `send_invite_email` is not + supported. Unknown fields are rejected with 422. Rows are validated together (duplicate ids or emails, + unknown teams, roles the caller may not grant), inserted in one statement, and each referenced team is + written once for all of its new members. + + Rows fail independently: a bad row is reported in `data` with `success: false` and an `error`, and the + other rows still get created. A user that was created but could not be added to one of its teams is + reported with `success: true`, `teams` listing where they did land, and `error` naming the failed team. + The whole request is refused with a 403 problem document only if creating the valid rows would exceed + the license seat limit. + + Example curl: + ``` + curl -X POST "http://localhost:4000/management/v1/users/bulk" \\ + -H "Content-Type: application/json" \\ + -H "Authorization: Bearer sk-1234" \\ + -d '{ + "users": [ + {"user_email": "a@example.com", "user_role": "internal_user", "teams": ["team-1"]}, + {"user_email": "b@example.com", "user_role": "internal_user", "auto_create_key": true} + ] + }' + ``` + + Returns `data` (one entry per input row, in order, with `user_id`, `user_email`, `success`, `teams`, + `key`, `error`) and `meta` with `total_requested`, `created` and `failed`. + """ + try: + from litellm.proxy.proxy_server import ( + _license_check, # pyright: ignore[reportPrivateUsage] # same proxy license singleton /user/new reads + litellm_proxy_admin_name, + prisma_client, + user_api_key_cache, + ) + + if prisma_client is None: + raise ManagementProblem( + ProblemDetail( + type=f"{PROBLEM_TYPE_BASE}database-not-connected", + title="Database not connected", + status=503, + detail=CommonProxyErrors.db_not_connected_error.value, + ) + ) + + return await bulk_create_users( + users=data.users, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + license_check=_license_check, + litellm_proxy_admin_name=litellm_proxy_admin_name, + user_api_key_cache=user_api_key_cache, + ) + + except ManagementProblem: + raise + except Exception: # noqa: BLE001 # a driver error answers as a problem document, not the OpenAI error shape + verbose_proxy_logger.exception("/management/v1/users/bulk: Exception occurred") + raise ManagementProblem( + ProblemDetail( + type=f"{PROBLEM_TYPE_BASE}internal-server-error", + title="Internal server error", + status=500, + detail="Failed to create users.", + ) + ) + + @router.post( "/users/bulk_delete", tags=["Internal User management"], # mutable-ok: FastAPI types `tags` as list[str], not Sequence diff --git a/litellm/proxy/management_endpoints/prompt_cache_prediction.py b/litellm/proxy/management_endpoints/prompt_cache_prediction.py new file mode 100644 index 00000000000..56e844214d6 --- /dev/null +++ b/litellm/proxy/management_endpoints/prompt_cache_prediction.py @@ -0,0 +1,278 @@ +import time +from collections.abc import Mapping +from types import MappingProxyType +from typing import Annotated, Final + +from fastapi import APIRouter, Depends, HTTPException, Request +from pydantic import BaseModel, JsonValue, TypeAdapter + +import litellm +from litellm._internal_context import current_billing_time, pinned_billing_time +from litellm.caching.caching import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.anthropic.prompt_cache_prediction import ( + PromptPrefix, + TokenCounter, + UnsupportedPredictionTarget, + cache_scope, + count_prompt_tokens, + parse_prompt, + resolve_prediction_target, + supported_prediction_headers, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.proxy.auth.auth_utils import get_cache_prediction_deployments +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.http_parsing_utils import ( + _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary +) +from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner +) +from litellm.proxy.hooks.prompt_cache_prediction import lookup +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.types.management_endpoints.prompt_cache_prediction import ( + CacheCostScenario, + CacheEvidence, + CachePredictionArm, + CachePredictionRequest, + CachePredictionResponse, + CacheTokenBuckets, +) +from litellm.types.router import Deployment +from litellm.utils import get_prompt_cache_min_tokens + +router: Final = APIRouter() +_REQUEST_DATA: Final = TypeAdapter(Mapping[str, object]) + + +class _CallerSettings(BaseModel): + config: Mapping[str, object] | None = None + + +def has_request_transforms() -> bool: + from litellm.proxy.hooks import PROXY_HOOKS + + builtins: Final = frozenset(PROXY_HOOKS.values()) + hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook") + callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger) + return any( + type(callback) not in builtins + and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks) + for callback in callbacks + ) + + +def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets: + return CacheTokenBuckets( + uncached_input_tokens=suffix_tokens, + cache_read_input_tokens=read_tokens, + cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0, + cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0, + ) + + +def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None: + cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens) + return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None + + +def _capacity_counter( + limiter: _PROXY_MaxParallelRequestsHandler_v3, + caller: UserAPIKeyAuth, + model_name: str, + request_data: Mapping[str, object], +) -> TokenCounter: + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + async with limiter.request_capacity(caller, model_name, request_data=request_data): + return await count_prompt_tokens(model, api_key, body) + + return count + + +def _capacity_request_data( + http_request: Request, caller: UserAPIKeyAuth, request_data: Mapping[str, object] +) -> Mapping[str, object]: + # The parsed-body cache retains only original top-level keys. Replay the + # shared idempotent tag merges on limiter-only data when auth added metadata. + data: Final = dict(request_data) # mutable-ok: the existing tag merge owners accept a dictionary out-param + LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(http_request, data, caller) # pyright: ignore[reportUnknownMemberType] # legacy tag owner takes the validated capacity dictionary + LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(data, caller) # pyright: ignore[reportUnknownMemberType] # legacy tag owner merges trusted key tags into capacity metadata + return MappingProxyType(data) + + +async def predict_arm( + deployment: Deployment, + body: Mapping[str, JsonValue], + prefix: PromptPrefix, + caller_key_hash: str, + cache: DualCache, + token_counter: TokenCounter, +) -> CachePredictionArm: + deployment_id: Final = deployment.model_info.id or "" + params: Final = deployment.litellm_params + unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model) + if deployment.model_info.blocked: + return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"})) + target: Final = resolve_prediction_target(params) + if isinstance(target, UnsupportedPredictionTarget): + return unknown.model_copy(update=MappingProxyType({"reason": target.reason})) + model: Final = target.model + api_key: Final = target.api_key + total_count: Final = await token_counter(model, api_key, body) + prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body) + if total_count is None or prefix_count is None or total_count < prefix_count: + return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"})) + scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model) + observation: Final = await lookup(cache, scope, prefix) + exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint + cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count + if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable): + return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"})) + suffix: Final = total_count - cacheable + evidence: Final = ( + CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at) + if observation is not None + else None + ) + if cacheable < get_prompt_cache_min_tokens(params.model): + disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count)) + if disabled is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="disabled", + reason="below_cache_minimum", + estimate=disabled, + cold=disabled, + warm=disabled, + token_count_source="anthropic_count_tokens", + ) + fresh: Final = observation is not None and observation.expires_at > time.time() + read: Final = observation.cached_tokens if fresh and observation is not None else 0 + with pinned_billing_time(current_billing_time()): + cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds)) + warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds)) + estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds)) + if cold is None or warm is None or estimate is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown", + reason=None if fresh else "observation_expired" if observation else "no_compatible_observation", + estimate=estimate, + cold=cold, + warm=warm, + evidence=evidence, + token_count_source="anthropic_count_tokens", + ) + + +@router.post( + "/cost/predict-cache", + tags=["Cost Tracking"], # mutable-ok: FastAPI requires a list for OpenAPI tags + response_model=CachePredictionResponse, +) +async def predict_cache_cost( + request: CachePredictionRequest, + http_request: Request, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> CachePredictionResponse: + """Compare the next native Anthropic request on two configured deployment IDs. + + Estimates use provider token counting and recent successful cache telemetry for this key. + Unknown cache state uses the cold scenario when prices/counts are available. Cache observations + do not guarantee retention. v0 supports one message-content breakpoint, text and client tools; + system/tool-only breakpoints, thinking, images, nondefault Anthropic versions, beta headers and + request transforms are unknown. + Each provider count consumes one RPM unit and holds concurrency capacity; a comparison uses + up to four counts. The legacy rate limiter returns unknown without contacting the provider. + This endpoint does not generate tokens, prewarm caches, choose a model or alter routing. + """ + from litellm.proxy.proxy_server import llm_router, proxy_logging_obj + + if llm_router is None: + raise HTTPException(status_code=503, detail="Model router is unavailable") + deployments: Final = get_cache_prediction_deployments( + current_deployment_id=request.current_deployment_id, + candidate_deployment_id=request.candidate_deployment_id, + llm_router=llm_router, + team_id=user_api_key_dict.team_id, + ) + if deployments is None: + raise HTTPException(status_code=404, detail="Deployment not found") + current, candidate = deployments + for deployment in (current, candidate): + await can_key_call_resolved_model( + model=deployment.model_name, + llm_model_list=llm_router.get_model_list(), + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + prefix: Final = parse_prompt(request.request) + caller: Final = user_api_key_dict.api_key + caller_settings: Final = _CallerSettings.model_validate(user_api_key_dict, from_attributes=True) + unsupported_transform: Final = bool(caller_settings.config) or has_request_transforms() + unsupported_headers: Final = not supported_prediction_headers(http_request.headers) + limiter: Final = proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if ( + prefix is None + or not caller + or unsupported_transform + or unsupported_headers + or not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3) + ): + reason: Final = ( + "unsupported_provider_headers" + if unsupported_headers + else "unsupported_request_transform" + if unsupported_transform + else "unsupported_prompt_shape" + if prefix is None + else "caller_identity_unavailable" + if not caller + else "limiter_unavailable" + ) + return CachePredictionResponse( + stay=CachePredictionArm(deployment_id=request.current_deployment_id, reason=reason), + switch=CachePredictionArm(deployment_id=request.candidate_deployment_id, reason=reason), + switch_delta=None, + cache_rebuild_penalty=None, + ) + request_data: Final = _capacity_request_data( + http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await _read_request_body(http_request)) + ) + stay: Final = await predict_arm( + current, + request.request, + prefix, + caller, + proxy_logging_obj.internal_usage_cache.dual_cache, + _capacity_counter(limiter, user_api_key_dict, current.model_name, request_data), + ) + switch: Final = ( + stay + if current.model_info.id == candidate.model_info.id + else await predict_arm( + candidate, + request.request, + prefix, + caller, + proxy_logging_obj.internal_usage_cache.dual_cache, + _capacity_counter(limiter, user_api_key_dict, candidate.model_name, request_data), + ) + ) + return CachePredictionResponse( + stay=stay, + switch=switch, + switch_delta=(switch.estimate.input_cost - stay.estimate.input_cost) + if switch.estimate is not None and stay.estimate is not None + else None, + cache_rebuild_penalty=(switch.estimate.input_cost - switch.warm.input_cost) + if switch.estimate is not None and switch.warm is not None + else None, + ) diff --git a/litellm/proxy/management_endpoints/router_weights.py b/litellm/proxy/management_endpoints/router_weights.py new file mode 100644 index 00000000000..99b368808c3 --- /dev/null +++ b/litellm/proxy/management_endpoints/router_weights.py @@ -0,0 +1,129 @@ +from abc import abstractmethod +from collections.abc import Mapping +from typing import Annotated, Final, Protocol + +from fastapi import HTTPException +from pydantic import BaseModel, BeforeValidator, ValidationError + +from litellm.repositories.prisma_protocols import TableActions +from litellm.types.router_weights import RouterWeights + + +class _StoredModel(Protocol): + @property + @abstractmethod + def model_id(self) -> str: + pass + + +class _ModelDb(Protocol): + @property + @abstractmethod + def litellm_proxymodeltable(self) -> TableActions[_StoredModel]: + pass + + +class _PrismaClient(Protocol): + @property + @abstractmethod + def db(self) -> _ModelDb: + pass + + +class _Router(Protocol): + @abstractmethod + def get_deployment(self, model_id: str) -> object | None: + pass + + +class _RouterWeightSettings(BaseModel): + weights: RouterWeights | None = None + + +class _RouterWeightModelInfo(BaseModel): + team_id: str | None = None + db_model: bool | None = None + team_public_model_name: str | None = None + + +def _router_weight_model_info(value: object) -> _RouterWeightModelInfo: + if isinstance(value, str): + return _RouterWeightModelInfo.model_validate_json(value) + return _RouterWeightModelInfo.model_validate(value or {}, from_attributes=True) + + +class _RouterWeightDeployment(BaseModel): + model_name: str + model_info: Annotated[_RouterWeightModelInfo, BeforeValidator(_router_weight_model_info)] + + +def _validate_router_weight_reference( + model_group: str, + deployment_id: str, + team_id: str | None, + stored: _RouterWeightDeployment | None, + configured: object | None, +) -> None: + reference: Final = ( + stored + if stored is not None + else ( + _RouterWeightDeployment.model_validate(configured, from_attributes=True) if configured is not None else None + ) + ) + if ( + reference is None + or (stored is None and reference.model_info.db_model) + or (reference.model_info.team_id is not None and reference.model_info.team_id != team_id) + ): + raise HTTPException(status_code=400, detail=f"Unknown deployment ID in router weights: {deployment_id}") + canonical_group: Final = ( + reference.model_info.team_public_model_name if reference.model_info.team_id is not None else None + ) or reference.model_name + if model_group != canonical_group: + raise HTTPException( + status_code=400, + detail=f"Deployment {deployment_id} does not belong to model group {model_group}", + ) + + +async def validate_router_settings_weights( + router_settings: BaseModel | Mapping[str, object] | None, + *, + team_id: str | None, + prisma_client: _PrismaClient | None, + llm_router: _Router | None, +) -> None: + try: + weights: Final = ( + _RouterWeightSettings.model_validate(router_settings, from_attributes=True).weights + if router_settings is not None + else None + ) + except ValidationError: + raise HTTPException( + status_code=400, + detail="Invalid router weights. Replace or clear router_settings.weights.", + ) from None + if not weights: + return + deployment_ids: Final = frozenset(deployment_id for group in weights.values() for deployment_id in group) + if not deployment_ids: + return + if prisma_client is None: + raise HTTPException(status_code=503, detail="Database unavailable while validating router weights") + stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many( + where={"model_id": {"in": list(deployment_ids)}} + ) + stored_by_id: Final = { + row.model_id: _RouterWeightDeployment.model_validate(row, from_attributes=True) for row in stored_models + } + for model_group, group_weights in weights.items(): + for deployment_id in group_weights: + _validate_router_weight_reference( + model_group, + deployment_id, + team_id, + stored_by_id.get(deployment_id), + llm_router.get_deployment(model_id=deployment_id) if llm_router is not None else None, + ) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 6b16692f7ad..2d04a4d1e04 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -112,6 +112,7 @@ from litellm.proxy.management_endpoints.common_utils import ( from litellm.proxy.management_endpoints.organization_endpoints import ( add_member_to_organization, ) +from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights from litellm.proxy.management_endpoints.tag_management_endpoints import ( get_daily_activity, ) @@ -431,27 +432,26 @@ async def _refresh_cached_team( ) +async def _can_manage_team( + team_obj: LiteLLM_TeamTable, + user_api_key_dict: UserAPIKeyAuth, +) -> bool: + """True for a proxy admin, an admin of this team, or an org admin for the team's organization.""" + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return True + + if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): + return True + + return await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj) + + async def _verify_team_access( team_obj: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth, ) -> None: - """ - Verify the caller is authorized to manage the given team. - - Access is granted if: - - Caller is a proxy admin, OR - - Caller is an org admin for the team's organization, OR - - Caller is a team admin of this team - - Raises HTTPException(403) otherwise. - """ - if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: - return - - if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): - return - - if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj): + """Raise HTTPException(403) unless the caller can manage the given team.""" + if await _can_manage_team(team_obj=team_obj, user_api_key_dict=user_api_key_dict): return raise HTTPException( @@ -1289,6 +1289,7 @@ async def new_team( create_audit_log_for_update, general_settings, litellm_proxy_admin_name, + llm_router, prisma_client, user_api_key_cache, ) @@ -1463,6 +1464,13 @@ async def new_team( user_api_key_dict=user_api_key_dict, ) + await validate_router_settings_weights( + data.router_settings, + team_id=data.team_id, + prisma_client=prisma_client, + llm_router=llm_router, + ) + ## ADD TO MODEL TABLE _model_id = None if data.model_aliases is not None and isinstance(data.model_aliases, dict): @@ -2076,6 +2084,13 @@ async def update_team( user_api_key_dict=user_api_key_dict, ) + await validate_router_settings_weights( + data.router_settings, + team_id=data.team_id, + prisma_client=prisma_client, + llm_router=llm_router, + ) + _existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None) enforce_output_token_estimates_are_admin_only( data=data, @@ -4370,6 +4385,20 @@ async def _hydrate_member_user_details( return tuple(hydrate(m) for m in members) +class _OrganizationModelsRow(BaseModel): + models: list[str] = [] # mutable-ok: pydantic field default + + +class _TeamRowWithOrganization(BaseModel): + litellm_organization_table: _OrganizationModelsRow | None = None + + +def _parent_organization_models(team_row: BaseModel) -> list[str] | None: + """Return the parent org's model allow-list, or None when the team has no org.""" + organization: Final = _TeamRowWithOrganization.model_validate(team_row.model_dump()).litellm_organization_table + return organization.models if organization is not None else None + + async def _resolve_team_access_group_resources( _team_info: TeamInfoResponseObjectTeamTable, ) -> TeamInfoResponseObjectTeamTable: @@ -4441,7 +4470,11 @@ async def team_info( try: team_info: BaseModel | None = await _team_db(prisma_client).find_unique( where={"team_id": team_id}, - include={"litellm_model_table": True, "object_permission": True}, + include={ + "litellm_model_table": True, + "object_permission": True, + "litellm_organization_table": True, + }, ) if team_info is None: raise Exception @@ -4450,9 +4483,12 @@ async def team_info( status_code=status.HTTP_404_NOT_FOUND, detail={"message": f"Team not found, passed team id: {team_id}."}, ) - await validate_membership( - user_api_key_dict=user_api_key_dict, - team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()), + team_table: Final = LiteLLM_TeamTable.model_validate(team_info.model_dump()) + await validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_table) + organization_models: Final[list[str] | None] = ( + _parent_organization_models(team_info) + if await _can_manage_team(team_obj=team_table, user_api_key_dict=user_api_key_dict) + else None ) ## GET ALL KEYS ## @@ -4512,7 +4548,10 @@ async def team_info( members=resolved_team_info.members_with_roles, ) hydrated_team_info: Final = resolved_team_info.model_copy( - update={"members_with_roles": hydrated_members} # mutable-ok: pydantic update payload + update={ # mutable-ok: pydantic update payload + "members_with_roles": hydrated_members, + "organization_models": organization_models, + } ) response_object: Final = TeamInfoResponseObject( diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 091dccf1433..329443148a2 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -3592,6 +3592,7 @@ class SSOAuthenticationHandler: verbose_proxy_logger.info("user_defined_values for creating ui key: %s", user_defined_values) response: Final = await generate_key_helper_fn( + llm_router=None, request_type="key", duration=LITELLM_UI_SESSION_DURATION, key_max_budget=litellm.max_ui_session_budget, diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py new file mode 100644 index 00000000000..56abe3b6a3f --- /dev/null +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -0,0 +1,871 @@ +"""Batched internal user creation behind `POST /management/v1/users/bulk`. + +The batch is validated with set queries, user rows land in one `create_many`, and every +referenced team is written once under its advisory lock instead of once per user. +""" + +import asyncio +import json +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeAlias, TypeVar + +from fastapi import HTTPException, Request +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing_extensions import ReadOnly, TypedDict + +from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid +from litellm.integrations.prometheus import PrometheusLogger +from litellm.proxy._types import ( + LiteLLM_TeamTable, + LitellmUserRoles, + Member, + NewUserRequestTeam, + OrganizationMemberAddRequest, + OrgMember, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state +from litellm.proxy.auth.litellm_license import LicenseCheck +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time +from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks +from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem +from litellm.proxy.management_endpoints.common_utils import ( + _is_user_org_admin_for_team, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses + _is_user_team_admin, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses + validate_budget_duration, +) +from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_internal_new_user_params, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # /user/new defaults; result validated below + check_if_default_team_set, +) +from litellm.proxy.management_endpoints.key_management_endpoints import ( + _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses + generate_key_helper_fn, # pyright: ignore[reportUnknownVariableType] # legacy untyped helper; result validated by _KEY_RESPONSE + metadata_json_with_limits, +) +from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add +from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL +from litellm.proxy.management_helpers.object_permission_utils import ( + _set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below +) +from litellm.proxy.management_helpers.utils import ( + _resolve_member_budget_id, # pyright: ignore[reportPrivateUsage] # shared with /team/member_add +) +from litellm.proxy.utils import PrismaClient +from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.types.proxy.management_endpoints.internal_user_endpoints import ( + BulkNewUserItem, + BulkNewUserMeta, + BulkNewUserResponse, + UserCreateResult, +) +from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail + +if TYPE_CHECKING: + from prisma import Prisma + from prisma import models as prisma_models + + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + +BULK_NEW_USER_CONCURRENCY: Final = 10 + +TeamRole: TypeAlias = Literal["user", "admin"] +KeyGenerator: TypeAlias = Callable[..., Awaitable[object]] +_T: Final = TypeVar("_T") + + +@dataclass(frozen=True, slots=True) +class _RowFailure: + index: int + user_id: str | None + user_email: str | None + error: str + + +@dataclass(frozen=True, slots=True) +class _PendingUser: + index: int + request: BulkNewUserItem + user_id: str + teams: tuple[NewUserRequestTeam, ...] + + +class _UserRow(BaseModel): + """The `/user/new` body after defaults and object permission were applied.""" + + model_config = ConfigDict(extra="ignore") + + user_id: str + user_email: str | None = None + user_alias: str | None = None + user_role: str | None = None + team_id: str | None = None + max_budget: float | None = None + spend: float | None = 0.0 + models: tuple[str, ...] | None = None + metadata: Mapping[str, object] | None = None + max_parallel_requests: int | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + budget_duration: str | None = None + allowed_cache_controls: tuple[str, ...] | None = None + sso_user_id: str | None = None + object_permission_id: str | None = None + model_max_budget: Mapping[str, object] | None = None + model_rpm_limit: Mapping[str, object] | None = None + model_tpm_limit: Mapping[str, object] | None = None + mcp_rpm_limit: Mapping[str, int] | None = None + tag_rpm_limit: Mapping[str, int] | None = None + guardrails: tuple[str, ...] | None = None + policies: tuple[str, ...] | None = None + prompts: tuple[str, ...] | None = None + duration: str | None = None + key_alias: str | None = None + aliases: Mapping[str, object] | None = None + config: Mapping[str, object] | None = None + permissions: Mapping[str, object] | None = None + blocked: bool | None = None + agent_id: str | None = None + budget_fallbacks: Mapping[str, tuple[str, ...]] | None = None + budget_limits: tuple[Mapping[str, object], ...] | None = None + organizations: tuple[str, ...] | None = None + + +_USER_ROW: Final = TypeAdapter(_UserRow) + + +@dataclass(frozen=True, slots=True) +class _PreparedUser: + pending: _PendingUser + row: _UserRow + + +@dataclass(frozen=True, slots=True) +class _TeamAssignment: + user_id: str + user_email: str | None + role: TeamRole + max_budget_in_team: float | None + + +@dataclass(frozen=True, slots=True) +class _TeamWrite: + """Outcome of one locked roster write. `failed` maps user ids to the reason they were not added.""" + + team_id: str + after: tuple[Member, ...] + added: frozenset[str] + failed: Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class _CreatedUser: + prepared: _PreparedUser + teams: tuple[str, ...] + key: str | None + errors: tuple[str, ...] + + +_ERROR_DETAIL: Final = TypeAdapter(Mapping[str, object]) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object]) + + +class _KeyResponse(BaseModel): + token: str + + +_KEY_RESPONSE: Final = TypeAdapter(_KeyResponse) + + +def _error_message(exc: BaseException) -> str: + if not isinstance(exc, HTTPException): + return str(exc) + try: + detail: Final = _ERROR_DETAIL.validate_python(exc.detail) + except ValidationError: + return str(exc.detail) + return str(detail.get("error", detail)) + + +def _requested_teams(item: BulkNewUserItem) -> tuple[NewUserRequestTeam, ...]: + if item.team_id is not None: + return (NewUserRequestTeam(team_id=item.team_id),) + teams: Final = item.teams if item.teams is not None else check_if_default_team_set() + if teams is None: + return () + return tuple(team if isinstance(team, NewUserRequestTeam) else NewUserRequestTeam(team_id=team) for team in teams) + + +def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str | None: + if ( + item.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN + ): + return ( + "Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). " + f"Attempted to create user with role: {item.user_role}. Your role: {user_api_key_dict.user_role}" + ) + try: + validate_budget_duration(item.budget_duration) + _check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict) + except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only + return _error_message(exc) + return None + + +def _normalized_email(email: str | None) -> str | None: + return email.strip().lower() if email else None + + +def _partition_rows( + users: Sequence[BulkNewUserItem], user_api_key_dict: UserAPIKeyAuth +) -> tuple[tuple[_PendingUser, ...], tuple[_RowFailure, ...]]: + """Assign ids, run the per-row checks and fail later rows that repeat an earlier row's id or email.""" + user_ids: Final = tuple(item.user_id or str(uuid.uuid4()) for item in users) + first_index_by_id: Final = MappingProxyType( + {user_id: index for index, user_id in reversed(tuple(enumerate(user_ids)))} + ) + first_index_by_email: Final = MappingProxyType( + { + email: index + for index, email in reversed(tuple(enumerate(_normalized_email(item.user_email) for item in users))) + if email is not None + } + ) + + def classify(index: int, item: BulkNewUserItem) -> _PendingUser | _RowFailure: + user_id: Final = user_ids[index] + email: Final = _normalized_email(item.user_email) + if first_index_by_id[user_id] != index: + return _RowFailure(index, user_id, item.user_email, f"Duplicate user_id in request: {user_id}") + if email is not None and first_index_by_email[email] != index: + return _RowFailure(index, user_id, item.user_email, f"Duplicate user_email in request: {item.user_email}") + error: Final = _row_error(item, user_api_key_dict) + if error is not None: + return _RowFailure(index, user_id, item.user_email, error) + return _PendingUser(index, item, user_id, _requested_teams(item)) + + outcomes: Final = tuple(classify(index, item) for index, item in enumerate(users)) + return ( + tuple(outcome for outcome in outcomes if isinstance(outcome, _PendingUser)), + tuple(outcome for outcome in outcomes if isinstance(outcome, _RowFailure)), + ) + + +def _user_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_UserTable]": + return UserRepository(prisma_client).table + + +async def _existing_user_conflicts( + prisma_client: PrismaClient, pending: Sequence[_PendingUser] +) -> tuple[frozenset[str], frozenset[str]]: + """Return the requested user ids and (lowercased) emails that already exist, using one query each.""" + user_ids: Final = sorted(user.user_id for user in pending) + emails: Final = sorted(frozenset(user.request.user_email for user in pending if user.request.user_email)) + if not user_ids: + return frozenset(), frozenset() + table: Final = _user_table(prisma_client) + id_filter: Final = {"user_id": {"in": user_ids}} # mutable-ok: Prisma query filters are dict-shaped + email_filter: Final = {"user_email": {"in": emails, "mode": "insensitive"}} # mutable-ok: Prisma filter + id_rows: Final = await table.find_many(where=id_filter) + email_rows: Final = await table.find_many(where=email_filter) if emails else () + return ( + frozenset(row.user_id for row in id_rows), + frozenset(lowered for row in email_rows if (lowered := _normalized_email(row.user_email)) is not None), + ) + + +async def _load_teams(prisma_client: PrismaClient, team_ids: frozenset[str]) -> Mapping[str, LiteLLM_TeamTable]: + if not team_ids: + return MappingProxyType({}) + rows: Final = await TeamRepository(prisma_client).table.find_many( + where={"team_id": {"in": sorted(team_ids)}} # mutable-ok: Prisma query filters are dict-shaped + ) + return MappingProxyType({row.team_id: LiteLLM_TeamTable.model_validate(row.model_dump()) for row in rows}) + + +async def _team_permission_error(team: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth) -> str | None: + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: + return None + if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team): + return None + if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team): + return None + return f"Call not allowed. User not proxy admin OR team admin. team_id={team.team_id}" + + +async def _unusable_teams( + prisma_client: PrismaClient, + pending: Sequence[_PendingUser], + user_api_key_dict: UserAPIKeyAuth, +) -> tuple[Mapping[str, LiteLLM_TeamTable], Mapping[str, str]]: + """Load every referenced team once and explain, per team id, why rows naming it cannot proceed.""" + team_ids: Final = frozenset(team.team_id for user in pending for team in user.teams) + teams: Final = await _load_teams(prisma_client, team_ids) + permission_errors: Final = await asyncio.gather( + *(_team_permission_error(team, user_api_key_dict) for team in teams.values()) + ) + missing: Final = tuple( + (team_id, f"Team id={team_id} does not exist") for team_id in team_ids if team_id not in teams + ) + denied: Final = tuple( + (team.team_id, error) + for team, error in zip(teams.values(), permission_errors, strict=True) + if error is not None + ) + return teams, MappingProxyType({team_id: error for team_id, error in (*missing, *denied)}) + + +def _db_failure( + user: _PendingUser, + existing_ids: frozenset[str], + existing_emails: frozenset[str], + team_errors: Mapping[str, str], +) -> _RowFailure | None: + email: Final = _normalized_email(user.request.user_email) + if user.user_id in existing_ids: + return _RowFailure(user.index, user.user_id, user.request.user_email, f"User id={user.user_id} already exists") + if email is not None and email in existing_emails: + return _RowFailure( + user.index, user.user_id, user.request.user_email, f"User email={user.request.user_email} already exists" + ) + errors: Final = tuple(team_errors[team.team_id] for team in user.teams if team.team_id in team_errors) + if errors: + return _RowFailure(user.index, user.user_id, user.request.user_email, "; ".join(errors)) + return None + + +async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _PreparedUser | _RowFailure: + try: + dumped: Final = user.request.model_dump(exclude={"user_id"}) # mutable-ok: pydantic IncEx takes a set + data: Final = {**dumped, "user_id": user.user_id} # mutable-ok: /user/new defaults helper mutates in place + data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request)) + with_permission: Final = _JSON_OBJECT.validate_python( + await _set_object_permission(data_json=data_json, prisma_client=prisma_client) # pyright: ignore[reportUnknownArgumentType] # validated by the adapter + ) + return _PreparedUser(user, _USER_ROW.validate_python(with_permission)) + except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only + verbose_proxy_logger.warning("/user/bulk_new: could not prepare row %d - %s", user.index, type(exc).__name__) + return _RowFailure(user.index, user.user_id, user.request.user_email, _error_message(exc)) + + +class _UserCreateData(TypedDict): + """One `LiteLLM_UserTable` row as `create_many` takes it; JSON columns are pre-serialized.""" + + user_id: ReadOnly[str] + user_email: ReadOnly[str | None] + user_alias: ReadOnly[str | None] + user_role: ReadOnly[str | None] + team_id: ReadOnly[str | None] + max_budget: ReadOnly[float | None] + spend: ReadOnly[float] + models: ReadOnly[tuple[str, ...]] + metadata: ReadOnly[str] + max_parallel_requests: ReadOnly[int | None] + tpm_limit: ReadOnly[int | None] + rpm_limit: ReadOnly[int | None] + budget_duration: ReadOnly[str | None] + budget_reset_at: ReadOnly[datetime | None] + allowed_cache_controls: ReadOnly[tuple[str, ...]] + sso_user_id: ReadOnly[str | None] + object_permission_id: ReadOnly[str | None] + teams: ReadOnly[tuple[str, ...]] + model_max_budget: ReadOnly[str] + + +def _user_create_payload(prepared: _PreparedUser) -> _UserCreateData: + row: Final = prepared.row + metadata_json: Final = metadata_json_with_limits( + row.metadata, + model_rpm_limit=row.model_rpm_limit, + model_tpm_limit=row.model_tpm_limit, + mcp_rpm_limit=row.mcp_rpm_limit, + tag_rpm_limit=row.tag_rpm_limit, + guardrails=row.guardrails, + policies=row.policies, + prompts=row.prompts, + ) + payload: Final[_UserCreateData] = { + "user_id": row.user_id, + "user_email": row.user_email, + "user_alias": row.user_alias, + "user_role": row.user_role, + "team_id": row.team_id, + "max_budget": row.max_budget, + "spend": row.spend or 0.0, + "models": row.models or (), + "metadata": metadata_json, + "max_parallel_requests": row.max_parallel_requests, + "tpm_limit": row.tpm_limit, + "rpm_limit": row.rpm_limit, + "budget_duration": row.budget_duration, + "budget_reset_at": get_budget_reset_time(row.budget_duration) if row.budget_duration else None, + "allowed_cache_controls": row.allowed_cache_controls or (), + "sso_user_id": row.sso_user_id, + "object_permission_id": row.object_permission_id, + "teams": tuple(team.team_id for team in prepared.pending.teams), + "model_max_budget": json.dumps(row.model_max_budget) if row.model_max_budget else "{}", + } + return payload + + +async def _bounded(limit: int, awaitables: Sequence[Awaitable[_T]]) -> tuple[_T | BaseException, ...]: + semaphore: Final = asyncio.Semaphore(limit) + + async def run(awaitable: Awaitable[_T]) -> _T: + async with semaphore: + return await awaitable + + return tuple(await asyncio.gather(*(run(awaitable) for awaitable in awaitables), return_exceptions=True)) + + +async def _insert_users( + prisma_client: PrismaClient, prepared: Sequence[_PreparedUser] +) -> tuple[tuple[_PreparedUser, ...], tuple[_RowFailure, ...]]: + """Insert every row in one statement. If that fails, retry rows one at a time so the error lands on its row.""" + if not prepared: + return (), () + table: Final = _user_table(prisma_client) + payloads: Final = tuple(_user_create_payload(user) for user in prepared) + try: + await table.create_many(data=payloads) + return tuple(prepared), () + except Exception as exc: # noqa: BLE001 # fall back to per-row inserts so the failing row can be identified + verbose_proxy_logger.warning("/user/bulk_new: create_many failed, retrying rows individually", exc_info=True) + outcome_unknown: Final = PrismaDBExceptionHandler.is_database_infrastructure_error(exc) + requested: Final = frozenset(payload["user_id"] for payload in payloads) + landed_rows: Final = await table.find_many(where={"user_id": {"in": list(requested)}}) # mutable-ok: Prisma filter + landed: Final = frozenset(row.user_id for row in landed_rows) + # create_many is one INSERT: after a lost response the full set is ours, any partial set belongs to another request + if outcome_unknown and landed == requested: + return tuple(prepared), () + taken: Final = tuple(user for user in prepared if user.row.user_id in landed) + retried: Final = tuple(user for user in prepared if user.row.user_id not in landed) + outcomes: Final = await _bounded( + BULK_NEW_USER_CONCURRENCY, tuple(table.create(data=_user_create_payload(user)) for user in retried) + ) + failed: Final = MappingProxyType( + { + **{ + user.row.user_id: _RowFailure( + user.pending.index, + user.pending.user_id, + user.row.user_email, + f"User id={user.row.user_id} already exists", + ) + for user in taken + }, + **{ + user.row.user_id: _RowFailure( + user.pending.index, user.pending.user_id, user.row.user_email, _error_message(outcome) + ) + for user, outcome in zip(retried, outcomes, strict=True) + if isinstance(outcome, BaseException) + }, + } + ) + return ( + tuple(user for user in prepared if user.row.user_id not in failed), + tuple(failed.values()), + ) + + +def _assignments_by_team(created: Sequence[_PreparedUser]) -> Mapping[str, tuple[_TeamAssignment, ...]]: + team_ids: Final = tuple(dict.fromkeys(team.team_id for user in created for team in user.pending.teams)) + return MappingProxyType( + { + team_id: tuple( + _TeamAssignment(user.pending.user_id, user.row.user_email, team.user_role, team.max_budget_in_team) + for user in created + for team in user.pending.teams + if team.team_id == team_id + ) + for team_id in team_ids + } + ) + + +class _MembershipData(TypedDict): + team_id: ReadOnly[str] + user_id: ReadOnly[str] + budget_id: ReadOnly[str | None] + + +class _RosterData(TypedDict): + members_with_roles: ReadOnly[str] + + +class _TeamsData(TypedDict): + teams: ReadOnly[tuple[str, ...]] + + +def _default_member_budget_id(team: LiteLLM_TeamTable) -> str | None: + metadata: Final = ( + _JSON_OBJECT.validate_python( + team.metadata # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter + ) + if team.metadata # pyright: ignore[reportUnknownMemberType] # same bare dict + else None + ) + budget_id: Final = metadata.get("team_member_budget_id") if metadata is not None else None + return budget_id if isinstance(budget_id, str) else None + + +def _team_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamTable]": + return tx.litellm_teamtable # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do + + +def _membership_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamMembership]": + return tx.litellm_teammembership # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do + + +async def _write_team_roster( + prisma_client: PrismaClient, + team: LiteLLM_TeamTable, + members: Sequence[_TeamAssignment], + user_api_key_dict: UserAPIKeyAuth, + litellm_proxy_admin_name: str, +) -> _TeamWrite: + """Add every new member to one team under its advisory lock: one roster rewrite and one membership insert.""" + try: + async with prisma_client.tx() as tx: + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team.team_id) + roster: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team.team_id) + if roster is None: + raise ValueError(f"Team id={team.team_id} does not exist") + already_present: Final = frozenset(member.user_id for member in roster if member.user_id) + new_members: Final = tuple(member for member in members if member.user_id not in already_present) + budget_ids: Final = tuple( + [ # mutable-ok: budgets are created one at a time on the transaction's single connection + await _resolve_member_budget_id( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + max_budget_in_team=member.max_budget_in_team, + allowed_models=team.default_team_member_models or None, + budget_duration=None, + default_team_budget_id=_default_member_budget_id(team), + tx=tx, # pyright: ignore[reportArgumentType] # MemberWriteTx lags the generated Prisma signatures, same as /team/member_add + ) + for member in new_members + ] + ) + await _membership_tx_db(tx).create_many( + data=tuple( + _MembershipData(team_id=team.team_id, user_id=member.user_id, budget_id=budget_id) + for member, budget_id in zip(new_members, budget_ids, strict=True) + ), + skip_duplicates=True, + ) + after: Final = ( + *roster, + *(Member(user_id=m.user_id, user_email=m.user_email, role=m.role) for m in new_members), + ) + await _team_tx_db(tx).update( + where={"team_id": team.team_id}, # mutable-ok: Prisma query filters are dict-shaped + data=_RosterData(members_with_roles=json.dumps(tuple(member.model_dump() for member in after))), + ) + return _TeamWrite( + team_id=team.team_id, + after=after, + added=frozenset(member.user_id for member in members), + failed=MappingProxyType({}), + ) + except Exception as exc: # noqa: BLE001 # the team write failure is reported on each affected row + verbose_proxy_logger.exception("/user/bulk_new: failed to add %d members to a team", len(members)) + message: Final = f"Failed to add user to team {team.team_id}: {_error_message(exc)}" + return _TeamWrite( + team_id=team.team_id, + after=(), + added=frozenset(), + failed=MappingProxyType({member.user_id: message for member in members}), + ) + + +async def _detach_failed_teams( + prisma_client: PrismaClient, created: Sequence[_PreparedUser], writes: Mapping[str, _TeamWrite] +) -> None: + """Users are inserted with `teams` already set; drop the teams whose roster write did not take them.""" + table: Final = _user_table(prisma_client) + updates: Final = tuple( + table.update( + where={"user_id": user.row.user_id}, # mutable-ok: Prisma query filters are dict-shaped + data=_TeamsData(teams=landed), + ) + for user in created + if (landed := _row_teams(user, writes)[0]) != tuple(team.team_id for team in user.pending.teams) + ) + for outcome in await _bounded(BULK_NEW_USER_CONCURRENCY, updates): + if isinstance(outcome, BaseException): + verbose_proxy_logger.warning( + "/user/bulk_new: could not detach failed teams from user - %s", type(outcome).__name__ + ) + + +async def _publish_team_writes(writes: Sequence[_TeamWrite], user_api_key_cache: "UserApiKeyCache") -> None: + prometheus_logger: Final = PrometheusLogger.get_instance() + for write in writes: + if prometheus_logger is None or not write.added: + continue + try: + prometheus_logger.set_team_members_metric( + LiteLLM_TeamTable( + team_id=write.team_id, + members_with_roles=write.after, # pyright: ignore[reportArgumentType] # pydantic coerces the tuple into the declared list + ) + ) + except Exception: # noqa: BLE001 # metrics are best-effort and must not fail the request + verbose_proxy_logger.debug("Prometheus: failed to emit team members metric", exc_info=True) + evictions: Final = await _bounded( + BULK_NEW_USER_CONCURRENCY, + tuple( + invalidate_team_member_spend_state( + user_id=user_id, team_id=write.team_id, user_api_key_cache=user_api_key_cache + ) + for write in writes + for user_id in write.added + ), + ) + for eviction in evictions: + if isinstance(eviction, BaseException): + verbose_proxy_logger.warning("/user/bulk_new: cache eviction failed - %s", type(eviction).__name__) + + +_KEY_FIELDS: Final = MappingProxyType( + { + name: True + for name in ( + "user_id", + "team_id", + "agent_id", + "duration", + "key_alias", + "models", + "aliases", + "config", + "permissions", + "blocked", + "spend", + "budget_fallbacks", + "budget_limits", + "metadata", + "max_parallel_requests", + "tpm_limit", + "rpm_limit", + "allowed_cache_controls", + "model_max_budget", + "model_rpm_limit", + "model_tpm_limit", + "mcp_rpm_limit", + "tag_rpm_limit", + "guardrails", + "policies", + "prompts", + "object_permission_id", + ) + } +) + + +async def _generate_key(prepared: _PreparedUser, generate_key: KeyGenerator) -> str: + response: Final = _KEY_RESPONSE.validate_python( + await generate_key( + request_type="key", table_name="key", **prepared.row.model_dump(include=_KEY_FIELDS, exclude_none=True) + ) + ) + return response.token + + +async def _add_to_organizations( + prepared: _PreparedUser, organizations: Sequence[str], user_api_key_dict: UserAPIKeyAuth +) -> None: + for organization_id in organizations: + await organization_member_add( + data=OrganizationMemberAddRequest( + organization_id=organization_id, + member=OrgMember(user_id=prepared.row.user_id, role=LitellmUserRoles.INTERNAL_USER), + ), + http_request=Request(scope={"type": "http", "path": "/user/bulk_new"}), # mutable-ok: ASGI scopes are dicts + user_api_key_dict=user_api_key_dict, + ) + + +async def _run_per_user( + created: Sequence[_PreparedUser], + select: Callable[[_PreparedUser], bool], + action: Callable[[_PreparedUser], Awaitable[_T]], +) -> Mapping[str, _T | BaseException]: + chosen: Final = tuple(user for user in created if select(user)) + outcomes: Final = await _bounded(BULK_NEW_USER_CONCURRENCY, tuple(action(user) for user in chosen)) + return MappingProxyType({user.row.user_id: outcome for user, outcome in zip(chosen, outcomes, strict=True)}) + + +async def _write_audit_logs( + prisma_client: PrismaClient, + created: Sequence[_PreparedUser], + user_api_key_dict: UserAPIKeyAuth, + litellm_proxy_admin_name: str, +) -> None: + if not created: + return + created_ids: Final = sorted(user.row.user_id for user in created) + created_filter: Final = {"user_id": {"in": created_ids}} # mutable-ok: Prisma query filters are dict-shaped + rows: Final = await _user_table(prisma_client).find_many(where=created_filter) + outcomes: Final = await _bounded( + BULK_NEW_USER_CONCURRENCY, + tuple( + UserManagementEventHooks.create_internal_user_audit_log( + user_id=row.user_id, + action="created", + litellm_changed_by=user_api_key_dict.user_id, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + before_value=None, + after_value=row.model_dump_json(exclude_none=True), + ) + for row in rows + ), + ) + for outcome in outcomes: + if isinstance(outcome, BaseException): + verbose_proxy_logger.warning( + "Unable to create audit log for user on `/user/bulk_new` - %s", type(outcome).__name__ + ) + + +def _row_teams(prepared: _PreparedUser, writes: Mapping[str, _TeamWrite]) -> tuple[tuple[str, ...], tuple[str, ...]]: + """Split a user's requested teams into the ones they landed in and the errors for the ones they did not.""" + requested: Final = tuple(team.team_id for team in prepared.pending.teams) + return ( + tuple(team_id for team_id in requested if prepared.row.user_id in writes[team_id].added), + tuple( + writes[team_id].failed[prepared.row.user_id] + for team_id in requested + if prepared.row.user_id in writes[team_id].failed + ), + ) + + +def _to_result(created: _CreatedUser) -> UserCreateResult: + return UserCreateResult( + user_id=created.prepared.row.user_id, + user_email=created.prepared.row.user_email, + success=True, + teams=created.teams, + key=created.key, + error="; ".join(created.errors) if created.errors else None, + ) + + +def _failure_result(failure: _RowFailure) -> UserCreateResult: + return UserCreateResult(user_id=failure.user_id, user_email=failure.user_email, success=False, error=failure.error) + + +async def bulk_create_users( + users: Sequence[BulkNewUserItem], + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, + license_check: LicenseCheck, + litellm_proxy_admin_name: str, + user_api_key_cache: "UserApiKeyCache", + generate_key: KeyGenerator = generate_key_helper_fn, +) -> BulkNewUserResponse: + """Create every valid row in `users`; rows that fail validation or a write are reported, not raised. + + Raises a 403 `ManagementProblem` only when the whole batch would push the deployment over its license seat + limit. + """ + pending, request_failures = _partition_rows(users, user_api_key_dict) + existing_ids, existing_emails = await _existing_user_conflicts(prisma_client, pending) + teams, team_errors = await _unusable_teams(prisma_client, pending, user_api_key_dict) + db_failures: Final = tuple( + failure + for user in pending + if (failure := _db_failure(user, existing_ids, existing_emails, team_errors)) is not None + ) + failed_indexes: Final = frozenset(failure.index for failure in db_failures) + creatable: Final = tuple(user for user in pending if user.index not in failed_indexes) + + billable_users: Final = await UserRepository(prisma_client).count_billable_users() + if creatable and license_check.is_over_limit(total_users=billable_users + len(creatable)): + raise ManagementProblem( + ProblemDetail( + type=f"{PROBLEM_TYPE_BASE}license-limit-exceeded", + title="License limit exceeded", + status=403, + detail="License is over limit. Please contact support@berri.ai to upgrade your license.", + ) + ) + + prepared_outcomes: Final = tuple([await _prepare_user(user, prisma_client) for user in creatable]) + prepare_failures: Final = tuple(o for o in prepared_outcomes if isinstance(o, _RowFailure)) + created, insert_failures = await _insert_users( + prisma_client, tuple(o for o in prepared_outcomes if isinstance(o, _PreparedUser)) + ) + + team_writes: Final = MappingProxyType( + { + team_id: await _write_team_roster( + prisma_client, teams[team_id], members, user_api_key_dict, litellm_proxy_admin_name + ) + for team_id, members in _assignments_by_team(created).items() + } + ) + await _detach_failed_teams(prisma_client, created, team_writes) + await _publish_team_writes(tuple(team_writes.values()), user_api_key_cache) + + keys: Final = await _run_per_user( + created, lambda user: user.pending.request.auto_create_key, lambda user: _generate_key(user, generate_key) + ) + org_outcomes: Final = await _run_per_user( + created, + lambda user: bool(user.row.organizations), + lambda user: _add_to_organizations(user, user.row.organizations or (), user_api_key_dict), + ) + await _write_audit_logs(prisma_client, created, user_api_key_dict, litellm_proxy_admin_name) + + def finish(prepared: _PreparedUser) -> _CreatedUser: + landed, team_failures = _row_teams(prepared, team_writes) + key_outcome: Final = keys.get(prepared.row.user_id) + org_outcome: Final = org_outcomes.get(prepared.row.user_id) + return _CreatedUser( + prepared=prepared, + teams=landed, + key=key_outcome if isinstance(key_outcome, str) else None, + errors=( + *team_failures, + *( + (f"Failed to create key: {_error_message(key_outcome)}",) + if isinstance(key_outcome, BaseException) + else () + ), + *( + (f"Failed to add user to organizations: {_error_message(org_outcome)}",) + if isinstance(org_outcome, BaseException) + else () + ), + ), + ) + + failures: Final = MappingProxyType( + { + failure.index: _failure_result(failure) + for failure in (*request_failures, *db_failures, *prepare_failures, *insert_failures) + } + ) + successes_by_index: Final = MappingProxyType({user.pending.index: _to_result(finish(user)) for user in created}) + results: Final = tuple( + failures[index] if index in failures else successes_by_index[index] for index in range(len(users)) + ) + successes: Final = sum(1 for result in results if result.success) + return BulkNewUserResponse( + data=results, + meta=BulkNewUserMeta(total_requested=len(users), created=successes, failed=len(users) - successes), + ) diff --git a/litellm/proxy/management_helpers/team_metadata_validation.py b/litellm/proxy/management_helpers/team_metadata_validation.py index 7bc66c240c7..76477ab2988 100644 --- a/litellm/proxy/management_helpers/team_metadata_validation.py +++ b/litellm/proxy/management_helpers/team_metadata_validation.py @@ -115,15 +115,15 @@ async def run_team_metadata_validation( "error": f"custom_team_metadata_validate is an Enterprise feature. {CommonProxyErrors.not_premium_user.value}" }, ) - if not ( - inspect.iscoroutinefunction(validator) or inspect.iscoroutinefunction(getattr(validator, "__call__", None)) - ): - raise HTTPException( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={ # mutable-ok: HTTPException.detail has no immutable form - "error": "custom_team_metadata_validate must be an async function" - }, - ) + if not inspect.iscoroutinefunction(validator): + validator_call: Final = getattr(validator, "__call__", None) # noqa: B004 # value unwrap for the functor check + if not inspect.iscoroutinefunction(validator_call): + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail={ # mutable-ok: HTTPException.detail has no immutable form + "error": "custom_team_metadata_validate must be an async function" + }, + ) try: raw_result: Final = await asyncio.wait_for(validator(payload), timeout=timeout_seconds) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py index e26f5f57532..0ff02b29c58 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py @@ -5,18 +5,105 @@ Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough e Supports different modalities: text, audio, video, and web search. """ +from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Any, Final +from itertools import chain, pairwise +from types import MappingProxyType +from typing import Final, Literal, TypeAlias from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.vertex_ai.gemini.grounding_requests import GroundingRequests, calculate_grounding_requests from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import ( BasePassthroughLoggingHandler, ) from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import ( PassThroughEndpointLoggingTypedDict, ) -from litellm.types.utils import LlmProviders, ModelResponse, Usage -from litellm.utils import get_model_info +from litellm.types.utils import ( + CompletionTokensDetailsWrapper, + CostBreakdown, + LlmProviders, + ModelResponse, + PromptTokensDetailsWrapper, + Usage, +) + +_NO_GROUNDING: Final = GroundingRequests(web_search_requests=None, google_maps_grounding_requests=None) + +_AGGREGATED_FIELDS: Final = frozenset( + { + "promptTokenCount", + "candidatesTokenCount", + "totalTokenCount", + "toolUsePromptTokenCount", + "promptTokensDetails", + "candidatesTokensDetails", + } +) + + +def _detail_entries(raw: object) -> tuple[Mapping[str, object], ...]: + """Narrow one turn's ``*TokensDetails`` value to the entries that are actually shaped like one.""" + return tuple(entry for entry in raw if isinstance(entry, Mapping)) if isinstance(raw, Sequence) else () + + +def _grounding_metadata(websocket_messages: Sequence[object]) -> tuple[Mapping[str, object], ...]: + """Collect every ``serverContent.groundingMetadata`` a session emitted. + + Live reports grounding in the server frames, never in ``usageMetadata``, so the per-query + charge has to be counted here rather than derived from the token totals. + """ + return tuple( + metadata + for message in websocket_messages + if isinstance(message, Mapping) + for server_content in (message.get("serverContent"),) + if isinstance(server_content, Mapping) + for metadata in (server_content.get("groundingMetadata"),) + if isinstance(metadata, Mapping) + ) + + +def _turns(websocket_messages: Sequence[object]) -> tuple[tuple[object, ...], ...]: + """Split a session at every ``usageMetadata`` frame; frames after the last one never got their usage.""" + closes: Final = tuple( + index + 1 + for index, message in enumerate(websocket_messages) + if isinstance(message, Mapping) and isinstance(message.get("usageMetadata"), dict) + ) + return tuple(tuple(websocket_messages[start:end]) for start, end in pairwise((0, *closes))) + + +def _session_grounding_requests(websocket_messages: Sequence[object]) -> GroundingRequests: + per_turn: Final = tuple( + calculate_grounding_requests(_grounding_metadata(turn)) for turn in _turns(websocket_messages) + ) + web_search_requests: Final = sum(requests.web_search_requests or 0 for requests in per_turn) + google_maps_grounding_requests: Final = sum(requests.google_maps_grounding_requests or 0 for requests in per_turn) + return GroundingRequests( + web_search_requests=web_search_requests or None, + google_maps_grounding_requests=google_maps_grounding_requests or None, + ) + + +_SummedField: TypeAlias = Literal[ + "input_cost", + "output_cost", + "tool_usage_cost", + "cache_read_cost", + "cache_creation_cost", + "reasoning_cost", + "original_cost", + "discount_amount", + "margin_fixed_amount", + "margin_total_amount", +] + + +def _summed(breakdowns: Sequence[CostBreakdown], field: _SummedField) -> float | None: + values: Final = tuple(value for breakdown in breakdowns if (value := breakdown.get(field)) is not None) + return sum(values) if values else None class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): @@ -48,186 +135,110 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): """Return the LLM provider name.""" return LlmProviders.VERTEX_AI + @staticmethod + def _resolve_detail_counts( + details: Sequence[Mapping[str, object]], + declared_total: object, + ) -> tuple[tuple[str, int], ...]: + """ + Pair each of one turn's ``*TokensDetails`` entries with its token count. + + Live sometimes names the modality that carries the rest of a turn without a + ``tokenCount``, and reading the absent key as zero drops those tokens from the + breakdown, so real audio ends up priced as text. A lone unpriced entry therefore takes + whatever the turn's declared count leaves over. Two or more cannot be told apart, so + they are left out and the cost calculator charges the remainder as text. + """ + priced: Final = tuple( + (str(detail.get("modality", "TEXT")), count) + for detail in details + if isinstance(count := detail.get("tokenCount"), int) + ) + unpriced: Final = tuple( + str(detail.get("modality", "TEXT")) for detail in details if not isinstance(detail.get("tokenCount"), int) + ) + if len(unpriced) != 1 or not isinstance(declared_total, int): + return priced + residual: Final = declared_total - sum(count for _, count in priced) + return priced if residual <= 0 else (*priced, (unpriced[0], residual)) + + @staticmethod + def _sum_by_modality(counts: Sequence[tuple[str, int]]) -> Mapping[str, int]: + """Total the (modality, tokenCount) pairs of one or more turns per modality.""" + return MappingProxyType({modality: sum(c for m, c in counts if m == modality) for modality, _ in counts}) + + @staticmethod + def _merged_modality_totals( + snapshots: Sequence[Mapping[str, object]], + count_key: str, + details_key: str, + ) -> Mapping[str, int]: + """Total every turn's per-modality counts, so the breakdown adds up the way the totals do.""" + return VertexAILivePassthroughLoggingHandler._sum_by_modality( + tuple( + chain.from_iterable( + VertexAILivePassthroughLoggingHandler._resolve_detail_counts( + _detail_entries(snapshot.get(details_key)), snapshot.get(count_key) + ) + for snapshot in snapshots + ) + ) + ) + @staticmethod def _extract_usage_metadata_from_websocket_messages( - websocket_messages: list[dict], + websocket_messages: Sequence[object], ) -> dict | None: """ Extract and aggregate usage metadata from a list of WebSocket messages. + Live emits one ``usageMetadata`` per turn and Google charges per turn for every token in + the session context window, which is the current turn's tokens plus all accumulated + tokens from previous turns, so the turns add up rather than restating each other. See + the Live API note under https://cloud.google.com/vertex-ai/generative-ai/pricing. + Args: websocket_messages: List of WebSocket messages from the Live API Returns: Dictionary containing aggregated usage metadata, or None if not found """ - all_usage_metadata: Final = [] + snapshots: Final = tuple( + metadata + for message in websocket_messages + if isinstance(message, Mapping) + for metadata in (message.get("usageMetadata"),) + if isinstance(metadata, dict) + ) - # Collect all usage metadata messages - for message in websocket_messages: - if isinstance(message, dict) and "usageMetadata" in message: - all_usage_metadata.append(message["usageMetadata"]) - - if not all_usage_metadata: + if not snapshots: return None - # If only one usage metadata, return it as-is - if len(all_usage_metadata) == 1: - return all_usage_metadata[0] - - # Aggregate multiple usage metadata messages - aggregated: Final[dict[str, Any]] = { - "promptTokenCount": 0, - "candidatesTokenCount": 0, - "totalTokenCount": 0, - "promptTokensDetails": [], - "candidatesTokensDetails": [], + prompt_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals( + snapshots, "promptTokenCount", "promptTokensDetails" + ) + candidate_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals( + snapshots, "candidatesTokenCount", "candidatesTokensDetails" + ) + return { + **{key: value for key, value in snapshots[0].items() if key not in _AGGREGATED_FIELDS}, + "promptTokenCount": sum(snapshot.get("promptTokenCount", 0) for snapshot in snapshots), + "candidatesTokenCount": sum(snapshot.get("candidatesTokenCount", 0) for snapshot in snapshots), + "totalTokenCount": sum(snapshot.get("totalTokenCount", 0) for snapshot in snapshots), + "toolUsePromptTokenCount": sum(snapshot.get("toolUsePromptTokenCount", 0) for snapshot in snapshots), + "promptTokensDetails": [ + {"modality": modality, "tokenCount": count} for modality, count in prompt_totals.items() if count > 0 + ], + "candidatesTokensDetails": [ + {"modality": modality, "tokenCount": count} for modality, count in candidate_totals.items() if count > 0 + ], } - # Aggregate token counts - for usage in all_usage_metadata: - aggregated["promptTokenCount"] += usage.get("promptTokenCount", 0) - aggregated["candidatesTokenCount"] += usage.get("candidatesTokenCount", 0) - aggregated["totalTokenCount"] += usage.get("totalTokenCount", 0) - - # Aggregate token details by modality - modality_totals: Final = {} - - for usage in all_usage_metadata: - # Process prompt tokens details - for detail in usage.get("promptTokensDetails", []): - modality = detail.get("modality", "TEXT") - token_count = detail.get("tokenCount", 0) - - if modality not in modality_totals: - modality_totals[modality] = {"prompt": 0, "candidate": 0} - modality_totals[modality]["prompt"] += token_count - - # Process candidate tokens details - for detail in usage.get("candidatesTokensDetails", []): - modality = detail.get("modality", "TEXT") - token_count = detail.get("tokenCount", 0) - - if modality not in modality_totals: - modality_totals[modality] = {"prompt": 0, "candidate": 0} - modality_totals[modality]["candidate"] += token_count - - # Convert aggregated modality totals back to details format - for modality, totals in modality_totals.items(): - if totals["prompt"] > 0: - aggregated["promptTokensDetails"].append({"modality": modality, "tokenCount": totals["prompt"]}) - if totals["candidate"] > 0: - aggregated["candidatesTokensDetails"].append({"modality": modality, "tokenCount": totals["candidate"]}) - - # Add any additional fields from the first usage metadata - first_usage: Final = all_usage_metadata[0] - for key, value in first_usage.items(): - if key not in aggregated: - aggregated[key] = value - - return aggregated - - @staticmethod - def _calculate_live_api_cost( - model: str, - usage_metadata: dict, - custom_llm_provider: str = "vertex_ai", - ) -> float: - """ - Calculate cost for Vertex AI Live API based on usage metadata. - - Args: - model: The model name (e.g., "gemini-2.0-flash-live-preview-04-09") - usage_metadata: Usage metadata from the Live API response - custom_llm_provider: The LLM provider (default: "vertex_ai") - - Returns: - Total cost in USD - """ - try: - # Get model pricing information - model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider) - - verbose_proxy_logger.debug("Vertex AI Live API model info for '%s': %s", model, model_info) - - # Check if pricing info is available - if not model_info or not model_info.get("input_cost_per_token"): - verbose_proxy_logger.error("No pricing info found for %s in local model pricing database", model) - return 0.0 - - total_cost = 0.0 - - # Extract token counts from usage metadata - prompt_token_count: Final = usage_metadata.get("promptTokenCount", 0) - candidates_token_count: Final = usage_metadata.get("candidatesTokenCount", 0) - - # Calculate base text token costs - input_cost_per_token: Final = model_info.get("input_cost_per_token", 0.0) - output_cost_per_token: Final = model_info.get("output_cost_per_token", 0.0) - - total_cost += prompt_token_count * input_cost_per_token - total_cost += candidates_token_count * output_cost_per_token - - # Handle modality-specific costs if present - prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", []) - candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", []) - - # Process prompt tokens by modality - for detail in prompt_tokens_details: - modality = detail.get("modality", "TEXT") - token_count = detail.get("tokenCount", 0) - - if modality == "AUDIO": - audio_cost_per_token = model_info.get("input_cost_per_audio_token", 0.0) - total_cost += token_count * audio_cost_per_token - elif modality == "VIDEO": - # Video tokens are typically per second, but we'll treat as per token for now - video_cost_per_token = model_info.get("input_cost_per_video_per_second", 0.0) - total_cost += token_count * video_cost_per_token - # TEXT tokens are already handled above - - # Process candidate tokens by modality - for detail in candidates_tokens_details: - modality = detail.get("modality", "TEXT") - token_count = detail.get("tokenCount", 0) - - if modality == "AUDIO": - audio_cost_per_token = model_info.get("output_cost_per_audio_token", 0.0) - total_cost += token_count * audio_cost_per_token - elif modality == "VIDEO": - # Video tokens are typically per second, but we'll treat as per token for now - video_cost_per_token = model_info.get("output_cost_per_video_per_second", 0.0) - total_cost += token_count * video_cost_per_token - # TEXT tokens are already handled above - - # Handle web search costs if present - tool_use_prompt_token_count: Final = usage_metadata.get("toolUsePromptTokenCount", 0) - if tool_use_prompt_token_count > 0: - # Web search typically has a fixed cost per request - web_search_cost: Final = model_info.get("web_search_cost_per_request", 0.0) - if isinstance(web_search_cost, (int, float)) and web_search_cost > 0: - total_cost += web_search_cost - else: - # Fallback to token-based pricing for tool use - total_cost += tool_use_prompt_token_count * input_cost_per_token - - verbose_proxy_logger.debug( - f"Vertex AI Live API cost calculation - Model: {model}, " - f"Prompt tokens: {prompt_token_count}, " - f"Candidate tokens: {candidates_token_count}, " - f"Total cost: ${total_cost:.6f}" - ) - - return total_cost - - except Exception as e: - verbose_proxy_logger.error("Error calculating Vertex AI Live API cost: %s", e) - return 0.0 - @staticmethod def _create_usage_object_from_metadata( usage_metadata: dict, model: str, + grounding_requests: GroundingRequests = _NO_GROUNDING, ) -> Usage: """ Create a LiteLLM Usage object from Live API usage metadata. @@ -235,48 +246,124 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): Args: usage_metadata: Usage metadata from the Live API response model: The model name + grounding_requests: The Search and Maps grounding requests summed over the session's + turns, matching the per-turn charge Returns: LiteLLM Usage object """ - prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0) - completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0) - total_tokens: Final = usage_metadata.get("totalTokenCount", 0) + prompt_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality( + VertexAILivePassthroughLoggingHandler._resolve_detail_counts( + _detail_entries(usage_metadata.get("promptTokensDetails")), usage_metadata.get("promptTokenCount") + ) + ) + candidates_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality( + VertexAILivePassthroughLoggingHandler._resolve_detail_counts( + _detail_entries(usage_metadata.get("candidatesTokensDetails")), + usage_metadata.get("candidatesTokenCount"), + ) + ) - # Create modality-specific token details if available - prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", []) - candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", []) - - # Extract text tokens from details - text_prompt_tokens = 0 - text_completion_tokens = 0 - - for detail in prompt_tokens_details: - if detail.get("modality") == "TEXT": - text_prompt_tokens = detail.get("tokenCount", 0) - break - - for detail in candidates_tokens_details: - if detail.get("modality") == "TEXT": - text_completion_tokens = detail.get("tokenCount", 0) - break - - # If no text tokens found in details, use total counts - if text_prompt_tokens == 0: - text_prompt_tokens = prompt_tokens - if text_completion_tokens == 0: - text_completion_tokens = completion_tokens + prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0) or sum(prompt_by_modality.values()) + completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0) or sum(candidates_by_modality.values()) return Usage( - prompt_tokens=text_prompt_tokens, - completion_tokens=text_completion_tokens, - total_tokens=total_tokens, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=usage_metadata.get("totalTokenCount", 0) or (prompt_tokens + completion_tokens), + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=prompt_by_modality.get("TEXT"), + audio_tokens=prompt_by_modality.get("AUDIO"), + image_tokens=prompt_by_modality.get("IMAGE"), + video_tokens=prompt_by_modality.get("VIDEO"), + tool_use_tokens=usage_metadata.get("toolUsePromptTokenCount") or None, + web_search_requests=grounding_requests.web_search_requests, + google_maps_grounding_requests=grounding_requests.google_maps_grounding_requests, + ), + completion_tokens_details=CompletionTokensDetailsWrapper( + text_tokens=candidates_by_modality.get("TEXT"), + audio_tokens=candidates_by_modality.get("AUDIO"), + image_tokens=candidates_by_modality.get("IMAGE"), + video_tokens=candidates_by_modality.get("VIDEO"), + ), ) + def _session_usage(self, websocket_messages: Sequence[object], model: str) -> Usage | None: + usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages) + if usage_metadata is None: + return None + return self._create_usage_object_from_metadata( + usage_metadata=usage_metadata, + grounding_requests=_session_grounding_requests(websocket_messages), + model=model, + ) + + def _turn_cost( + self, + turn: Sequence[object], + model: str, + logging_obj: LiteLLMLoggingObj, + ) -> tuple[float, CostBreakdown] | None: + usage: Final = self._session_usage(turn, model) + if usage is None: + return None + cost: Final = logging_obj._response_cost_calculator( # pyright: ignore[reportPrivateUsage] # the call's own calculator keeps custom pricing and the deployment's region in step with the spend row + result=ModelResponse(model=model, usage=usage), + litellm_model_name=model, + ) + if cost is None: + return None + breakdown: Final = logging_obj.cost_breakdown + return None if breakdown is None else (cost, breakdown) + + def _session_cost( + self, + websocket_messages: Sequence[object], + model: str, + logging_obj: LiteLLMLoggingObj, + ) -> float | None: + """Price each turn on its own tokens and grounding, so two grounded turns pay the query fee twice. + + The fixed cost margin is a flat per-request fee, so the session's single spend row carries it once + rather than once per turn. + """ + turn_costs: Final = tuple(self._turn_cost(turn, model, logging_obj) for turn in _turns(websocket_messages)) + priced: Final = tuple(turn_cost for turn_cost in turn_costs if turn_cost is not None) + if not priced or len(priced) != len(turn_costs): + return None + breakdowns: Final = tuple(breakdown for _, breakdown in priced) + first: Final = breakdowns[0] + fixed_margin: Final = first.get("margin_fixed_amount") or 0.0 + duplicated_fixed_margin: Final = fixed_margin * (len(priced) - 1) + total_cost: Final = sum(cost for cost, _ in priced) - duplicated_fixed_margin + summed_margin_total: Final = _summed(breakdowns, "margin_total_amount") + margin_total_amount: Final = ( + None if summed_margin_total is None else summed_margin_total - duplicated_fixed_margin + ) + logging_obj.set_cost_breakdown( + input_cost=_summed(breakdowns, "input_cost") or 0.0, + output_cost=_summed(breakdowns, "output_cost") or 0.0, + total_cost=total_cost, + cost_for_built_in_tools_cost_usd_dollar=_summed(breakdowns, "tool_usage_cost") or 0.0, + original_cost=_summed(breakdowns, "original_cost"), + discount_percent=first.get("discount_percent"), + discount_amount=_summed(breakdowns, "discount_amount"), + margin_percent=first.get("margin_percent"), + margin_fixed_amount=first.get("margin_fixed_amount"), + margin_total_amount=margin_total_amount, + cache_read_cost=_summed(breakdowns, "cache_read_cost"), + cache_creation_cost=_summed(breakdowns, "cache_creation_cost"), + reasoning_cost=_summed(breakdowns, "reasoning_cost"), + service_tier=first.get("service_tier"), + data_residency=first.get("data_residency"), + vertex_location=first.get("vertex_location"), + ) + return total_cost + def vertex_ai_live_passthrough_handler( self, - websocket_messages: list[dict], - logging_obj, + websocket_messages: Sequence[object], + logging_obj: LiteLLMLoggingObj, url_route: str, start_time: datetime, end_time: datetime, @@ -300,34 +387,25 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): """ try: # Extract model from request body or kwargs - model: Final = kwargs.get("model", "gemini-2.0-flash-live-preview-04-09") + requested_model: Final = kwargs.get("model") + model: Final = ( + requested_model if isinstance(requested_model, str) else "gemini-2.0-flash-live-preview-04-09" + ) custom_llm_provider: Final = kwargs.get("custom_llm_provider", "vertex_ai") verbose_proxy_logger.debug( "Vertex AI Live API model: %s, custom_llm_provider: %s", model, custom_llm_provider ) - # Extract usage metadata from WebSocket messages - usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages) + usage: Final = self._session_usage(websocket_messages, model) - if not usage_metadata: + if usage is None: verbose_proxy_logger.warning("No usage metadata found in Vertex AI Live API WebSocket messages") return { "result": None, "kwargs": kwargs, } - # Calculate cost using Live API specific pricing - response_cost: Final = self._calculate_live_api_cost( - model=model, - usage_metadata=usage_metadata, - custom_llm_provider=custom_llm_provider, - ) - - # Create Usage object for standard LiteLLM logging - usage: Final = self._create_usage_object_from_metadata( - usage_metadata=usage_metadata, - model=model, - ) + response_cost: Final = self._session_cost(websocket_messages, model, logging_obj) # Create a mock ModelResponse for standard logging litellm_model_response: Final = ModelResponse( @@ -338,9 +416,9 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): usage=usage, choices=[], ) + if response_cost is not None: + litellm_model_response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads the cost off the response's hidden params; the constructor's hidden_params kwarg is reset by pydantic - # Update kwargs with cost information - kwargs["response_cost"] = response_cost kwargs["model"] = model kwargs["custom_llm_provider"] = custom_llm_provider @@ -348,12 +426,15 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): import re allowed_pattern: Final = re.compile(r"^[A-Za-z0-9._\-:]+$") - safe_model: Final = model if isinstance(model, str) and allowed_pattern.match(model) else "[REDACTED]" + safe_model: Final = model if allowed_pattern.match(model) else "[REDACTED]" verbose_proxy_logger.debug( - f"Vertex AI Live API passthrough cost tracking - " - f"Model: {safe_model}, Cost: ${response_cost:.6f}, " - f"Prompt tokens: {usage.prompt_tokens}, " - f"Completion tokens: {usage.completion_tokens}" + "Vertex AI Live API passthrough cost tracking - Model: %s, " + "Prompt tokens: %s %s, Completion tokens: %s %s", + safe_model, + usage.prompt_tokens, + usage.prompt_tokens_details, + usage.completion_tokens, + usage.completion_tokens_details, ) return { diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index ea4ede7e513..b66c295d1aa 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -2090,6 +2090,22 @@ def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Calla return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload +def _resolved_vertex_live_setup( + setup_data: Mapping[str, object], setup_model_rewriter: Callable[[str], str] | None +) -> Mapping[str, object]: + """ + Give the model extractor the same fully qualified path the upstream will receive. + + Clients may name a bare gateway alias, which the rewriter turns into a ``projects/...`` path before + it reaches Vertex. The extractor only reads a path containing ``/models/``, so running it on the raw + frame logs the session as ``unknown`` at no cost, which is precisely the supported client form + """ + setup_model: Final = setup_data.get("model") + if setup_model_rewriter is None or not isinstance(setup_model, str): + return setup_data + return {**setup_data, "model": setup_model_rewriter(setup_model)} + + def _truncated_close_reason(reason: str) -> str: """ Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character @@ -2314,7 +2330,9 @@ async def websocket_passthrough_request( setup_data, ) if isinstance(setup_data, dict) and "model" in setup_data: - extracted_model = _extract_model_from_vertex_ai_setup(setup_data) + extracted_model = _extract_model_from_vertex_ai_setup( + _resolved_vertex_live_setup(setup_data, setup_model_rewriter) + ) if extracted_model: kwargs["model"] = extracted_model kwargs["custom_llm_provider"] = "vertex_ai-language-models" diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 4be0235adbb..b310fc661c4 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -270,6 +270,24 @@ class PassThroughStreamingHandler: - Vertex AI - OpenAI """ + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + _is_message_stop_chunk, # pyright: ignore[reportPrivateUsage] # both native stream paths share terminal-event detection + _is_provider_error_chunk, # pyright: ignore[reportPrivateUsage] # provider errors must not become cache evidence + ) + + # Transport reads can split event names and JSON payloads. Recognize terminal + # events only after the shared SSE framer has reassembled the collected bytes. + complete_frames, incomplete_tail = split_complete_sse_frames( + b"".join(raw_bytes) if endpoint_type == EndpointType.ANTHROPIC else b"" + ) + litellm_logging_obj.model_call_details[ # rebind-ok: stamp evidence on the per-request state read by callbacks + "prompt_cache_response_complete" + ] = ( + endpoint_type == EndpointType.ANTHROPIC + and not incomplete_tail.strip() + and _is_message_stop_chunk(complete_frames) + and not _is_provider_error_chunk(complete_frames) + ) try: ( standard_logging_response_object, diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index ad45781d5d2..ed193c7f434 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -58,15 +58,6 @@ class UndeliverableStreamRewrite(Exception): self.guardrail_name: Final = guardrail_name -class UnappliableRequestRewrite(Exception): - def __init__(self, guardrail_name: str) -> None: - super().__init__( - f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, " - "so the request was rejected rather than sent unrewritten" - ) - self.guardrail_name: Final = guardrail_name - - def _tool_call_shape(tool_call: object) -> tuple[object, object]: plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call function: Final = plain.get("function") if isinstance(plain, Mapping) else None diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 04145eb68ef..963fca104f8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -476,9 +476,10 @@ from litellm.proxy.hooks.prompt_injection_detection import ( from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.list_api.common import ( - PROBLEM_TYPE_BASE, ManagementProblem, + ValidationErrorDetail, problem_response, + request_validation_problem, ) from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.logging_endpoints.callback_logs_endpoints import ( @@ -601,7 +602,6 @@ from litellm.proxy.spend_tracking.spend_event_producer import ( SpendEventProducer, build_spend_event_producer, ) -from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail try: from litellm.proxy.enterprise_billing.billing_metrics import ( @@ -928,6 +928,7 @@ def cleanup_router_config_variables(): user_custom_auth_path, \ user_custom_key_generate, \ user_custom_key_update, \ + user_custom_key_policy, \ user_custom_sso, \ user_custom_ui_sso_sign_in_handler, \ use_background_health_checks, \ @@ -945,6 +946,7 @@ def cleanup_router_config_variables(): user_custom_auth_path = None user_custom_key_generate = None user_custom_key_update = None + user_custom_key_policy = None TEAM_METADATA_VALIDATOR_REGISTRY.set(None) TEAM_METADATA_SCHEMA_REGISTRY.set(()) user_custom_sso = None @@ -1787,40 +1789,13 @@ class _ExceptionRow(TypedDict, total=False): exception_counts: Mapping[str, int] -class _ValidationErrorDetail(TypedDict): - type: ReadOnly[str] - loc: ReadOnly[tuple[int | str, ...]] - msg: ReadOnly[str] - - -def _is_length_error_of_rejected_items(error: _ValidationErrorDetail, errors: Sequence[_ValidationErrorDetail]) -> bool: - """pydantic counts only items that validated, so a bad item also trips the parent's min_length.""" - return error["type"] == "too_short" and any( - len(other["loc"]) > len(error["loc"]) and other["loc"][: len(error["loc"])] == error["loc"] for other in errors - ) - - @app.exception_handler(RequestValidationError) async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError): if request.url.path.startswith(MANAGEMENT_V1_PREFIX): - raw_errors: Final[Sequence[_ValidationErrorDetail]] = exc.errors() - validation_errors: Final = tuple( - error for error in raw_errors if not _is_length_error_of_rejected_items(error, raw_errors) - ) - in_body: Final = any(error["loc"] and error["loc"][0] == "body" for error in validation_errors) - status: Final = 422 if in_body else 400 - _close_dangling_otel_server_span(request, status, exc=exc) - return problem_response( - ProblemDetail( - type=f"{PROBLEM_TYPE_BASE}{'invalid-request-body' if in_body else 'invalid-query-parameter'}", - title="Invalid request body" if in_body else "Invalid query parameter", - status=status, - detail="; ".join( - f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in validation_errors - ) - or "The request is invalid.", - ) - ) + validation_errors: Final[Sequence[ValidationErrorDetail]] = exc.errors() + problem: Final = request_validation_problem(validation_errors) + _close_dangling_otel_server_span(request, problem.status, exc=exc) + return problem_response(problem) _close_dangling_otel_server_span(request, 422, exc=exc) return JSONResponse( status_code=422, @@ -2382,6 +2357,7 @@ user_custom_key_generate = None _pkce_no_redis_warning_emitted: bool = False _cp_no_redis_warning_emitted: bool = False user_custom_key_update = None +user_custom_key_policy = None user_custom_sso = None user_custom_ui_sso_sign_in_handler = None use_background_health_checks = None @@ -4269,6 +4245,7 @@ _DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: Final[dict[str, tuple[str, ...]]] = { "custom_auth", "custom_key_generate", "custom_key_update", + "custom_key_policy", "custom_team_metadata_validate", "custom_sso", "custom_ui_sso_sign_in_handler", @@ -5418,6 +5395,7 @@ class ProxyConfig: user_custom_auth_path, \ user_custom_key_generate, \ user_custom_key_update, \ + user_custom_key_policy, \ user_custom_sso, \ user_custom_ui_sso_sign_in_handler, \ use_background_health_checks, \ @@ -5955,6 +5933,10 @@ class ProxyConfig: if custom_key_update is not None: user_custom_key_update = get_instance_fn(value=custom_key_update, config_file_path=config_file_path) + custom_key_policy: Final = general_settings.get("custom_key_policy", None) + if custom_key_policy is not None: + user_custom_key_policy = get_instance_fn(value=custom_key_policy, config_file_path=config_file_path) + custom_team_metadata_validate: Final = general_settings.get("custom_team_metadata_validate", None) TEAM_METADATA_VALIDATOR_REGISTRY.set( get_instance_fn(value=custom_team_metadata_validate, config_file_path=config_file_path) @@ -9559,6 +9541,7 @@ class ProxyStartupEvent: gate the first duration window. """ await generate_key_helper_fn( + llm_router=llm_router, request_type="user", table_name="user", user_id=LITELLM_PROXY_BUDGET_NAME, @@ -16303,6 +16286,7 @@ async def _generate_onboarding_ui_session_token(user_obj: _UserTableRow) -> str: global master_key, general_settings response: Final = await generate_key_helper_fn( + llm_router=llm_router, request_type="key", **{ "user_role": user_obj.user_role, diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 3ef21996b9c..4bcdf6aad22 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -159,7 +159,7 @@ def _get_spend_logs_metadata( requester_ip_address=None, additional_usage_values=None, applied_guardrails=None, - status=None or "success", + status="success", error_information=None, proxy_server_request=None, batch_models=None, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a4875321a2a..48a6baa3587 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4,6 +4,7 @@ import copy import hashlib import inspect import json +import math import os import smtplib import ssl @@ -6405,7 +6406,7 @@ class PrismaClient: return None try: value: Final = float(response_time_ms) - return value if value == value and value not in (float("inf"), float("-inf")) else None + return value if math.isfinite(value) else None except (ValueError, TypeError): verbose_proxy_logger.warning("Invalid response_time_ms value: %s", response_time_ms) return None diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 44c47af57f4..d67e4555a29 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -30,6 +30,7 @@ from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import CallTypes, LlmProviders from litellm.utils import ProviderConfigManager +from ..litellm_core_utils.credential_accessor import CredentialAccessor from ..litellm_core_utils.get_litellm_params import get_litellm_params from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ..llms.azure.common_utils import get_azure_ad_token @@ -54,6 +55,17 @@ xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() base_llm_http_handler = BaseLLMHTTPHandler() _EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({}) +_EMPTY_AUTH_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) + + +def _model_params_with_stored_credentials(model_params: Mapping[str, Any]) -> Mapping[str, Any]: + credential_name: Final = model_params.get("litellm_credential_name") + credential_values: Final = ( + CredentialAccessor.get_credential_values(credential_name) + if isinstance(credential_name, str) + else _EMPTY_MODEL_PARAMS + ) + return MappingProxyType({**credential_values, **model_params}) def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]: @@ -591,13 +603,15 @@ def _azure_realtime_health_protocol( def _realtime_health_check_auth_headers( custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any] -) -> Mapping[str, str | None]: - if custom_llm_provider != "azure": - return MappingProxyType({"api-key": api_key}) - return azure_realtime.get_auth_headers( - api_key=api_key, - azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))), - ) +) -> Mapping[str, str]: + if custom_llm_provider == "azure": + return azure_realtime.get_auth_headers( + api_key=api_key, + azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))), + ) + if api_key is None: + return _EMPTY_AUTH_HEADERS + return MappingProxyType({"Authorization": f"Bearer {api_key}"}) async def _realtime_health_check( @@ -629,34 +643,46 @@ async def _realtime_health_check( """ import websockets + resolved_params: Final = _model_params_with_stored_credentials(model_params or _EMPTY_MODEL_PARAMS) + resolved_api_key: Final = cast( # cast-ok: provider parameters expose optional string credentials + str | None, api_key or resolved_params.get("api_key") + ) + resolved_api_base: Final = cast( # cast-ok: provider parameters expose optional string endpoints + str | None, api_base or resolved_params.get("api_base") + ) + resolved_api_version: Final = cast( # cast-ok: provider parameters expose optional string versions + str | None, api_version or resolved_params.get("api_version") + ) url: str | None = None auth_headers: Final = _realtime_health_check_auth_headers( custom_llm_provider=custom_llm_provider, - api_key=api_key, - model_params=model_params or _EMPTY_MODEL_PARAMS, + api_key=resolved_api_key, + model_params=resolved_params, ) if custom_llm_provider == "azure": resolved_protocol, azure_query_params = _azure_realtime_health_protocol( model=model, realtime_protocol=realtime_protocol, - model_params=model_params or _EMPTY_MODEL_PARAMS, + model_params=resolved_params, ) url = azure_realtime._construct_url( - api_base=api_base or "", + api_base=resolved_api_base or "", model=model, - api_version=api_version or "2024-10-01-preview", + api_version=resolved_api_version or "2024-10-01-preview", realtime_protocol=resolved_protocol, query_params=azure_query_params, ) elif custom_llm_provider == "openai": url = openai_realtime._construct_url( - api_base=api_base or "https://api.openai.com/", + api_base=resolved_api_base or "https://api.openai.com/", query_params={"model": model}, ) elif custom_llm_provider == "xai": - url = xai_realtime._construct_url(api_base=api_base or "https://api.x.ai/v1", query_params={"model": model}) + url = xai_realtime._construct_url( + api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model} + ) elif custom_llm_provider == "vertex_ai": - vertex_model_params: Final = model_params or {} + vertex_model_params: Final = dict(resolved_params) resolved_location: Final = vertex_llm_base.get_vertex_region( vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params), model=model, @@ -675,19 +701,19 @@ async def _realtime_health_check( project=resolved_project, location=resolved_location, ) - url = vertex_realtime_config.get_complete_url(api_base=api_base, model=model) - ssl_context = get_shared_realtime_ssl_context() + url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model) + vertex_ssl_context: Final = get_shared_realtime_ssl_context() headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None) async with websockets.connect( url, additional_headers=headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, + ssl=vertex_ssl_context, ): return True else: raise ValueError(f"Unsupported model: {model}") - ssl_context = get_shared_realtime_ssl_context() + ssl_context: Final = get_shared_realtime_ssl_context() async with websockets.connect( url, additional_headers=auth_headers, diff --git a/litellm/responses/additional_tools.py b/litellm/responses/additional_tools.py new file mode 100644 index 00000000000..ea0d7af350c --- /dev/null +++ b/litellm/responses/additional_tools.py @@ -0,0 +1,65 @@ +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Final, cast # noqa: TID251 # validating the openai tool union strips vendor keys from raw tools + +from pydantic import BaseModel, ValidationError + +from litellm._logging import verbose_logger +from litellm.types.llms.openai import ALL_RESPONSES_API_TOOL_PARAMS, ResponseInputParam + +ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools" + + +class _InputItemType(BaseModel): + type: str = "" + + +class _AdditionalToolsItem(BaseModel): + tools: tuple[dict[str, object], ...] = () + + +@dataclass(frozen=True, slots=True) +class HoistedAdditionalTools: + input: str | ResponseInputParam + tools: tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...] + hoisted: tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...] + + +def _is_additional_tools_item(item: object) -> bool: + try: + return _InputItemType.model_validate(item).type == ADDITIONAL_TOOLS_INPUT_ITEM_TYPE + except ValidationError: + return False + + +def _tools_of_item(item: object) -> tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...]: + try: + parsed: Final = _AdditionalToolsItem.model_validate(item) + except ValidationError: + return () + return tuple( + cast( + "ALL_RESPONSES_API_TOOL_PARAMS", tool + ) # cast-ok: nested tools carry the same raw tool JSON as top-level tools + for tool in parsed.tools + ) + + +def hoist_additional_tools( + input: str | ResponseInputParam, + tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None, +) -> HoistedAdditionalTools: + existing: Final = tuple(tools or ()) + if isinstance(input, str): + return HoistedAdditionalTools(input=input, tools=existing, hoisted=()) + items: Final = tuple(item for item in input if _is_additional_tools_item(item)) + if not items: + return HoistedAdditionalTools(input=input, tools=existing, hoisted=()) + hoisted: Final = tuple(tool for item in items for tool in _tools_of_item(item)) + verbose_logger.debug( + "Responses API: hoisting %d tool(s) out of %d 'additional_tools' input item(s) into the top-level tools param.", + len(hoisted), + len(items), + ) + remaining_input: Final = [item for item in input if not _is_additional_tools_item(item)] + return HoistedAdditionalTools(input=remaining_input, tools=(*existing, *hoisted), hoisted=hoisted) diff --git a/litellm/responses/litellm_completion_transformation/custom_tools.py b/litellm/responses/litellm_completion_transformation/custom_tools.py index 4aa489d9e50..7888a07e248 100644 --- a/litellm/responses/litellm_completion_transformation/custom_tools.py +++ b/litellm/responses/litellm_completion_transformation/custom_tools.py @@ -39,15 +39,38 @@ def openai_shaped_tool_call_item_id(item_type: str, tool_id: str) -> str: return f"{prefix}_{tool_id}" +class _ToolNameFields(BaseModel): + type: str = "" + name: str = "" + tools: tuple[object, ...] = () + + +def _tool_name_fields_of(tool: object) -> _ToolNameFields | None: + try: + return _ToolNameFields.model_validate(tool) + except ValidationError: + return None + + +def _custom_tool_name_of(tool: object) -> str | None: + parsed: Final = _tool_name_fields_of(tool) + if parsed is None or parsed.type != "custom" or not parsed.name: + return None + return parsed.name + + +def _nested_tools_of(tool: object) -> tuple[object, ...]: + parsed: Final = _tool_name_fields_of(tool) + if parsed is None or parsed.type != "namespace": + return () + return parsed.tools + + def extract_custom_tool_names(tools: Sequence[object] | None) -> set[str]: - """Extract names of tools originally defined as ``type: "custom"``.""" - if not tools: - return set() - names: Final[set[str]] = set() - for tool in tools: - if isinstance(tool, dict) and tool.get("type") == "custom" and "name" in tool: - names.add(tool["name"]) - return names + """Extract names of ``type: "custom"`` tools, at the top level or one level inside a ``namespace`` tool.""" + top_level: Final = tuple(tools or ()) + nested: Final = tuple(nested_tool for tool in top_level for nested_tool in _nested_tools_of(tool)) + return {name for tool in (*top_level, *nested) if (name := _custom_tool_name_of(tool)) is not None} def is_custom_tool_call(tool_name: str, custom_tool_names: set[str]) -> bool: @@ -143,7 +166,7 @@ def validated_allowed_callers(value: object) -> list[str] | None: raise ValueError("allowed_callers must be a list of strings") from exc -def _grammar_suffix(fmt: object) -> str: +def custom_tool_grammar_suffix(fmt: object) -> str: try: parsed: Final = _CustomToolFormat.model_validate(fmt) except ValidationError: @@ -167,7 +190,9 @@ def convert_custom_tool_to_function_tool(tool: Mapping[str, object]) -> ChatComp raw_name: Final = tool.get("name") name: Final = raw_name if isinstance(raw_name, str) else "" raw_description: Final = tool.get("description") - description = (raw_description if isinstance(raw_description, str) else "") + _grammar_suffix(tool.get("format")) + description: Final = (raw_description if isinstance(raw_description, str) else "") + custom_tool_grammar_suffix( + tool.get("format") + ) allowed_callers: Final = validated_allowed_callers(tool.get("allowed_callers")) function_chunk: Final = ChatCompletionToolParamFunctionChunk( name=name, diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index a0e8cd278e6..505b5b09433 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -6,6 +6,7 @@ from collections.abc import Coroutine, Mapping from typing import Final import litellm +from litellm.responses.additional_tools import hoist_additional_tools from litellm.responses.litellm_completion_transformation.streaming_iterator import ( LiteLLMCompletionStreamingIterator, ) @@ -37,11 +38,16 @@ class LiteLLMCompletionTransformationHandler: | BaseResponsesAPIStreamingIterator | Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator] ): + hoisted: Final = hoist_additional_tools(input, responses_api_request.get("tools")) + bridged_input: Final = hoisted.input + bridged_request: Final[ResponsesAPIOptionalRequestParams] = ( + {**responses_api_request, "tools": list(hoisted.tools)} if hoisted.hoisted else responses_api_request + ) litellm_completion_request: Final[dict] = ( LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( model=model, - input=input, - responses_api_request=responses_api_request, + input=bridged_input, + responses_api_request=bridged_request, custom_llm_provider=custom_llm_provider, stream=stream, extra_headers=extra_headers, @@ -52,8 +58,8 @@ class LiteLLMCompletionTransformationHandler: if _is_async: return self.async_response_api_handler( litellm_completion_request=litellm_completion_request, - request_input=input, - responses_api_request=responses_api_request, + request_input=bridged_input, + responses_api_request=bridged_request, **kwargs, ) @@ -70,8 +76,8 @@ class LiteLLMCompletionTransformationHandler: responses_api_response: Final[ResponsesAPIResponse] = ( LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( chat_completion_response=litellm_completion_response, - request_input=input, - responses_api_request=responses_api_request, + request_input=bridged_input, + responses_api_request=bridged_request, ) ) @@ -81,8 +87,8 @@ class LiteLLMCompletionTransformationHandler: return LiteLLMCompletionStreamingIterator( model=model, litellm_custom_stream_wrapper=litellm_completion_response, - request_input=input, - responses_api_request=responses_api_request, + request_input=bridged_input, + responses_api_request=bridged_request, custom_llm_provider=custom_llm_provider, litellm_metadata=kwargs.get("litellm_metadata", {}), ) diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 126b976e2c5..c28b5558c75 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -8,6 +8,7 @@ from litellm.main import stream_chunk_builder from litellm.responses.litellm_completion_transformation.custom_tools import ( build_tool_call_item_kwargs, extract_custom_tool_names, + is_custom_tool_call, serialize_tool_call_arguments, ) from litellm.responses.litellm_completion_transformation.transformation import ( @@ -166,6 +167,14 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return tool_name, namespace return fn_name, None + def _tool_call_item_kwargs(self, call_id: str, fn_name: str, arguments: str, status: str) -> dict[str, str]: + item_kwargs: Final = build_tool_call_item_kwargs(call_id, fn_name, arguments, status, self._custom_tool_names) + if is_custom_tool_call(fn_name, self._custom_tool_names): + return item_kwargs + tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name) + namespace_kwargs: Final = {"namespace": tool_namespace} if tool_namespace else {} + return {**item_kwargs, "name": tool_name, **namespace_kwargs} + def _is_reasoning_end(self, chunk): delta: Final = chunk.choices[0].delta @@ -244,17 +253,13 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): else: fn_name = str(getattr(fn, "name", "") or "") fn_args_delta = serialize_tool_call_arguments(getattr(fn, "arguments", "")) - tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name) output_index = self._get_or_assign_tool_output_index(call_id) if call_id not in self._tool_args_by_call_id: self._tool_args_by_call_id[call_id] = "" self._sequence_number += 1 - names = self._custom_tool_names - item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names) + item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, "", "in_progress") self._tool_item_id_by_call_id[call_id] = item_kwargs["id"] - if tool_namespace: - item_kwargs["namespace"] = tool_namespace event = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=output_index, @@ -315,7 +320,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): else: fn_name = str(getattr(fn, "name", "") or "") fn_args = serialize_tool_call_arguments(getattr(fn, "arguments", "")) - tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name) web_search_call = self._web_search_calls.get(call_id) if web_search_call is not None: if call_id not in self._queued_web_search_call_ids: @@ -330,11 +334,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if is_new_tool_call: self._tool_args_by_call_id[call_id] = "" self._sequence_number += 1 - names = self._custom_tool_names - item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names) + item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, "", "in_progress") self._tool_item_id_by_call_id[call_id] = item_kwargs["id"] - if tool_namespace: - item_kwargs["namespace"] = tool_namespace event = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=output_index, @@ -376,11 +377,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._pending_tool_events.append(done_event) self._sequence_number += 1 - names = self._custom_tool_names - item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, final_args, "completed", names) + item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, final_args, "completed") item_kwargs["id"] = self._tool_item_id_by_call_id.setdefault(call_id, item_kwargs["id"]) - if tool_namespace: - item_kwargs["namespace"] = tool_namespace item_done_event = OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, output_index=output_index, diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 64324c6cad8..01fb6cb483d 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -110,6 +110,7 @@ NamespaceTool: TypeAlias = Mapping[str, object] ResponseTools: TypeAlias = Sequence[Mapping[str, object]] | None ChatToolParam: TypeAlias = ChatCompletionToolParam | OpenAIMcpServerTool NAMESPACE_DESCRIPTION_SEPARATOR: Final = "\n\n" +NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: Final = frozenset({"function", "custom"}) @dataclass(frozen=True, slots=True) @@ -1891,9 +1892,21 @@ class LiteLLMCompletionResponsesConfig: namespace_tool: NamespaceTool, nested: bool, ) -> ChatCompletionToolParam | None: - if nested and namespace_tool.get("type") != "function": + tool_type: Final = namespace_tool.get("type") + if nested and tool_type not in NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: return None + raw_description: Final = str(namespace_tool.get("description") or "") + description: Final = ( + f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}{raw_description}" + if nested and namespace_description and raw_description + else namespace_description + if nested and namespace_description + else raw_description + ) + if nested and tool_type == "custom": + return convert_custom_tool_to_function_tool({**namespace_tool, "description": description}) + raw_parameters: Final = namespace_tool.get("parameters") parameters: Final = ( MappingProxyType(raw_parameters) if isinstance(raw_parameters, Mapping) else MappingProxyType({}) @@ -1902,14 +1915,6 @@ class LiteLLMCompletionResponsesConfig: parameters if parameters and "type" in parameters else MappingProxyType({**parameters, "type": "object"}) ) tool_name: Final = str(namespace_tool.get("name") or "") - raw_description: Final = str(namespace_tool.get("description") or "") - description: Final = ( - f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}{raw_description}" - if nested and namespace_description and raw_description - else namespace_description - if nested and namespace_description - else raw_description - ) chat_tool_name: Final = f"{namespace}__{tool_name}" if nested else tool_name function: Final = ChatCompletionToolParamFunctionChunk( name=chat_tool_name, @@ -2826,6 +2831,22 @@ class LiteLLMCompletionResponsesConfig: if cache_write_tokens is not None else MappingProxyType({}) ) + # The cost path reads the grounding counters off the input details, and a realtime + # session's usage is rebuilt from its own response.done, so dropping them here bills + # no per-query grounding fee at all. + grounding_request_counts: Final[Mapping[str, int]] = MappingProxyType( + { + counter: count + for counter, count in ( + ("web_search_requests", getattr(prompt_details, "web_search_requests", None)), + ( + "google_maps_grounding_requests", + getattr(prompt_details, "google_maps_grounding_requests", None), + ), + ) + if count is not None + } + ) response_usage.input_tokens_details = InputTokensDetails( cached_tokens=prompt_details.cached_tokens if prompt_details.cached_tokens is not None else 0, text_tokens=prompt_details.text_tokens, @@ -2834,6 +2855,7 @@ class LiteLLMCompletionResponsesConfig: cached_tokens_details if isinstance(cached_tokens_details, CachedTokensDetails) else None ), **cache_write_extra, + **grounding_request_counts, ) # Translate completion_tokens_details to output_tokens_details diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 40ff88fc557..b39e130242d 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -221,6 +221,13 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None ) +def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool: + if isinstance(mapped_exception, litellm.ContentPolicyViolationError): + return True + status_code: Final = getattr(mapped_exception, "status_code", None) + return not isinstance(status_code, int) or status_code >= 500 or status_code == 429 + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -521,15 +528,8 @@ class BaseResponsesAPIStreamingIterator: getattr(self.completed_response, "response", None) if self.completed_response else None ) error_info: Final = getattr(response_obj, "error", None) if response_obj else None - error_message, error_type, error_code = _error_event_fields(error_info) self._record_failed_response_usage(response_obj) - exception: Final = litellm.APIError( - status_code=_status_code_for_error_fields(error_type, error_code), - message=error_message, - llm_provider=self.custom_llm_provider or "", - model=self.model or "", - ) - self._handle_failure(exception) + self._handle_failure(self._map_error_event_exception(error_info)) def _record_failed_response_usage(self, response_obj: ResponsesAPIResponse | None) -> None: if response_obj is None or self.logging_obj is None: @@ -551,6 +551,28 @@ class BaseResponsesAPIStreamingIterator: self.logging_obj._response_cost_calculator(result=response_obj) or 0.0 ) + def _map_error_event_exception(self, error_obj: object) -> Exception: + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + error_message, error_type, error_code = _error_event_fields(error_obj) + status_code: Final = _status_code_for_error_fields(error_type, error_code) + error_body: Final = {"message": error_message, "type": error_type, "code": error_code} + provider_exception: Final = BaseLLMException( + status_code=status_code, + message=f"Error code: {status_code} - {{'error': {error_body}}}", + body=error_body, + ) + try: + return litellm.exception_type( + model=self.model or "", + custom_llm_provider=self.custom_llm_provider or "", + original_exception=provider_exception, + completion_kwargs={}, + extra_kwargs={}, + ) + except Exception as mapped_exception: + return mapped_exception + def _maybe_raise_for_error_event(self, result: object) -> None: chunk_type: Final = getattr(result, "type", None) if chunk_type not in ("error", "response.failed"): @@ -562,15 +584,8 @@ class BaseResponsesAPIStreamingIterator: else getattr(result, "error", None) ) - error_message, error_type, error_code = _error_event_fields(error_obj) - status_code: Final = _status_code_for_error_fields(error_type, error_code) - mapped_exception: Final = litellm.APIError( - status_code=status_code, - message=error_message, - llm_provider=self.custom_llm_provider or "", - model=self.model or "", - ) - if 400 <= status_code < 500 and status_code != 429: + mapped_exception: Final = self._map_error_event_exception(error_obj) + if not _mid_stream_fallback_eligible(mapped_exception): raise mapped_exception raise MidStreamFallbackError( message=str(mapped_exception), diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index d63e3ddf0aa..41a3ded7022 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -1183,6 +1183,10 @@ class ResponseAPILoggingUtils: response_api_usage.input_tokens_details, "cached_tokens_details", None ), cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None), + web_search_requests=getattr(response_api_usage.input_tokens_details, "web_search_requests", None), + google_maps_grounding_requests=getattr( + response_api_usage.input_tokens_details, "google_maps_grounding_requests", None + ), ) completion_tokens_details: CompletionTokensDetailsWrapper | None = None output_tokens_details: Final[OutputTokensDetails | None] = getattr( diff --git a/litellm/router.py b/litellm/router.py index 3bbfc03919a..a414bdabf80 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3268,8 +3268,15 @@ class Router: kwargs=initial_kwargs, metadata_variable_name="litellm_metadata", ) + # The content-policy dispatch branch matches on the trigger's own type, so a refusal's + # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. + fallback_trigger: Final[Exception] = ( + e.original_exception + if isinstance(e.original_exception, litellm.ContentPolicyViolationError) + else e + ) fallback_response = await self.async_function_with_fallbacks_common_utils( - e=e, + e=fallback_trigger, disable_fallbacks=False, fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, @@ -4105,16 +4112,16 @@ class Router: models: Final = [m.strip() for m in model.split(",")] async def _async_completion_no_exceptions( - model: str, messages: list[dict[str, str]], stream: bool, **kwargs: Any + model_name: str, messages: list[dict[str, str]], stream: bool, **kwargs: Any ) -> ModelResponse | CustomStreamWrapper | Exception: """ Wrapper around self.acompletion that catches exceptions and returns them as a result """ try: - result = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) + result = await self.acompletion(model=model_name, messages=messages, stream=stream, **kwargs) return result except asyncio.CancelledError: - verbose_router_logger.debug("Received 'task.cancel'. Cancelling call w/ model=%s.", model) + verbose_router_logger.debug("Received 'task.cancel'. Cancelling call w/ model=%s.", model_name) raise except Exception as e: return e @@ -4141,9 +4148,9 @@ class Router: except KeyError: pass - for model in models: + for model_name in models: task = asyncio.create_task( - _async_completion_no_exceptions(model=model, messages=messages, stream=stream, **kwargs) + _async_completion_no_exceptions(model_name=model_name, messages=messages, stream=stream, **kwargs) ) pending_tasks.append(task) @@ -4842,6 +4849,7 @@ class Router: model=model, messages=messages, specific_deployment=kwargs.pop("specific_deployment", None), + request_kwargs=kwargs, ) data: Final = deployment["litellm_params"].copy() @@ -5156,13 +5164,11 @@ class Router: return healthy_deployments[0] # Use simple_shuffle for weighted selection - return cast( - GuardrailTypedDict, - simple_shuffle( - llm_router_instance=self, - healthy_deployments=healthy_deployments, - model=guardrail_name, - ), + return simple_shuffle( + resolve_model_alias=self._get_model_from_alias, + healthy_deployments=healthy_deployments, + model=guardrail_name, + request_kwargs=None, ) async def _ageneric_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs): @@ -8371,7 +8377,8 @@ class Router: def log_retry(self, kwargs: dict, e: Exception) -> dict: """ - When a retry or fallback happens, record which model group, deployment and attempt just failed and why + When a retry or fallback happens, record which model group, deployment and attempt just failed and why, + and count it toward the request-wide num_retries_per_request cap """ from litellm.types.router import RetryAttemptRecord @@ -8395,7 +8402,10 @@ class Router: else () ) breadcrumbs: Final = (*kept_breadcrumbs, attempt_record) + earlier: Final = request_metadata.get("request_retry_count") + request_retry_count: Final = (earlier if type(earlier) is int and 0 <= earlier else 0) + 1 kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict + kwargs[_metadata_var]["request_retry_count"] = request_retry_count # rebind-ok: same dict, read by the cap return kwargs def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int: @@ -13038,9 +13048,10 @@ class Router: start_time: Final = time.time() if strategy == "simple-shuffle": return simple_shuffle( - llm_router_instance=self, + resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, + request_kwargs=request_kwargs, ) deployment: Final = await self._select_deployment_async( strategy=strategy, @@ -13183,9 +13194,10 @@ class Router: start_time: Final = time.perf_counter() if strategy == "simple-shuffle": return simple_shuffle( - llm_router_instance=self, + resolve_model_alias=self._get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, + request_kwargs=request_kwargs, ) deployment: Final = await self._select_deployment_async( strategy=strategy, @@ -13881,9 +13893,10 @@ class Router: # if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm ############## Check 'weight' param set for weighted pick ################# return simple_shuffle( - llm_router_instance=self, + resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, + request_kwargs=request_kwargs, ) deployment: Final = self._select_deployment_sync( strategy=strategy, @@ -13951,6 +13964,7 @@ class Router: messages=messages, input=input, specific_deployment=specific_deployment, + request_kwargs=request_kwargs, ) strategy, strategy_selector = self._get_routing_context(model, request_kwargs) @@ -14033,9 +14047,10 @@ class Router: # 6. Apply load balancing strategy if strategy == "simple-shuffle": return simple_shuffle( - llm_router_instance=self, + resolve_model_alias=self._get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, + request_kwargs=request_kwargs, ) deployment: Final = self._select_deployment_sync( strategy=strategy, diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 9deccc9a468..1000f0f479f 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -16,6 +16,8 @@ Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter from __future__ import annotations import asyncio +import hashlib +import json import random import re import time @@ -28,6 +30,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast from pydantic import BaseModel, TypeAdapter, ValidationError, create_model from litellm._logging import verbose_router_logger +from litellm.caching.affinity_cache import claim_affinity_pin from litellm.constants import ( EMPTY_MAPPING, INTERNAL_CALL_ORIGIN_METADATA_KEY, @@ -55,6 +58,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import ( TierSuccessPredictor, resolve_tier_artifact, ) +from litellm.router_utils.pre_call_checks.deployment_affinity_check import DeploymentAffinityCheck from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionImageObject, @@ -1119,10 +1123,10 @@ class _ContextWindowPlacement(NamedTuple): class _SessionAffinityPin(NamedTuple): model: str - tier: ComplexityTier | None + tier: ComplexityTier | str | None -def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None: +def _parse_session_affinity_pin(value: object, active_tiers: tuple[str, ...]) -> _SessionAffinityPin | None: if isinstance(value, str): return _SessionAffinityPin(model=value, tier=None) parts: Final[tuple[object, object] | None] = ( @@ -1137,8 +1141,11 @@ def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None: model, tier_value = parts if not isinstance(model, str): return None - tier: Final = ComplexityTier(tier_value) if isinstance(tier_value, str) else None - return _SessionAffinityPin(model=model, tier=tier) + if tier_value is None: + return _SessionAffinityPin(model=model, tier=None) + if not isinstance(tier_value, str) or tier_value not in active_tiers: + return None + return _SessionAffinityPin(model=model, tier=_built_in_tier_or_none(tier_value) or tier_value) def _session_affinity_cache_value(model: str, tier: ComplexityTier | str | None) -> Mapping[str, str | None]: @@ -1195,6 +1202,10 @@ class ComplexityRouter(CustomLogger): if default_model: self.config.default_model = default_model + self._tier_affinity_config = hashlib.sha256( + self.config.model_dump_json(include=MappingProxyType({"tiers": True, "tier_model_configs": True})).encode() + ).hexdigest() + # Checked here rather than on the config model because the deployment's # complexity_router_default_model arrives outside complexity_router_config and is # applied just above, so a validator on the model would reject a deployment that @@ -2259,6 +2270,51 @@ class ComplexityRouter(CustomLogger): def _tier_pools(self) -> dict[str, list[str]]: return {tier: (models if isinstance(models, list) else [models]) for tier, models in self.config.tiers.items()} + async def _pin_model_for_tier( + self, + tier: ComplexityTier | str, + model: str, + candidates: tuple[str, ...], + request_kwargs: dict[str, object], # mutable-ok: adaptive feedback metadata must follow the selected model + retained_pin: _SessionAffinityPin | None = None, + ) -> str: + if not self._uses_deployment_pin or model not in candidates: + return model + retained_model: Final = ( + retained_pin.model + if retained_pin is not None + and retained_pin.tier is not None + and _tier_name(retained_pin.tier) == _tier_name(tier) + else None + ) + if retained_model is not None and retained_model in candidates: + self._restamp_adaptive_choice(request_kwargs, model, retained_model) + return retained_model + session_id: Final = self._get_session_id_from_request_kwargs(request_kwargs) + if session_id is None: + return model + caller: Final = DeploymentAffinityCheck.get_user_key_from_request_kwargs(request_kwargs) + identity: Final = (self.model_name, self._tier_affinity_config, caller, session_id, _tier_name(tier)) + cache_identity: Final = ( + (*identity, ("replay_fallback", retained_model)) if retained_model is not None else identity + ) + cache_key: Final = ( + "complexity_router_tier_model_affinity:v1:" + + hashlib.sha256(json.dumps(cache_identity).encode()).hexdigest() + ) + winner: Final = await claim_affinity_pin( + self.litellm_router_instance.cache, + cache_key, + MappingProxyType({"model": model}), + self.config.session_affinity_ttl_seconds, + eligible_values=tuple(MappingProxyType({"model": candidate}) for candidate in candidates), + ) + pinned: Final[object] = winner.get("model") if isinstance(winner, Mapping) else None + if not isinstance(pinned, str) or pinned not in candidates: + return model + self._restamp_adaptive_choice(request_kwargs, model, pinned) + return pinned + async def _pick_model_for_tier( self, tier: ComplexityTier | str, @@ -2266,11 +2322,18 @@ class ComplexityRouter(CustomLogger): resolved_messages: list[dict[str, Any]] | None, request_kwargs: dict, allowed_models: tuple[str, ...] | None = None, + retained_pin: _SessionAffinityPin | None = None, ) -> str: if not self.config.plugins: - if allowed_models is not None: - return self._pick_from_tier_value(allowed_models, _tier_name(tier)) - return self.get_model_for_tier(tier) + candidates: Final = ( + allowed_models if allowed_models is not None else tuple(self._tier_pools().get(_tier_name(tier), ())) + ) + selected: Final = ( + self._pick_from_tier_value(allowed_models, _tier_name(tier)) + if allowed_models is not None + else self.get_model_for_tier(tier) + ) + return await self._pin_model_for_tier(tier, selected, candidates, request_kwargs, retained_pin) from litellm.types.router import RoutingContext @@ -2369,6 +2432,40 @@ class ComplexityRouter(CustomLogger): self._adaptive_chosen_model_key = ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY return self.adaptive_router + def _adaptive_candidate_models( + self, + classified_tier: ComplexityTier | str, + hard_floor: ComplexityTier | str | None = None, + hard_ceiling: ComplexityTier | str | None = None, + fit_filter: frozenset[str] | None = None, + ) -> tuple[str, ...]: + pools: Final = self._tier_pools() + candidates: Final = ( + tuple(pools.get(_tier_name(classified_tier), ())) + if self.config.adaptive_eligible == "classified_tier" + else tuple(dict.fromkeys(chain.from_iterable(pools.values()))) + ) + floor: Final = self._active_tier_severity(hard_floor) if hard_floor is not None else None + ceiling: Final = self._active_tier_severity(hard_ceiling) if hard_ceiling is not None else None + return tuple( + model + for model in _allowed(candidates, fit_filter) + if ( + floor is None + or any( + self._active_tier_severity(tier) >= floor + for tier in self._model_tiers.get(model, (classified_tier,)) + ) + ) + and ( + ceiling is None + or any( + self._active_tier_severity(tier) <= ceiling + for tier in self._model_tiers.get(model, (classified_tier,)) + ) + ) + ) + def _soft_floor_pick( self, classified_tier: ComplexityTier | str, @@ -2436,34 +2533,17 @@ class ComplexityRouter(CustomLogger): ], } return chosen_model - if self.config.adaptive_eligible == "classified_tier": - candidates = list(classified_candidates) - if not candidates: - return self._fitting_tier_fallback(classified_tier, fit_filter) - else: - candidates = list(_allowed(tuple(adaptive.config.available_models), fit_filter)) + candidates: Final = self._adaptive_candidate_models(classified_tier, fit_filter=fit_filter) all_costs: Final = [adaptive.model_to_cost.get(m, 0.0) for m in candidates] quality_weight: Final = self.config.adaptive_weights.quality cost_weight: Final = self.config.adaptive_weights.cost penalty_weight: Final = self.config.tier_distance_penalty - floor_severity: Final = self._active_tier_severity(hard_floor) if hard_floor is not None else None - ceiling_severity: Final = self._active_tier_severity(hard_ceiling) if hard_ceiling is not None else None best_model: str | None = None best_score = float("-inf") candidate_scores: Final[list[dict[str, object]]] = [] - for model in candidates: - if floor_severity is not None and all( - self._active_tier_severity(model_tier) < floor_severity - for model_tier in self._model_tiers.get(model, (classified_tier,)) - ): - continue - if ceiling_severity is not None and all( - self._active_tier_severity(model_tier) > ceiling_severity - for model_tier in self._model_tiers.get(model, (classified_tier,)) - ): - continue + for model in self._adaptive_candidate_models(classified_tier, hard_floor, hard_ceiling, fit_filter): cell = adaptive._cells[(request_type, model)] quality_sample = thompson_sample(cell) cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs) @@ -2644,8 +2724,6 @@ class ComplexityRouter(CustomLogger): """Prompt content the resolved message list never carries: the Responses API's `instructions`, the /v1/messages top-level `system` block, and tool definitions. A coding agent's context is dominated by these.""" - import json - instructions: Final = request_kwargs.get("instructions") proxy_request: Final = request_kwargs.get("proxy_server_request") body: Final = proxy_request.get("body") if isinstance(proxy_request, Mapping) else None @@ -2831,19 +2909,21 @@ class ComplexityRouter(CustomLogger): ) return higher_tiers[0] if higher_tiers else tier - def _escalated_pin(self, pinned_model: str) -> str | None: + def _escalated_pin(self, pinned_model: str, tier: ComplexityTier | str | None = None) -> _SessionAffinityPin | None: """Bump a session's pinned model to the next-higher configured tier. Returns None when the pin no longer maps to any configured tier, signalling a full reclassification instead. """ - pinned_tier: Final = self._tier_for_model(pinned_model) + pinned_tier: Final = tier if tier is not None else self._tier_for_model(pinned_model) if pinned_tier is None: return None escalated_tier: Final = self._escalate_tier(pinned_tier) if escalated_tier == pinned_tier: - return pinned_model - return self.get_model_for_tier(escalated_tier) + return _SessionAffinityPin(pinned_model, pinned_tier) + return _SessionAffinityPin( + self.get_model_for_tier(escalated_tier), _built_in_tier_or_none(_tier_name(escalated_tier)) + ) def _vision_verdicts(self, model_name: str) -> tuple[bool | None, ...]: """Declared vision support per deployment serving the name: True, False, or None when @@ -2907,6 +2987,7 @@ class ComplexityRouter(CustomLogger): resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives context_fit: _RequestContextFit | None = None, + retained_pin: _SessionAffinityPin | None = None, ) -> PreRoutingHookResponse: """Replace a routed model that cannot accept this request's image input. @@ -2955,6 +3036,7 @@ class ComplexityRouter(CustomLogger): repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them request_kwargs, allowed_models=tuple(entry for entry in pools.get(capable, ()) if entry in eligible), + retained_pin=retained_pin, ) elif self._modality_default_model_usable(request_kwargs, resolved_messages, eligible): new_tier = None @@ -3098,6 +3180,7 @@ class ComplexityRouter(CustomLogger): resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives context_fit: _RequestContextFit | None = None, + retained_pin: _SessionAffinityPin | None = None, ) -> PreRoutingHookResponse: """Try compatible tier recovery before the default, preserving request policy and fit.""" decision: Final = response.routing_decision @@ -3155,6 +3238,7 @@ class ComplexityRouter(CustomLogger): repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them request_kwargs, allowed_models=live, + retained_pin=retained_pin, ) except ValueError as exc: verbose_router_logger.debug( @@ -3247,8 +3331,13 @@ class ComplexityRouter(CustomLogger): """The adaptive feedback loop reads its chosen-model marker from request metadata; a gate rewrite must move the marker with the model or rewards land on the displaced one.""" metadata: Final = request_kwargs.get("metadata") - if isinstance(metadata, dict) and metadata.get("adaptive_router_chosen_model") == old_model: + if not isinstance(metadata, dict): + return + if metadata.get("adaptive_router_chosen_model") == old_model: metadata["adaptive_router_chosen_model"] = new_model + decision: Final = metadata.get("adaptive_router_decision") + if isinstance(decision, dict) and decision.get("chosen_model") == old_model: + decision["chosen_model"] = new_model def _lexical_tier_override(self, user_message: str) -> KeywordOverride | None: """When keyword_tier_rules match literally, the most-severe matched tier wins. @@ -3561,25 +3650,42 @@ class ComplexityRouter(CustomLogger): if cache_key is not None and pin_replay_allowed: pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key) - pinned_pin: Final = _parse_session_affinity_pin(pinned_value) + pinned_pin: Final = _parse_session_affinity_pin(pinned_value, self.config.tier_names()) if pinned_pin is not None: - routed_model: str | None = pinned_pin.model - pin_escalation_keyword: str | None = None - if self.escalation_keywords: - user_message: Final = ( - _newest_turn_ask(resolved_messages, marker_pairs) if resolved_messages else None + user_message: Final = _newest_turn_ask(resolved_messages, marker_pairs) if resolved_messages else None + pin_escalation_keyword: Final = ( + self._matched_escalation_keyword(user_message) if user_message is not None else None + ) + selected_pin: Final = ( + self._escalated_pin(pinned_pin.model, pinned_pin.tier) + if pin_escalation_keyword is not None + else _SessionAffinityPin( + pinned_pin.model, + pinned_pin.tier if pinned_pin.tier is not None else self._tier_for_model(pinned_pin.model), ) - if user_message is not None: - pin_escalation_keyword = self._matched_escalation_keyword(user_message) - if pin_escalation_keyword is not None: - routed_model = self._escalated_pin(pinned_pin.model) - if routed_model is not None: - escalated: Final = routed_model != pinned_pin.model - resolved_pin_tier: Final = ( - pinned_pin.tier - if not escalated and pinned_pin.tier is not None - else self._tier_for_model(routed_model) + ) + if selected_pin is not None: + escalated: Final = selected_pin.model != pinned_pin.model or ( + pin_escalation_keyword is not None + and pinned_pin.tier is not None + and selected_pin.tier != pinned_pin.tier ) + resolved_pin_tier: Final = selected_pin.tier + session_model: Final = ( + await self._pin_model_for_tier( + resolved_pin_tier, + selected_pin.model, + tuple(self._tier_pools().get(_tier_name(resolved_pin_tier), ())), + request_kwargs, + ) + if escalated and resolved_pin_tier is not None + else selected_pin.model + ) + retained_pin: Final = _SessionAffinityPin(session_model, resolved_pin_tier) + if resolved_pin_tier is not None: + await self._pin_model_for_tier( + resolved_pin_tier, session_model, (session_model,), request_kwargs + ) # The floor outranks the pin because plan mode is a transient state of the # session, not a request to move it: the turns carrying the sentinel route at # the floor, and the stored pin deliberately keeps the session's own model so @@ -3590,16 +3696,28 @@ class ComplexityRouter(CustomLogger): plan_floored: Final = ( pinned_tier is not None and self._apply_plan_mode_floor(pinned_tier) != pinned_tier ) - session_model: Final = routed_model - if plan_floored and pinned_tier is not None: - routed_model = self.get_model_for_tier(self._apply_plan_mode_floor(pinned_tier)) - pin_source_tier: Final = self._tier_for_model(routed_model) + floor_model: Final = ( + await self._pick_model_for_tier( + self._apply_plan_mode_floor(pinned_tier), + messages, + resolved_messages, + request_kwargs, + retained_pin=retained_pin, + ) + if plan_floored and pinned_tier is not None + else session_model + ) + pin_source_tier: Final = ( + self._apply_plan_mode_floor(pinned_tier) + if plan_floored and pinned_tier is not None + else resolved_pin_tier + ) pin_placement: Final = ( await self._context_window_placement( pin_source_tier, resolved_messages, request_kwargs, - pool_override=(routed_model,), + pool_override=(floor_model,), context_fit=context_fit, ) if pin_source_tier is not None @@ -3612,11 +3730,18 @@ class ComplexityRouter(CustomLogger): and _tier_name(pin_placement.tier) != _tier_name(pin_source_tier) else None ) - if pin_placement is not None and pin_context_original_tier is not None: - # The stored pin below keeps the session's own model on purpose. - routed_model = self._pick_from_tier_value( - pin_placement.allowed_models, _tier_name(pin_placement.tier) + routed_model: Final = ( + await self._pick_model_for_tier( + pin_placement.tier, + messages, + resolved_messages, + request_kwargs, + allowed_models=pin_placement.allowed_models, + retained_pin=retained_pin, ) + if pin_placement is not None and pin_context_original_tier is not None + else floor_model + ) # Refresh the TTL on every hit so an active session doesn't lose its # pin mid-conversation just because it outlives the original write. await self.litellm_router_instance.cache.async_set_cache( @@ -3644,7 +3769,7 @@ class ComplexityRouter(CustomLogger): routed_pin_tier: Final = ( pin_placement.tier if pin_placement is not None and pin_context_original_tier is not None - else (self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier) + else pin_source_tier ) session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model) has_original_messages: Final = messages is not None and len(messages) > 0 @@ -3671,12 +3796,14 @@ class ComplexityRouter(CustomLogger): resolved_messages, request_kwargs, context_fit, + retained_pin, ), messages, input, resolved_messages, request_kwargs, context_fit, + retained_pin, ) ) @@ -3961,13 +4088,21 @@ class ComplexityRouter(CustomLogger): housekeeping_ceiling: Final = tier if outcome.cause == "housekeeping" else None # A context-escalated tier becomes the hard floor: a floor the bandit can slide # under is not a floor. - routed_model = self._soft_floor_pick( + adaptive_floor: Final = tier if context_original_tier is not None else plan_floor + adaptive_fit: Final = context_placement.holdable_models if context_placement is not None else None + sampled_model: Final = self._soft_floor_pick( tier, ask, request_kwargs, - hard_floor=tier if context_original_tier is not None else plan_floor, + hard_floor=adaptive_floor, hard_ceiling=housekeeping_ceiling, - fit_filter=context_placement.holdable_models if context_placement is not None else None, + fit_filter=adaptive_fit, + ) + routed_model = await self._pin_model_for_tier( # rebind-ok: reuse the eligible tier winner + tier, + sampled_model, + self._adaptive_candidate_models(tier, adaptive_floor, housekeeping_ceiling, adaptive_fit), + request_kwargs, ) adaptive: Final = self._ensure_adaptive_router() if adaptive is not None: diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 1f1b5a5cc4b..bfc83f8dcab 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1256,20 +1256,16 @@ class ComplexityRouterConfig(BaseModel): deployment_affinity: bool = Field( default=True, description=( - "When True and a session_id is resolvable on the request, pin the deployment chosen " - "inside each routed model group and reuse it whenever the session returns to that " - "group, without pinning which group the session routes to. Independent of " - "session_affinity, which pins the model group instead (and always carries this " - "deployment pin with it): with session_affinity off, " - "every turn is still classified on its own merits while a session that escalates to a " - "stronger tier and comes back still lands on the deployment it used before, which is " - "what keeps a provider prompt cache warm. Pins are held per model group, so switching " - "tiers does not disturb the pin left behind in the previous group. On by default " - "because re-shuffling a conversation across deployments of the same model discards " - "that cache for no benefit; set False to keep every turn load-balanced across the " - "group, which is what a deployment set with tight per-deployment rate limits wants. " - "Inert when no session_id is resolvable, since there is nothing to key a pin on, and " - "suppressed when plugins are configured, for the same reason session_affinity is." + "When True and a client session_id is resolvable, reuse the session's chosen model " + "for each classified tier and its deployment within each model group. With " + "session_affinity off, every turn is still classified: moving to another tier leaves " + "the previous tier's model pin intact for a later return. Pins yield to current " + "candidate, context, modality, and availability constraints. Adaptive selection chooses " + "the initial model from its eligible pool, then reuses that choice per tier. This " + "reduces avoidable provider prompt-cache misses; it does not guarantee cache hits. " + "Set False to select models and load-balance deployments on every turn, unless " + "session_affinity or user_turn classification requires a pin. Inert without a client " + "session_id and suppressed when plugins are configured." ), ) session_affinity_ttl_seconds: int = Field( @@ -1277,7 +1273,7 @@ class ComplexityRouterConfig(BaseModel): gt=0, description=( "TTL for the session affinity pin; refreshed on every cache hit. Bounds both the " - "session_affinity model pin and the deployment_affinity deployment pin, so it measures " + "session_affinity model pin and the deployment_affinity per-tier model and deployment pins, so it measures " "idle time for the session's routing decisions rather than total session length" ), ) diff --git a/litellm/router_strategy/simple_shuffle.py b/litellm/router_strategy/simple_shuffle.py index 860e89cea22..4f2c5e8d933 100644 --- a/litellm/router_strategy/simple_shuffle.py +++ b/litellm/router_strategy/simple_shuffle.py @@ -1,71 +1,67 @@ -""" -Returns a random deployment from the list of healthy deployments. +"""Choose among eligible deployments using request weights, then global metrics.""" -If weights are provided, it will return a deployment based on the weights. - -""" +from __future__ import annotations +import logging import random -from typing import TYPE_CHECKING, Any, Final +from collections.abc import Callable, Mapping, Sequence +from itertools import chain +from typing import Final, TypeVar -from litellm._logging import verbose_router_logger +from litellm.types.router_weights import validate_router_weights -if TYPE_CHECKING: - from litellm.router import Router as _Router +_DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object]) +_ROUTER_LOGGER: Final = logging.getLogger("LiteLLM Router") - LitellmRouter = _Router -else: - LitellmRouter = Any + +def _metric_weight(deployment: Mapping[str, object], metric: str) -> float: + params: Final = deployment.get("litellm_params") + value: Final = params.get(metric) if isinstance(params, Mapping) else None + if value is None: + return 0.0 + if isinstance(value, (int, float)): + return float(value) + raise TypeError(f"Deployment {metric} must be numeric") + + +def _scoped_weights( + deployments: Sequence[Mapping[str, object]], + model: str, + request_kwargs: Mapping[str, object] | None, +) -> tuple[float, ...]: + settings: Final = validate_router_weights((request_kwargs or {}).get("_router_weights")) + model_weights: Final = settings.get(model) if settings is not None else None + if not model_weights: + return () + return tuple( + model_weights.get(str(info.get("id")), 0.0) if isinstance(info, Mapping) else 0.0 + for deployment in deployments + for info in (deployment.get("model_info"),) + ) def simple_shuffle( - llm_router_instance: LitellmRouter, - healthy_deployments: list[Any] | dict[Any, Any], + resolve_model_alias: Callable[[str], str | None], + healthy_deployments: Sequence[_DeploymentT], model: str, -) -> dict: - """ - Returns a random deployment from the list of healthy deployments. - - If weights are provided, it will return a deployment based on the weights. - - If users pass `rpm` or `tpm`, we do a random weighted pick - based on `rpm`/`tpm`. - - Args: - llm_router_instance: LitellmRouter instance - healthy_deployments: List of healthy deployments - model: Model name - - Returns: - Dict: A single healthy deployment - """ - - ############## Check if 'weight' or 'rpm' or 'tpm' param set for a weighted pick ################# - for weight_by in ["weight", "rpm", "tpm"]: - if any(m["litellm_params"].get(weight_by) is not None for m in healthy_deployments): - weights = [m["litellm_params"].get(weight_by, 0) for m in healthy_deployments] - verbose_router_logger.debug("\nweight %s", weights) - total_weight = sum(weights) - if total_weight <= 0: - # All remaining candidates have weight 0 for this metric (e.g. - # after a weighted-failover exclusion left only zero-weight - # backups). Skip to the next metric (rpm/tpm) which may still - # provide a meaningful weighted pick; if none do, we fall - # through to the uniform random pick at the end. - continue - weights = [weight / total_weight for weight in weights] - verbose_router_logger.debug("\n weights %s by %s", weights, weight_by) - # Perform weighted random pick - selected_index = random.choices(range(len(weights)), weights=weights)[0] - verbose_router_logger.debug("\n selected index, %s", selected_index) - deployment = healthy_deployments[selected_index] - verbose_router_logger.info( - "get_available_deployment for model: %s, Selected deployment: %s for model: %s", - model, - llm_router_instance.print_deployment(deployment) or deployment[0], - model, - ) - return deployment or deployment[0] - - ############## No RPM/TPM passed, we do a random pick ################# - item: Final = random.choice(healthy_deployments) - return item or item[0] + request_kwargs: Mapping[str, object] | None, +) -> _DeploymentT: + resolved_model: Final = resolve_model_alias(model) or model + weight_sets: Final = chain( + (_scoped_weights(healthy_deployments, resolved_model, request_kwargs),), + ( + tuple(_metric_weight(deployment, metric) for deployment in healthy_deployments) + for metric in ("weight", "rpm", "tpm") + ), + ) + for weights in weight_sets: + largest = max(weights, default=0.0) + if largest <= 0: + continue + normalized = tuple(weight / largest for weight in weights) + if sum(normalized) <= 0: + continue + selected = random.choices(healthy_deployments, weights=normalized)[0] + _ROUTER_LOGGER.info("Selected deployment for model %s: %s", model, selected.get("model_info")) + return selected + return random.choice(healthy_deployments) diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index c7eb46046ef..3b88ac2eb00 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -13,13 +13,13 @@ where routing to a consistent deployment is still beneficial. """ import hashlib -import json from collections.abc import Mapping, Sequence from typing import Any, Final, cast -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_router_logger +from litellm.caching.affinity_cache import claim_affinity_pin, claim_affinity_pin_in_memory, set_local_affinity_pin from litellm.caching.dual_cache import DualCache from litellm.constants import SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger, Span @@ -28,8 +28,8 @@ from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import CallTypes -class DeploymentAffinityCacheValue(TypedDict): - model_id: str +class DeploymentAffinityCacheValue(TypedDict, closed=True): + model_id: ReadOnly[str] VALID_MODEL_GROUP_AFFINITY_FLAGS: Final = frozenset( @@ -60,19 +60,6 @@ def warn_on_unknown_model_group_affinity_flags(model_group_affinity_config: Mapp ) -_CLAIM_PIN_SCRIPT: Final = """ -local current = redis.call('GET', KEYS[1]) -if current == false then - redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2]) - return ARGV[1] -end -if current == ARGV[1] then - redis.call('EXPIRE', KEYS[1], ARGV[2]) -end -return current -""" - - class DeploymentAffinityCheck(CustomLogger): """ Router deployment affinity callback. @@ -255,34 +242,33 @@ class DeploymentAffinityCheck(CustomLogger): return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{hashed_user_key}:{session_id}" @staticmethod - def _get_session_id_from_metadata_dict(metadata: dict) -> str | None: + def _get_session_id_from_metadata_dict(metadata: Mapping[object, object]) -> str | None: session_id: Final = metadata.get("session_id") if session_id is None or metadata.get(SESSION_ID_GENERATED_METADATA_KEY): return None return str(session_id) @staticmethod - def _iter_metadata_dicts(request_kwargs: dict) -> list[dict]: + def _iter_metadata_dicts(request_kwargs: Mapping[str, object]) -> tuple[Mapping[object, object], ...]: """ Return all metadata dicts available on the request. Depending on the endpoint, Router may populate `metadata` or `litellm_metadata`. Users may also send one or both, so we check both (rather than using `or`). """ - metadata_dicts: Final[list[dict]] = [] - for key in ("litellm_metadata", "metadata"): - md = request_kwargs.get(key) - if isinstance(md, dict): - metadata_dicts.append(md) - return metadata_dicts + return tuple( + cast(Mapping[object, object], metadata) # cast-ok: isinstance proves mapping shape; values remain opaque + for key in ("litellm_metadata", "metadata") + if isinstance(metadata := request_kwargs.get(key), dict) + ) @staticmethod - def _first_metadata_value(metadata_dicts: Sequence[dict], key: str) -> str | None: + def _first_metadata_value(metadata_dicts: Sequence[Mapping[object, object]], key: str) -> str | None: value: Final = next((metadata[key] for metadata in metadata_dicts if metadata.get(key) is not None), None) return None if value is None else str(value) @classmethod - def _get_user_key_from_request_kwargs(cls, request_kwargs: dict) -> str | None: + def get_user_key_from_request_kwargs(cls, request_kwargs: Mapping[str, object]) -> str | None: """ Extract a stable affinity key from request kwargs. @@ -334,74 +320,17 @@ class DeploymentAffinityCheck(CustomLogger): return None def _set_local_pin(self, cache_key: str, value: object, ttl_seconds: int) -> None: - """The one owner of authoritative local pin writes: a plain set keeps a live - key's original expiry (`allow_ttl_override`), so the entry is replaced to make - the TTL real. Every local pin write goes through here so the redis-winner sync - and the pod-local claim can never disagree about expiry again.""" - self.cache.in_memory_cache.delete_cache(cache_key) - self.cache.in_memory_cache.set_cache(cache_key, value, ttl=ttl_seconds) + set_local_affinity_pin(self.cache, cache_key, value, ttl_seconds) async def _claim_pin(self, cache_key: str, pin_value: DeploymentAffinityCacheValue, ttl_seconds: int) -> str | None: - """First-writer-wins pin write: store `pin_value` only when the key is absent and - return the deployment id the key holds afterwards, so a caller learns whether it won - by comparing against its own id, and None when the stored value is one no reader can - interpret. Concurrent claimers converge on the - first write instead of the last. Re-claiming with the stored value refreshes its - TTL, the same keepalive the complexity router's model pin documents: an active - session must not lose its pin mid-conversation just because it outlives the - original write, so the affinity TTL (the Router's - `deployment_affinity_ttl_seconds`, or a pre-routing hook's per-request - `session_affinity_ttl_seconds` override) bounds idle time, not total - session length. On Redis one Lua script does the get-or-set-or-refresh - atomically (same registration seam the rate limiters use) and the in-memory - tier is synchronized to the winner; without Redis, and whenever Redis is - unreachable, the pod-local check-and-set below stands in and is atomic because it - runs synchronously on the event loop. Degrading to a pod-local claim rather than - propagating the fault is what keeps same-pod stickiness through a Redis blip: the - caller only logs this result, so an escaping error would leave the session with no - pin at all and reshuffle every turn for the outage, which is worse than losing - cross-pod agreement. The redis tier is - resolved per call because the proxy attaches it after Router construction - (`Router._update_redis_cache`); the compiled script is cached per event loop - underneath the registration seam. - """ - redis_cache: Final = self.cache.redis_cache - if redis_cache is not None: - try: - claim_script: Final = redis_cache.async_register_script(_CLAIM_PIN_SCRIPT) - raw: Final = await claim_script(keys=(cache_key,), args=(json.dumps(pin_value), int(ttl_seconds))) - decoded: Final = raw.decode("utf-8") if isinstance(raw, bytes) else raw - if not isinstance(decoded, str): - return pin_value["model_id"] - try: - winner: object = json.loads(decoded) - except json.JSONDecodeError: - winner = decoded - self._set_local_pin(cache_key=cache_key, value=winner, ttl_seconds=ttl_seconds) - return self._pinned_model_id(winner) - except Exception as e: # noqa: BLE001 # any Redis/Lua failure degrades to the pod-local claim, never unpins - verbose_router_logger.debug( - "DeploymentAffinityCheck: redis pin claim failed, falling back to pod-local claim. error=%s", e - ) - - return self._claim_pin_in_memory(cache_key=cache_key, pin_value=pin_value, ttl_seconds=ttl_seconds) + winner: Final = await claim_affinity_pin(self.cache, cache_key, pin_value, ttl_seconds) + return self._pinned_model_id(winner) def _claim_pin_in_memory( self, cache_key: str, pin_value: DeploymentAffinityCacheValue, ttl_seconds: int ) -> str | None: - """Pod-local half of the claim, used when no Redis tier is attached and as the - fallback when the Redis claim fails. Mirrors the Lua script exactly, including - the keepalive: re-claiming with the stored value slides the idle window through - `_set_local_pin`. Both branches stay synchronous, hence atomic on the event - loop.""" - existing: Final = self.cache.in_memory_cache.get_cache(cache_key) - if existing is not None: - existing_model_id: Final = self._pinned_model_id(existing) - if existing_model_id == pin_value["model_id"]: - self._set_local_pin(cache_key=cache_key, value=pin_value, ttl_seconds=ttl_seconds) - return existing_model_id - self._set_local_pin(cache_key=cache_key, value=pin_value, ttl_seconds=ttl_seconds) - return pin_value["model_id"] + winner: Final = claim_affinity_pin_in_memory(self.cache, cache_key, pin_value, ttl_seconds) + return self._pinned_model_id(winner) @staticmethod def _find_deployment_by_model_id(healthy_deployments: list[dict], model_id: str) -> dict | None: @@ -465,7 +394,7 @@ class DeploymentAffinityCheck(CustomLogger): enable_session_id or self._get_marker_session_affinity_ttl(request_kwargs=request_kwargs) is not None ) user_key: Final = ( - self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs) + self.get_user_key_from_request_kwargs(request_kwargs=request_kwargs) if (session_affinity_active or enable_user_key) else None ) @@ -580,7 +509,7 @@ class DeploymentAffinityCheck(CustomLogger): return None user_key: Final = ( - self._get_user_key_from_request_kwargs(request_kwargs=kwargs) + self.get_user_key_from_request_kwargs(request_kwargs=kwargs) if (enable_user_key or session_affinity_active) else None ) diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py index cdd70e6baf2..cb5c3089685 100644 --- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py @@ -48,6 +48,7 @@ from litellm.exceptions import ( ServiceUnavailableError, ) from litellm.integrations.custom_logger import CustomLogger, Span +from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.prompt_templates.common_utils import ( encrypted_content_of_block, strip_encrypted_reasoning_from_messages, @@ -215,11 +216,11 @@ class EncryptedContentAffinityCheck(CustomLogger): @staticmethod def _encryption_boundary_key( litellm_params: object, - ) -> tuple | None: + ) -> tuple[object, object] | None: """ - ``(api_base, api_key)`` pair identifying an Azure resource. Two - deployments sharing both are interchangeable for ``encrypted_content`` - follow-ups; Azure rejects content produced by any other resource. + ``(api_base, api_key)`` identifies an upstream encryption boundary. + The values are resolved from the deployment and its named credential + without modifying the deployment. Accepts any object exposing dict-style ``.get(key, default)``: plain dicts (the common case in ``healthy_deployments``) as well as @@ -234,9 +235,25 @@ class EncryptedContentAffinityCheck(CustomLogger): return None api_base: Final = getter("api_base") api_key: Final = getter("api_key") - if not api_base or not api_key: + credential_name: Final = getter("litellm_credential_name") + credential_values: Final[Mapping[str, object] | None] = ( + CredentialAccessor.get_credential_values(credential_name) + if isinstance(credential_name, str) and credential_name + else None + ) + effective_api_base: Final = ( + credential_values.get("api_base") + if credential_values is not None and "api_base" in credential_values + else api_base + ) + effective_api_key: Final = ( + credential_values.get("api_key") + if credential_values is not None and "api_key" in credential_values + else api_key + ) + if not effective_api_base or not effective_api_key: return None - return (api_base, api_key) + return (effective_api_base, effective_api_key) def _find_deployments_on_same_encryption_boundary( self, diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 69cb88bfa2f..ab2f9edea1d 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -552,6 +552,16 @@ class BedrockGuardrailConfigModel(BaseModel): "still rejects is bisected automatically, so this value only trades round trips against " "batch size and cannot fail a request on its own.", ) + contextual_grounding_from_messages: bool = Field( + default=False, + description="ApplyGuardrail: when True, post-call scans of a request with no grounding_source / " + "query content parts send the system and developer messages as the grounding source and " + "the latest user message as the query, so the guardrail's contextual grounding policy can " + "score the response. Bedrock bills contextual grounding units for these scans and rejects " + "queries, sources and responses over its contextual grounding length limits, so leave this " + "off for guardrails without a contextual grounding policy. Default False: plain messages " + "are never sent as grounding context.", + ) class BedrockGuardrailStreamingParams(BaseModel): diff --git a/litellm/types/management_endpoints/prompt_cache_prediction.py b/litellm/types/management_endpoints/prompt_cache_prediction.py new file mode 100644 index 00000000000..3789607b021 --- /dev/null +++ b/litellm/types/management_endpoints/prompt_cache_prediction.py @@ -0,0 +1,67 @@ +from collections.abc import Mapping +from typing import Annotated, Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, JsonValue, StrictInt + +TokenCount: TypeAlias = Annotated[StrictInt, Field(ge=0)] + + +class CacheTokenBuckets(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + uncached_input_tokens: TokenCount = 0 + cache_read_input_tokens: TokenCount = 0 + cache_creation_5m_input_tokens: TokenCount = 0 + cache_creation_1h_input_tokens: TokenCount = 0 + + @property + def total_tokens(self) -> int: + return ( + self.uncached_input_tokens + + self.cache_read_input_tokens + + self.cache_creation_5m_input_tokens + + self.cache_creation_1h_input_tokens + ) + + +class CacheEvidence(BaseModel): + model_config = ConfigDict(frozen=True) + + observed_at: float + expires_at: float + source: Literal["provider_usage"] = "provider_usage" + confidence: Literal["observed"] = "observed" + + +class CacheCostScenario(BaseModel): + tokens: CacheTokenBuckets + input_cost: float + + +class CachePredictionArm(BaseModel): + deployment_id: str + model: str | None = None + cache_state: Literal["warm", "partial", "stale", "unknown", "disabled"] = "unknown" + reason: str | None = None + estimate: CacheCostScenario | None = None + cold: CacheCostScenario | None = None + warm: CacheCostScenario | None = None + evidence: CacheEvidence | None = None + token_count_source: Literal["anthropic_count_tokens"] | None = None + + +class CachePredictionRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + current_deployment_id: str = Field(min_length=1, max_length=256) + candidate_deployment_id: str = Field(min_length=1, max_length=256) + request: Mapping[str, JsonValue] + + +class CachePredictionResponse(BaseModel): + stay: CachePredictionArm + switch: CachePredictionArm + switch_delta: float | None + cache_rebuild_penalty: float | None + pricing_basis: Literal["input_before_discounts_and_margins"] = "input_before_discounts_and_margins" + cache_guarantee: Literal[False] = False diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 4a868c48352..44e2cc2404f 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -1,4 +1,5 @@ -from typing import Any, Final, Literal +from collections.abc import Mapping, Sequence +from typing import Any, Final, Literal, cast # noqa: TID251 # JSON chat rows have no typed constructor across roles from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict @@ -158,12 +159,21 @@ def coerce_stream_holdback_value(value: Any) -> int: return 0 +def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None: + if not isinstance(value, list): + return None + if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value): + return None + return cast("Sequence[AllMessageValues]", value) # cast-ok: JSON rows checked for a role, the same trust texts get + + class GenericGuardrailAPIResponse: """Response model for the Generic Guardrail API""" texts: list[str] | None images: list[str] | None tools: list[GuardrailToolParam] | None + structured_messages: Sequence[AllMessageValues] | None action: str blocked_reason: str | None stream_holdback_chars: list[int] | None @@ -176,12 +186,14 @@ class GenericGuardrailAPIResponse: images: list[str] | None = None, tools: list[GuardrailToolParam] | None = None, stream_holdback_chars: list[int] | None = None, + structured_messages: Sequence[AllMessageValues] | None = None, ) -> None: self.action = action self.blocked_reason = blocked_reason self.texts = texts self.images = images self.tools = tools + self.structured_messages = structured_messages # Number of trailing chars, indexed the same as ``texts``, that the # framework must withhold from streaming emission until the next # processing round (word-boundary safety for text transformations). @@ -200,4 +212,5 @@ class GenericGuardrailAPIResponse: images=data.get("images"), tools=data.get("tools"), stream_holdback_chars=stream_holdback_chars, + structured_messages=structured_messages_from_response(data.get("structured_messages")), ) diff --git a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py index 2b41893663d..43e3899d523 100644 --- a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py @@ -1,4 +1,4 @@ -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import Any, Final, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator @@ -6,6 +6,7 @@ from typing_extensions import ReadOnly, TypedDict from litellm.proxy._types import ( LiteLLM_UserTableWithKeyCount, + NewUserRequest, UpdateUserRequest, UpdateUserRequestNoUserIDorEmail, ) @@ -13,6 +14,8 @@ from litellm.types.proxy.management_endpoints.management_v1 import ResourceRespo MAX_BULK_DELETE_USERS: Final = 500 +MAX_BULK_NEW_USERS: Final = 500 + class InsensitiveContains(TypedDict): contains: ReadOnly[str] @@ -108,3 +111,50 @@ class UserDeleteResult(BaseModel): class BulkDeleteUsersResponse(ResourceResponse[tuple[UserDeleteResult, ...]]): """`{data: [...]}` with one `UserDeleteResult` per requested user, in request order.""" + + +class BulkNewUserItem(NewUserRequest): + """One row of `POST /management/v1/users/bulk`: the `/user/new` body, with keys opt-in and invite emails + unsupported. Unknown fields are rejected, as on every `/management/v1` request body.""" + + model_config = ConfigDict(extra="forbid", protected_namespaces=()) + + auto_create_key: bool = False + + @field_validator("send_invite_email") + @classmethod + def reject_invite_email(cls, value: bool | None) -> bool | None: + if value: + raise ValueError("send_invite_email is not supported on /management/v1/users/bulk; invite users separately") + return value + + +class BulkNewUserRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + users: Sequence[BulkNewUserItem] = Field(min_length=1, max_length=MAX_BULK_NEW_USERS) + + +class UserCreateResult(BaseModel): + """Outcome for one row of `POST /management/v1/users/bulk`. `teams` lists the teams the user was actually + added to.""" + + user_id: str | None = None + user_email: str | None = None + success: bool + teams: tuple[str, ...] | None = None + key: str | None = None + error: str | None = None + + +class BulkNewUserMeta(BaseModel): + total_requested: int + created: int + failed: int + + +class BulkNewUserResponse(BaseModel): + """`data` holds one result per input row, in input order.""" + + data: tuple[UserCreateResult, ...] + meta: BulkNewUserMeta diff --git a/litellm/types/proxy/management_endpoints/key_management_endpoints.py b/litellm/types/proxy/management_endpoints/key_management_endpoints.py index 9fb5bea81e3..63bbaa5ba4e 100644 --- a/litellm/types/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/key_management_endpoints.py @@ -1,9 +1,12 @@ from datetime import datetime -from typing import Any, Final, Literal +from typing import Any, Final, Literal, TypeAlias from pydantic import BaseModel, ConfigDict, model_validator from typing_extensions import ReadOnly, TypedDict +from litellm.models.verification_token import LiteLLM_VerificationToken +from litellm.proxy._types import GenerateKeyRequest, RegenerateKeyRequest, UpdateKeyRequest +from litellm.types.llms.base import LiteLLMPydanticObjectBase from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains @@ -123,3 +126,24 @@ class BulkUpdateTeamKeysRequest(BaseModel): if not has_key_ids and not self.all_keys_in_team: raise ValueError("Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`.") return self + + +CustomKeyPolicyOperation: TypeAlias = Literal["generate", "update", "regenerate"] + + +class CustomKeyPolicyRequest(LiteLLMPydanticObjectBase): + """What `general_settings.custom_key_policy` receives. + + `effective_key` is the verification token row as it will be written: the existing row overlaid with the + requested changes, with `duration` resolved to `expires` and `budget_duration` to `budget_reset_at`. Values the + proxy fills in after the policy stay at their defaults: `token`, `key_name`, `created_by`, `updated_by` and the + soft-budget `budget_id` on generate, the rotated token on regenerate, and the `object_permission` relation on + every operation (`object_permission_id` is set; read `request.object_permission` for the requested change). + """ + + model_config = ConfigDict(protected_namespaces=(), frozen=True) + + operation: CustomKeyPolicyOperation + existing_key: LiteLLM_VerificationToken | None + effective_key: LiteLLM_VerificationToken + request: GenerateKeyRequest | UpdateKeyRequest | RegenerateKeyRequest diff --git a/litellm/types/router.py b/litellm/types/router.py index 0aefc07ae4b..49109e4fbbe 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -15,6 +15,7 @@ from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_c from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.litellm_core_utils.core_helpers import normalize_drop_params +from litellm.types.router_weights import RouterWeights if TYPE_CHECKING: from litellm.router import Router @@ -146,6 +147,7 @@ class UpdateRouterConfig(BaseModel): context_window_fallbacks: list[dict] | None = None model_group_alias: dict[str, str | dict] | None = {} enable_tag_filtering: bool | None = None + weights: RouterWeights | None = None tag_routing_prefix: str | None = None optional_pre_call_checks: OptionalPreCallChecks | None = None diff --git a/litellm/types/router_weights.py b/litellm/types/router_weights.py new file mode 100644 index 00000000000..fa156661564 --- /dev/null +++ b/litellm/types/router_weights.py @@ -0,0 +1,30 @@ +from collections.abc import Mapping +from typing import Annotated, Final + +from pydantic import AfterValidator, Field, TypeAdapter + + +def _validate_positive_router_weights(weights: Mapping[str, Mapping[str, float]]) -> Mapping[str, Mapping[str, float]]: + if any(group and not any(weight > 0 for weight in group.values()) for group in weights.values()): + raise ValueError("Each nonempty weights group must contain at least one positive weight") + return weights + + +RouterWeightIdentifier = Annotated[str, Field(strict=True, min_length=1, pattern=r"\S")] +RouterWeight = Annotated[float, Field(strict=True, ge=0, allow_inf_nan=False)] +RouterWeights = Annotated[ + dict[RouterWeightIdentifier, dict[RouterWeightIdentifier, RouterWeight]], + AfterValidator(_validate_positive_router_weights), +] +_ROUTER_WEIGHTS_ADAPTER: Final[TypeAdapter[RouterWeights | None]] = TypeAdapter(RouterWeights | None) +_ROUTER_SETTINGS_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) + + +def validate_router_weights(value: object) -> RouterWeights | None: + return _ROUTER_WEIGHTS_ADAPTER.validate_python(value) + + +def validate_router_settings_dict(value: object) -> dict[str, object]: + settings: Final = _ROUTER_SETTINGS_DICT_ADAPTER.validate_python(value) + validate_router_weights(settings.get("weights")) + return settings diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 1d73542c9bb..8bab7c349ff 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -283,8 +283,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): input_cost_per_video_token: float | None # for gemini omni models with video input input_cost_per_audio_per_second: float | None # only for vertex ai models input_cost_per_video_per_second: float | None # only for vertex ai models + input_cost_per_audio_token_batches: ReadOnly[float | None] + input_cost_per_image_token_batches: ReadOnly[float | None] input_cost_per_second: float | None # for OpenAI Speech models input_cost_per_token_batches: float | None + input_cost_per_video_token_batches: ReadOnly[float | None] output_cost_per_token_batches: float | None output_cost_per_token: Required[float | None] output_cost_per_token_flex: float | None # OpenAI flex service tier pricing @@ -3583,7 +3586,10 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): input_cost_per_video_per_second_above_128k_tokens: float | None = None input_cost_per_video_per_second_above_15s_interval: float | None = None input_cost_per_video_per_second_above_8s_interval: float | None = None + input_cost_per_audio_token_batches: float | None = None + input_cost_per_image_token_batches: float | None = None input_cost_per_token_batches: float | None = None + input_cost_per_video_token_batches: float | None = None output_cost_per_token_batches: float | None = None output_cost_per_token_flex: float | None = None output_cost_per_token_priority: float | None = None @@ -3761,11 +3767,22 @@ all_litellm_params = ( "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", diff --git a/litellm/utils.py b/litellm/utils.py index 5a0c14e5614..18df5e2abf7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1208,30 +1208,13 @@ def _dispatch_success_logging( is_litellm_internal_call: bool, ) -> None: if not is_litellm_internal_call: - if getattr(logging_obj, "_defer_async_logging", False): - - def _enqueue_deferred_logging() -> None: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging - else: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) + _schedule_async_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) logging_obj.handle_sync_success_callbacks_for_async_calls( result=result, @@ -1240,6 +1223,43 @@ def _dispatch_success_logging( ) +def _schedule_async_success_logging( + logging_obj: LiteLLMLoggingObject, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + is_completion_with_fallbacks: bool, +) -> None: + """Fire the async success log for ``result`` now, or park it on the logging object while + the proxy defers logging past its post-call guardrails. + + Nested @client wrappers (Anthropic Messages over the chat adapter, chat over the Responses + bridge) each exit through here with the same logging object and their own shape of the same + response. The immediate path already logs one request once, since the first task marks + ``has_logged_async_success`` and the later ones skip. The deferred slot keeps the same + first-wins rule: the innermost wrapper's provider-shaped result is the one the spend log + reads usage from, and a later wrapper never swaps in its client-shaped translation. + """ + + def _enqueue_async_logging() -> None: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + if not getattr(logging_obj, "_defer_async_logging", False): + _enqueue_async_logging() + return + if getattr(logging_obj, "_enqueue_deferred_logging", None) is not None: + return + logging_obj._enqueue_deferred_logging = _enqueue_async_logging + + async def _client_async_logging_helper( logging_obj: LiteLLMLoggingObject, result, @@ -5923,10 +5943,13 @@ def _get_model_info_helper( input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None), input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None), input_cost_per_video_token=_model_info.get("input_cost_per_video_token", None), + input_cost_per_audio_token_batches=_model_info.get("input_cost_per_audio_token_batches", None), + input_cost_per_image_token_batches=_model_info.get("input_cost_per_image_token_batches", None), input_cost_per_image=_model_info.get("input_cost_per_image", None), input_cost_per_audio_per_second=_model_info.get("input_cost_per_audio_per_second", None), input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None), input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"), + input_cost_per_video_token_batches=_model_info.get("input_cost_per_video_token_batches", None), output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"), output_cost_per_token=_output_cost_per_token, output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None), @@ -6260,7 +6283,7 @@ def function_to_dict(input_function) -> dict: "enum": param_enum, } - parameters[param_name] = dict([(k, v) for k, v in param_dict.items() if isinstance(v, str)]) + parameters[param_name] = {k: v for k, v in param_dict.items() if isinstance(v, str)} # Check if the parameter has no default value (i.e., it's required) if param.default == param.empty: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 685e63dcc56..9f91cf82f41 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -11216,12 +11216,15 @@ "babbage-002": { "deprecation_date": "2026-09-28", "input_cost_per_token": 4e-07, + "input_cost_per_token_batches": 2e-07, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 4e-07 + "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "bedrock/*/1-month-commitment/cohere.command-light-text-v14": { "input_cost_per_second": 0.001902, @@ -13286,7 +13289,9 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 0.0001, + "source": "https://developers.openai.com/api/docs/pricing" }, "claude-haiku-4-5-20251001": { "deprecation_date": "2026-10-15", @@ -13334,7 +13339,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-3-7-sonnet-20250219": { "cache_creation_input_token_cost": 3.75e-06, @@ -13493,7 +13499,8 @@ "supports_native_structured_output": true, "supports_tool_choice": true, "supports_vision": true, - "prompt_cache_min_tokens": 1024 + "prompt_cache_min_tokens": 1024, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-5-20250929": { "deprecation_date": "2026-09-29", @@ -13567,7 +13574,7 @@ }, "supports_output_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-6": { "deprecation_date": "2027-02-17", @@ -13603,7 +13610,8 @@ "prompt_cache_min_tokens": 1024, "provider_specific_entry": { "us": 1.1 - } + }, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, @@ -13780,7 +13788,8 @@ "supports_tool_choice": true, "supports_vision": true, "supports_output_config": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6": { "deprecation_date": "2027-02-05", @@ -13817,7 +13826,8 @@ "supports_output_config": true, "supports_max_reasoning_effort": true, "supports_speed": true, - "prompt_cache_min_tokens": 4096 + "prompt_cache_min_tokens": 4096, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6-20260205": { "deprecation_date": "2027-02-05", @@ -13892,7 +13902,8 @@ }, "supports_output_config": true, "supports_speed": true, - "prompt_cache_min_tokens": 2048 + "prompt_cache_min_tokens": 2048, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-7-20260416": { "deprecation_date": "2027-04-16", @@ -13970,7 +13981,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-fable-5-1": { "deprecation_date": "2027-09-01", @@ -14011,7 +14022,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://platform.claude.com/docs/en/models/fable-5-1/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5": { "deprecation_date": "2027-07-24", @@ -14052,7 +14063,7 @@ "supports_output_config": true, "supports_speed": true, "prompt_cache_min_tokens": 512, - "source": "https://docs.anthropic.com/en/docs/about-claude/models/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-8": { "deprecation_date": "2027-05-28", @@ -14092,7 +14103,8 @@ }, "supports_output_config": true, "supports_speed": true, - "prompt_cache_min_tokens": 1024 + "prompt_cache_min_tokens": 1024, + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-sonnet-4-20250514": { "deprecation_date": "2026-06-15", @@ -19361,12 +19373,15 @@ "davinci-002": { "deprecation_date": "2026-09-28", "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 2e-06 + "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "deepgram/base": { "input_cost_per_second": 0.00020833, @@ -22413,15 +22428,18 @@ "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.45e-07, - "input_cost_per_token": 1.74e-06, + "cache_read_input_token_cost": 6e-07, + "cache_read_input_token_cost_priority": 6e-07, + "input_cost_per_token": 1.2e-06, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 3.48e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22430,14 +22448,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 4.4e-08, + "cache_read_input_token_cost_priority": 5.5e-08, "input_cost_per_token": 1.32e-06, + "input_cost_per_token_priority": 1.65e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 3.96e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 4.95e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22532,14 +22553,17 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p2": { "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost_priority": 1.75e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22548,14 +22572,17 @@ }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { "cache_read_input_token_cost": 1.5e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.8e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 7.2e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22634,14 +22661,17 @@ }, "fireworks_ai/accounts/fireworks/models/kimi-k2p6": { "cache_read_input_token_cost": 1.6e-07, + "cache_read_input_token_cost_priority": 2.2e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22650,14 +22680,17 @@ }, "fireworks_ai/accounts/fireworks/models/kimi-k2p7-code": { "cache_read_input_token_cost": 1.9e-07, + "cache_read_input_token_cost_priority": 2.85e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.425e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22783,14 +22816,17 @@ }, "fireworks_ai/accounts/fireworks/models/minimax-m2p7": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 6e-07, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 196608, "max_output_tokens": 196608, "max_tokens": 196608, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22799,14 +22835,17 @@ }, "fireworks_ai/accounts/fireworks/models/minimax-m3": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 9e-08, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 512000, "max_output_tokens": 512000, "max_tokens": 512000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.8e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22882,15 +22921,18 @@ "supports_vision": false }, "fireworks_ai/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.45e-07, - "input_cost_per_token": 1.74e-06, + "cache_read_input_token_cost": 6e-07, + "cache_read_input_token_cost_priority": 6e-07, + "input_cost_per_token": 1.2e-06, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 3.48e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22946,14 +22988,17 @@ }, "fireworks_ai/glm-5p2": { "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost_priority": 1.75e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -22962,14 +23007,17 @@ }, "fireworks_ai/gpt-oss-120b": { "cache_read_input_token_cost": 1.5e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.8e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 7.2e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23008,14 +23056,17 @@ }, "fireworks_ai/kimi-k2p6": { "cache_read_input_token_cost": 1.6e-07, + "cache_read_input_token_cost_priority": 2.2e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23040,14 +23091,17 @@ }, "fireworks_ai/kimi-k2p7-code": { "cache_read_input_token_cost": 1.9e-07, + "cache_read_input_token_cost_priority": 2.85e-07, "input_cost_per_token": 9.5e-07, + "input_cost_per_token_priority": 1.425e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23086,14 +23140,17 @@ }, "fireworks_ai/minimax-m2p7": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 6e-07, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 1.2e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 196608, "max_output_tokens": 196608, "max_tokens": 196608, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.2e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23102,14 +23159,17 @@ }, "fireworks_ai/minimax-m3": { "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_priority": 9e-08, "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 512000, "max_output_tokens": 512000, "max_tokens": 512000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 1.8e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23125,7 +23185,7 @@ "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -23374,26 +23434,28 @@ "ft:babbage-002": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1.6e-06, - "input_cost_per_token_batches": 2e-07, + "input_cost_per_token_batches": 8e-07, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.6e-06, - "output_cost_per_token_batches": 2e-07 + "output_cost_per_token_batches": 9e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "ft:davinci-002": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1.2e-05, - "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_batches": 6e-06, "litellm_provider": "text-completion-openai", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", "output_cost_per_token": 1.2e-05, - "output_cost_per_token_batches": 1e-06 + "output_cost_per_token_batches": 6e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "ft:gpt-3.5-turbo": { "deprecation_date": "2026-10-23", @@ -23406,6 +23468,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_batches": 3e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_system_messages": true, "supports_tool_choice": true }, @@ -23462,14 +23525,15 @@ "ft:gpt-4o-2024-08-06": { "cache_read_input_token_cost": 1.875e-06, "input_cost_per_token": 3.75e-06, - "input_cost_per_token_batches": 1.875e-06, + "input_cost_per_token_batches": 2.225e-06, "litellm_provider": "openai", "max_input_tokens": 128000, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-05, - "output_cost_per_token_batches": 7.5e-06, + "output_cost_per_token_batches": 1.25e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -23507,6 +23571,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "output_cost_per_token_batches": 6e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -23526,6 +23591,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23544,6 +23610,7 @@ "mode": "chat", "output_cost_per_token": 3.2e-06, "output_cost_per_token_batches": 1.6e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23563,6 +23630,7 @@ "mode": "chat", "output_cost_per_token": 8e-07, "output_cost_per_token_batches": 4e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -23582,6 +23650,7 @@ "mode": "chat", "output_cost_per_token": 1.6e-05, "output_cost_per_token_batches": 8e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -23592,15 +23661,18 @@ "gemini-2.0-flash": { "cache_read_input_token_cost": 2.5e-08, "deprecation_date": "2026-06-01", - "input_cost_per_audio_token": 7e-07, - "input_cost_per_token": 1e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_character": 3.75e-08, + "input_cost_per_token": 1.5e-07, + "input_cost_per_token_batches": 7.5e-08, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 4e-07, - "source": "https://ai.google.dev/pricing#2_0flash", + "output_cost_per_token": 6e-07, + "output_cost_per_token_batches": 3e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image", @@ -23669,13 +23741,16 @@ "cache_read_input_token_cost": 1.875e-08, "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 7.5e-08, + "input_cost_per_character": 1.875e-08, "input_cost_per_token": 7.5e-08, + "input_cost_per_token_batches": 3.75e-08, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", + "output_cost_per_token_batches": 1.5e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image", @@ -23736,8 +23811,11 @@ } }, "gemini-2.5-flash": { + "cache_read_input_audio_token_cost": 1e-07, "deprecation_date": "2026-10-20", "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 3e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, "litellm_provider": "vertex_ai-language-models", @@ -23747,7 +23825,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23780,6 +23858,12 @@ "search_context_size_high": 0.035 }, "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, "supports_image_size": false }, "gemini-2.5-flash-image": { @@ -23787,6 +23871,9 @@ "cache_read_input_token_cost": 3e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 32768, "max_output_tokens": 32768, @@ -23796,8 +23883,10 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, "rpm": 100000, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23828,10 +23917,19 @@ "supports_image_size": false }, "gemini-3-pro-image": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 1e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 3.6e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -23840,8 +23938,12 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3-pro-image", + "output_cost_per_token_flex": 6e-06, + "output_cost_per_token_priority": 2.16e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23909,9 +24011,13 @@ "web_search_billing_unit": "per_query" }, "gemini-3.1-flash-image": { + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 2.5e-08, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -23920,7 +24026,9 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -23987,9 +24095,11 @@ }, "gemini-3.1-flash-lite-image": { "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_flex": 1.25e-08, "input_cost_per_image": 0.00028, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 4096, @@ -23999,6 +24109,7 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_token": 1.5e-06, "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -24073,6 +24184,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "gemini-3.1-flash-lite": { + "cache_read_input_audio_token_cost": 5e-08, "deprecation_date": "2027-05-07", "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, @@ -24092,7 +24204,7 @@ "output_cost_per_token_batches": 7.5e-07, "output_cost_per_token_flex": 7.5e-07, "output_cost_per_token_priority": 2.7e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24134,7 +24246,7 @@ "deprecation_date": "2027-07-21", "cache_read_input_token_cost": 3e-08, "cache_read_input_token_cost_flex": 1.5e-08, - "cache_read_input_token_cost_priority": 5e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, "input_cost_per_token_flex": 1.5e-07, @@ -24149,7 +24261,7 @@ "output_cost_per_token_batches": 1.25e-06, "output_cost_per_token_flex": 1.25e-06, "output_cost_per_token_priority": 4.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24188,6 +24300,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "deep-research-pro-preview-12-2025": { + "cache_read_input_token_cost": 2e-07, "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -24200,7 +24313,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24222,8 +24335,11 @@ "supports_web_search": true }, "gemini-2.5-flash-lite": { + "cache_read_input_audio_token_cost": 3e-08, "deprecation_date": "2026-10-20", "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_flex": 1e-08, + "cache_read_input_token_cost_priority": 1.8e-08, "input_cost_per_audio_token": 3e-07, "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-language-models", @@ -24233,7 +24349,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24266,6 +24382,12 @@ "search_context_size_high": 0.035 }, "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_batches": 5e-08, + "input_cost_per_token_flex": 5e-08, + "input_cost_per_token_priority": 1.8e-07, + "output_cost_per_token_batches": 2e-07, + "output_cost_per_token_flex": 2e-07, + "output_cost_per_token_priority": 7.2e-07, "supports_image_size": false }, "gemini-2.5-flash-lite-preview-09-2025": { @@ -24417,7 +24539,7 @@ "mode": "realtime", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_token": 2e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/vertex_ai/live" ], @@ -24448,7 +24570,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "gemini_native_audio": true + "gemini_native_audio": true, + "input_cost_per_image_token": 3e-06 }, "gemini/gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -24548,6 +24671,9 @@ "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 4.5e-07, + "cache_read_input_token_cost_flex": 1.25e-07, + "cache_read_input_token_cost_priority": 2.25e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, "litellm_provider": "vertex_ai-language-models", @@ -24557,7 +24683,7 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions" @@ -24587,7 +24713,15 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "google_maps_grounding_cost_per_query": 0.025 + "google_maps_grounding_cost_per_query": 0.025, + "input_cost_per_token_above_200k_tokens_priority": 4.5e-06, + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "input_cost_per_token_priority": 2.25e-06, + "output_cost_per_token_above_200k_tokens_priority": 2.7e-05, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "output_cost_per_token_priority": 1.8e-05 }, "gemini-3-pro-preview": { "deprecation_date": "2026-03-26", @@ -24662,7 +24796,7 @@ "output_cost_per_token_above_200k_tokens": 1.8e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_image": 0.00012, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24696,13 +24830,16 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_low": 0.014, "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_flex": 1e-06, + "output_cost_per_token_flex": 6e-06 }, "gemini-3.1-pro-preview-customtools": { "prompt_cache_min_tokens": 4096, @@ -24813,7 +24950,9 @@ "web_search_billing_unit": "per_query" }, "vertex_ai/gemini-3-flash-preview": { + "cache_read_input_audio_token_cost": 1e-07, "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, "input_cost_per_token": 5e-07, "input_cost_per_audio_token": 1e-06, "litellm_provider": "vertex_ai", @@ -24822,7 +24961,7 @@ "max_tokens": 65535, "mode": "chat", "output_cost_per_token": 3e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24859,7 +24998,11 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06 }, "vertex_ai/gemini-3.5-flash": { "prompt_cache_min_tokens": 4096, @@ -24875,7 +25018,7 @@ "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24938,7 +25081,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -24995,7 +25138,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25052,7 +25195,7 @@ "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25109,7 +25252,7 @@ "output_cost_per_token_above_200k_tokens": 1.8e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_image": 0.00012, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -25143,13 +25286,16 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "cache_read_input_token_cost_priority": 3.6e-07, "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_low": 0.014, "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_flex": 1e-06, + "output_cost_per_token_flex": 6e-06 }, "vertex_ai/gemini-3.1-pro-preview-customtools": { "prompt_cache_min_tokens": 4096, @@ -25214,13 +25360,15 @@ "cache_read_input_token_cost": 1.25e-07, "input_cost_per_audio_token": 7e-07, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 1048576, "max_output_tokens": 65535, "max_tokens": 65535, "mode": "chat", + "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2e-05, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text" ], @@ -25427,7 +25575,7 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/computer-use", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_modalities": [ "text", "image" @@ -25453,10 +25601,14 @@ "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models" }, "gemini-embedding-2-preview": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25467,25 +25619,33 @@ "uses_embed_content": true }, "gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai-embedding-models", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, "vertex_ai/gemini-embedding-2-preview": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25497,17 +25657,21 @@ "uses_embed_content": true }, "vertex_ai/gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, @@ -25539,10 +25703,14 @@ }, "gemini/gemini-embedding-2-preview": { "deprecation_date": "2026-08-10", - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "gemini", "max_input_tokens": 8192, "max_tokens": 8192, @@ -25555,10 +25723,14 @@ "tpm": 10000000 }, "gemini/gemini-embedding-2": { - "input_cost_per_audio_per_second": 0.00016, - "input_cost_per_image": 0.00012, + "input_cost_per_audio_token": 6.5e-06, + "input_cost_per_audio_token_batches": 3.25e-06, + "input_cost_per_image_token": 4.5e-07, + "input_cost_per_image_token_batches": 2.25e-07, "input_cost_per_token": 2e-07, - "input_cost_per_video_per_second": 0.00079, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_video_token": 1.2e-05, + "input_cost_per_video_token_batches": 6e-06, "litellm_provider": "gemini", "max_input_tokens": 8192, "max_tokens": 8192, @@ -27141,7 +27313,9 @@ "google_maps_grounding_cost_per_query": 0.014 }, "gemini-3-flash-preview": { + "cache_read_input_audio_token_cost": 1e-07, "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 5e-07, "litellm_provider": "vertex_ai-language-models", @@ -27151,7 +27325,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 3e-06, "output_cost_per_token": 3e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27189,7 +27363,11 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "google_maps_grounding_cost_per_query": 0.014 + "google_maps_grounding_cost_per_query": 0.014, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06 }, "gemini-omni-flash-preview": { "input_cost_per_audio_token": 1.5e-06, @@ -27202,7 +27380,7 @@ "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, "output_cost_per_video_token": 1.75e-05, - "source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/gemini/omni-flash-preview", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions" ], @@ -27235,7 +27413,7 @@ "mode": "chat", "output_cost_per_reasoning_token": 9e-06, "output_cost_per_token": 9e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27298,7 +27476,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27355,7 +27533,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -27412,7 +27590,7 @@ "output_cost_per_token": 3.75e-06, "output_cost_per_token_batches": 1.875e-06, "output_cost_per_token_flex": 1.875e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -28870,6 +29048,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_system_messages": true, @@ -28878,12 +29057,15 @@ "gpt-3.5-turbo-0125": { "deprecation_date": "2026-10-23", "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -28893,12 +29075,15 @@ "gpt-3.5-turbo-1106": { "deprecation_date": "2026-09-28", "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "openai", "max_input_tokens": 16385, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -28926,7 +29111,8 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "completion", - "output_cost_per_token": 2e-06 + "output_cost_per_token": 2e-06, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-3.5-turbo-instruct-0914": { "input_cost_per_token": 1.5e-06, @@ -28981,12 +29167,15 @@ "gpt-4-0613": { "deprecation_date": "2026-10-23", "input_cost_per_token": 3e-05, + "input_cost_per_token_batches": 1.5e-05, "litellm_provider": "openai", "max_input_tokens": 8192, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_system_messages": true, @@ -29027,12 +29216,15 @@ "gpt-4-turbo-2024-04-09": { "deprecation_date": "2026-10-23", "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "openai", "max_input_tokens": 128000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 3e-05, + "output_cost_per_token_batches": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29076,6 +29268,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29118,6 +29311,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29160,6 +29354,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29202,6 +29397,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29240,6 +29436,7 @@ "output_cost_per_token": 4e-07, "output_cost_per_token_batches": 2e-07, "output_cost_per_token_priority": 8e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29277,6 +29474,7 @@ "output_cost_per_token": 4e-07, "output_cost_per_token_priority": 8e-07, "output_cost_per_token_batches": 2e-07, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -29313,6 +29511,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_batches": 5e-06, "output_cost_per_token_priority": 1.7e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29335,6 +29534,7 @@ "output_cost_per_token": 1.5e-05, "output_cost_per_token_batches": 7.5e-06, "output_cost_per_token_priority": 2.625e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29357,6 +29557,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_priority": 1.7e-05, "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29380,6 +29581,7 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_priority": 1.7e-05, "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29454,6 +29656,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29490,6 +29693,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions" ], @@ -29524,6 +29728,7 @@ "mode": "chat", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29561,6 +29766,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29598,6 +29804,7 @@ "mode": "chat", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses", @@ -29659,7 +29866,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": false, - "deprecation_date": "2027-01-20" + "deprecation_date": "2027-01-20", + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini": { "cache_read_input_token_cost": 7.5e-08, @@ -29687,7 +29895,8 @@ "search_context_size_high": 0.025, "search_context_size_low": 0.025, "search_context_size_medium": 0.025 - } + }, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini-2024-07-18": { "cache_read_input_token_cost": 7.5e-08, @@ -29708,6 +29917,7 @@ "search_context_size_low": 0.025, "search_context_size_medium": 0.025 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -29856,15 +30066,18 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 5e-05, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-4o-mini-tts": { - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -29996,7 +30209,9 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "input_cost_per_second": 0.0001, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-image-1.5": { "cache_read_input_token_cost": 1.25e-06, @@ -30006,7 +30221,10 @@ "mode": "image_generation", "output_cost_per_token": 1e-05, "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations" ], @@ -30021,7 +30239,10 @@ "mode": "image_generation", "output_cost_per_token": 1e-05, "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations" ], @@ -30034,7 +30255,9 @@ "litellm_provider": "openai", "mode": "image_generation", "input_cost_per_image_token": 8e-06, + "input_cost_per_token_batches": 2.5e-06, "output_cost_per_image_token": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -30481,6 +30704,7 @@ "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -30489,6 +30713,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, "search_context_cost_per_query": { @@ -30496,6 +30721,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -30525,6 +30751,7 @@ }, "gpt-5.1": { "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_priority": 2.5e-06, @@ -30565,11 +30792,17 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, "supports_minimal_reasoning_effort": false }, "gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_flex": 6.25e-08, "cache_read_input_token_cost_priority": 2.5e-07, "input_cost_per_token": 1.25e-06, "input_cost_per_token_priority": 2.5e-06, @@ -30610,6 +30843,11 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 6.25e-07, + "input_cost_per_token_flex": 6.25e-07, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": false, "supports_minimal_reasoning_effort": false }, @@ -30661,6 +30899,7 @@ }, "gpt-5.2": { "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_flex": 8.75e-08, "cache_read_input_token_cost_priority": 3.5e-07, "input_cost_per_token": 1.75e-06, "input_cost_per_token_priority": 3.5e-06, @@ -30702,11 +30941,17 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 8.75e-07, + "input_cost_per_token_flex": 8.75e-07, + "output_cost_per_token_batches": 7e-06, + "output_cost_per_token_flex": 7e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, "gpt-5.2-2025-12-11": { "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_flex": 8.75e-08, "cache_read_input_token_cost_priority": 3.5e-07, "input_cost_per_token": 1.75e-06, "input_cost_per_token_priority": 3.5e-06, @@ -30748,6 +30993,11 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "input_cost_per_token_batches": 8.75e-07, + "input_cost_per_token_flex": 8.75e-07, + "output_cost_per_token_batches": 7e-06, + "output_cost_per_token_flex": 7e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -30841,17 +31091,20 @@ }, "gpt-5.2-pro": { "input_cost_per_token": 2.1e-05, + "input_cost_per_token_batches": 1.05e-05, "litellm_provider": "openai", "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "output_cost_per_token": 0.000168, + "output_cost_per_token_batches": 8.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -30880,17 +31133,20 @@ }, "gpt-5.2-pro-2025-12-11": { "input_cost_per_token": 2.1e-05, + "input_cost_per_token_batches": 1.05e-05, "litellm_provider": "openai", "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", "output_cost_per_token": 0.000168, + "output_cost_per_token_batches": 8.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -30956,6 +31212,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31092,6 +31349,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31160,6 +31418,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31227,6 +31486,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -31290,7 +31550,7 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": true, - "source": "https://developers.openai.com/api/docs/models/gpt-5.6-cyber", + "source": "https://developers.openai.com/api/docs/pricing", "supports_computer_use": true, "supports_parallel_function_calling": true }, @@ -31460,7 +31720,7 @@ "reasoning_effort_levels": [ "medium" ], - "source": "https://developers.openai.com/api/docs/models/chat-latest", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -31539,7 +31799,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 5e-06, "output_cost_per_token_above_272k_tokens_flex": 2.25e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.5-2026-04-23": { "cache_read_input_token_cost": 5e-07, @@ -31596,7 +31857,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 5e-06, "output_cost_per_token_above_272k_tokens_flex": 2.25e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.5-pro": { "input_cost_per_token": 3e-05, @@ -31619,6 +31881,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -31667,6 +31930,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -31744,7 +32008,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 2.5e-06, "output_cost_per_token_above_272k_tokens_flex": 1.125e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-2026-03-05": { "cache_read_input_token_cost": 2.5e-07, @@ -31796,7 +32061,8 @@ "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 2.5e-06, "output_cost_per_token_above_272k_tokens_flex": 1.125e-05, - "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 + "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-pro": { "input_cost_per_token": 3e-05, @@ -31845,7 +32111,8 @@ "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 3e-05, - "output_cost_per_token_above_272k_tokens_flex": 0.000135 + "output_cost_per_token_above_272k_tokens_flex": 0.000135, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-pro-2026-03-05": { "input_cost_per_token": 3e-05, @@ -31894,7 +32161,8 @@ "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "input_cost_per_token_above_272k_tokens_flex": 3e-05, - "output_cost_per_token_above_272k_tokens_flex": 0.000135 + "output_cost_per_token_above_272k_tokens_flex": 0.000135, + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, @@ -31945,6 +32213,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -31997,6 +32266,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32046,6 +32316,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32095,6 +32366,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "default_reasoning_effort": "none", + "source": "https://developers.openai.com/api/docs/pricing", "supports_xhigh_reasoning_effort": true, "supports_minimal_reasoning_effort": false }, @@ -32113,6 +32385,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -32155,6 +32428,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/batch", "/v1/responses" @@ -32187,6 +32461,7 @@ "cache_read_input_token_cost_priority": 2.5e-07, "deprecation_date": "2026-12-11", "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 6.25e-07, "input_cost_per_token_flex": 6.25e-07, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -32195,6 +32470,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, "search_context_cost_per_query": { @@ -32202,6 +32478,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32526,6 +32803,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses" ], @@ -32556,6 +32834,7 @@ "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, "input_cost_per_token_flex": 1.25e-07, "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "openai", @@ -32564,6 +32843,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, "output_cost_per_token_flex": 1e-06, "output_cost_per_token_priority": 3.6e-06, "search_context_cost_per_query": { @@ -32571,6 +32851,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32604,6 +32885,7 @@ "cache_read_input_token_cost_priority": 4.5e-08, "deprecation_date": "2026-12-11", "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, "input_cost_per_token_flex": 1.25e-07, "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "openai", @@ -32612,6 +32894,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_batches": 1e-06, "output_cost_per_token_flex": 1e-06, "output_cost_per_token_priority": 3.6e-06, "search_context_cost_per_query": { @@ -32619,6 +32902,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32650,6 +32934,7 @@ "cache_read_input_token_cost": 5e-09, "cache_read_input_token_cost_flex": 2.5e-09, "input_cost_per_token": 5e-08, + "input_cost_per_token_batches": 2.5e-08, "input_cost_per_token_flex": 2.5e-08, "input_cost_per_token_priority": 2.5e-06, "litellm_provider": "openai", @@ -32658,12 +32943,14 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, "output_cost_per_token_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32696,6 +32983,7 @@ "cache_read_input_token_cost_flex": 2.5e-09, "deprecation_date": "2026-12-11", "input_cost_per_token": 5e-08, + "input_cost_per_token_batches": 2.5e-08, "input_cost_per_token_priority": 2.5e-06, "input_cost_per_token_flex": 2.5e-08, "litellm_provider": "openai", @@ -32704,12 +32992,14 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_batches": 2e-07, "output_cost_per_token_flex": 2e-07, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", @@ -32742,9 +33032,11 @@ "deprecation_date": "2026-10-23", "input_cost_per_image_token": 1e-05, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "openai", "mode": "image_generation", "output_cost_per_image_token": 4e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32755,9 +33047,11 @@ "deprecation_date": "2026-12-01", "input_cost_per_image_token": 2.5e-06, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "openai", "mode": "image_generation", "output_cost_per_image_token": 8e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -32778,6 +33072,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32811,6 +33106,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32844,6 +33140,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 2.4e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32879,6 +33176,7 @@ "output_cost_per_token": 2.4e-05, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32914,6 +33212,7 @@ "output_cost_per_token": 2.4e-06, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32939,6 +33238,7 @@ "cache_read_input_token_cost": 6e-08, "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1e-05, + "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, "litellm_provider": "openai", "max_input_tokens": 32000, @@ -32947,6 +33247,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -32981,6 +33282,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 1.6e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -37700,12 +38002,15 @@ "cache_read_input_token_cost": 7.5e-06, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.5e-05, + "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -37720,12 +38025,15 @@ "cache_read_input_token_cost": 7.5e-06, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.5e-05, + "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 6e-05, + "output_cost_per_token_batches": 3e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -37747,6 +38055,7 @@ "mode": "responses", "output_cost_per_token": 0.0006, "output_cost_per_token_batches": 0.0003, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -37780,6 +38089,7 @@ "mode": "responses", "output_cost_per_token": 0.0006, "output_cost_per_token_batches": 0.0003, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -37807,6 +38117,7 @@ "cache_read_input_token_cost_flex": 2.5e-07, "cache_read_input_token_cost_priority": 8.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 3.5e-06, "litellm_provider": "openai", @@ -37815,6 +38126,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8e-06, + "output_cost_per_token_batches": 4e-06, "output_cost_per_token_flex": 4e-06, "output_cost_per_token_priority": 1.4e-05, "search_context_cost_per_query": { @@ -37822,6 +38134,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/chat/completions", @@ -37851,6 +38164,7 @@ "cache_read_input_token_cost_priority": 8.75e-07, "deprecation_date": "2026-12-11", "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 3.5e-06, "litellm_provider": "openai", @@ -37859,6 +38173,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8e-06, + "output_cost_per_token_batches": 4e-06, "output_cost_per_token_flex": 4e-06, "output_cost_per_token_priority": 1.4e-05, "search_context_cost_per_query": { @@ -37866,6 +38181,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/chat/completions", @@ -37975,12 +38291,15 @@ "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -37993,12 +38312,15 @@ "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "litellm_provider": "openai", "max_input_tokens": 200000, "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_prompt_caching": true, @@ -38022,6 +38344,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -38059,6 +38382,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/responses", "/v1/batch" @@ -38082,10 +38406,11 @@ }, "o4-mini": { "cache_read_input_token_cost": 2.75e-07, - "cache_read_input_token_cost_flex": 1.375e-07, + "cache_read_input_token_cost_flex": 1.38e-07, "cache_read_input_token_cost_priority": 5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "input_cost_per_token_flex": 5.5e-07, "input_cost_per_token_priority": 2e-06, "litellm_provider": "openai", @@ -38094,6 +38419,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, "output_cost_per_token_flex": 2.2e-06, "output_cost_per_token_priority": 8e-06, "search_context_cost_per_query": { @@ -38101,6 +38427,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -38113,10 +38440,11 @@ }, "o4-mini-2025-04-16": { "cache_read_input_token_cost": 2.75e-07, - "cache_read_input_token_cost_flex": 1.375e-07, + "cache_read_input_token_cost_flex": 1.38e-07, "cache_read_input_token_cost_priority": 5e-07, "deprecation_date": "2026-10-23", "input_cost_per_token": 1.1e-06, + "input_cost_per_token_batches": 5.5e-07, "input_cost_per_token_flex": 5.5e-07, "input_cost_per_token_priority": 2e-06, "litellm_provider": "openai", @@ -38125,6 +38453,7 @@ "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.4e-06, + "output_cost_per_token_batches": 2.2e-06, "output_cost_per_token_flex": 2.2e-06, "output_cost_per_token_priority": 8e-06, "search_context_cost_per_query": { @@ -38132,6 +38461,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, @@ -43294,7 +43624,8 @@ "mode": "embedding", "output_cost_per_token": 0.0, "output_cost_per_token_batches": 0.0, - "output_vector_size": 3072 + "output_vector_size": 3072, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-3-small": { "input_cost_per_token": 2e-08, @@ -43305,7 +43636,8 @@ "mode": "embedding", "output_cost_per_token": 0.0, "output_cost_per_token_batches": 0.0, - "output_vector_size": 1536 + "output_vector_size": 1536, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-ada-002": { "input_cost_per_token": 1e-07, @@ -43314,7 +43646,8 @@ "max_tokens": 8191, "mode": "embedding", "output_cost_per_token": 0.0, - "output_vector_size": 1536 + "output_vector_size": 1536, + "source": "https://developers.openai.com/api/docs/pricing" }, "text-embedding-ada-002-v2": { "input_cost_per_token": 1e-07, @@ -43485,7 +43818,7 @@ "input_cost_per_token": 1.2e-06, "output_cost_per_token": 1.2e-06, "max_input_tokens": 131072, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen2.5-7B-Instruct-Turbo": { "litellm_provider": "together_ai", @@ -43497,7 +43830,7 @@ "input_cost_per_token": 3e-07, "output_cost_per_token": 3e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { "deprecation_date": "2026-07-10", @@ -43544,7 +43877,7 @@ "max_input_tokens": 256000, "mode": "chat", "output_cost_per_token": 2e-06, - "source": "https://www.together.ai/models/qwen3-coder-480b-a35b-instruct", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43606,7 +43939,7 @@ }, "mode": "chat", "output_cost_per_token": 1.7e-06, - "source": "https://www.together.ai/models/deepseek-v3-1", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43630,7 +43963,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.04e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43664,6 +43997,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 5.9e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43686,6 +44020,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43697,6 +44032,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 1.8e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43713,7 +44049,7 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 2e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { "deprecation_date": "2026-04-02", @@ -43725,7 +44061,7 @@ "input_cost_per_token": 1e-07, "output_cost_per_token": 3e-07, "max_input_tokens": 32768, - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { "deprecation_date": "2026-04-16", @@ -43733,6 +44069,7 @@ "litellm_provider": "together_ai", "mode": "chat", "output_cost_per_token": 6e-07, + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43759,7 +44096,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://www.together.ai/models/gpt-oss-120b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43773,7 +44110,7 @@ "max_input_tokens": 131072, "mode": "chat", "output_cost_per_token": 2e-07, - "source": "https://www.together.ai/models/gpt-oss-20b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43793,7 +44130,7 @@ "max_input_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.1e-06, - "source": "https://www.together.ai/models/glm-4-5-air", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43809,7 +44146,7 @@ }, "mode": "chat", "output_cost_per_token": 2.2e-06, - "source": "https://www.together.ai/models/glm-4-6", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43826,7 +44163,7 @@ }, "mode": "chat", "output_cost_per_token": 2e-06, - "source": "https://www.together.ai/models/glm-4-7", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43874,7 +44211,7 @@ }, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://www.together.ai/models/qwen3-next-80b-a3b-instruct", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43890,7 +44227,7 @@ }, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://www.together.ai/models/qwen3-next-80b-a3b-thinking", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -43904,7 +44241,7 @@ "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 3.6e-06, - "source": "https://www.together.ai/models/qwen3-5-397b-a17b", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -43919,7 +44256,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -43944,7 +44281,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 2.5e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_reasoning": true, @@ -43959,7 +44296,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 3e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_reasoning": true }, "together_ai/Qwen/Qwen3.7-Max": { @@ -43970,7 +44307,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 7.5e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/Qwen/Qwen3.7-Plus": { @@ -43980,7 +44317,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 1.28e-06, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3.8-2.4T-A95B": { "cache_read_input_token_cost": 2.5e-07, @@ -43990,7 +44327,7 @@ "max_tokens": 1010000, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/arize-ai/qwen-2-1.5b-instruct": { @@ -44000,7 +44337,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 1e-07, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { "cache_read_input_token_cost": 3e-08, @@ -44010,7 +44347,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.8e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44042,7 +44379,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 3.96e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44067,7 +44404,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 9.7e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -44103,7 +44440,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/moonshotai/Kimi-K2.7-Code": { @@ -44115,7 +44452,7 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44136,7 +44473,7 @@ "high", "max" ], - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44154,7 +44491,7 @@ "max_tokens": 512288, "mode": "chat", "output_cost_per_token": 3.6e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44180,7 +44517,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 4.05e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44196,7 +44533,7 @@ "max_tokens": 524288, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_prompt_caching": true }, "together_ai/zai-org/GLM-5.2": { @@ -44208,7 +44545,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44225,7 +44562,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44242,7 +44579,7 @@ "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-07, - "source": "https://docs.together.ai/docs/serverless-models", + "source": "https://api.together.ai/v1/models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_prompt_caching": true, @@ -44255,6 +44592,7 @@ "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", "mode": "audio_speech", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ] @@ -44263,6 +44601,7 @@ "input_cost_per_character": 3e-05, "litellm_provider": "openai", "mode": "audio_speech", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ] @@ -47605,6 +47944,9 @@ "cache_read_input_token_cost": 3e-08, "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 32768, "max_output_tokens": 32768, @@ -47614,8 +47956,10 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, "rpm": 100000, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-generation#edit-an-image", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47647,10 +47991,19 @@ "supports_image_size": false }, "vertex_ai/gemini-3-pro-image": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, + "cache_read_input_token_cost_flex": 1e-07, + "cache_read_input_token_cost_priority": 3.6e-07, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 3.6e-06, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -47659,9 +48012,13 @@ "output_cost_per_image": 0.134, "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, + "output_cost_per_token_flex": 6e-06, + "output_cost_per_token_priority": 2.16e-05, "supports_reasoning": false, - "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -47680,9 +48037,13 @@ "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, "vertex_ai/gemini-3.1-flash-image": { + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_flex": 2.5e-08, "deprecation_date": "2027-05-28", "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 32768, @@ -47691,8 +48052,10 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, + "output_cost_per_token_batches": 1.5e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_reasoning": false, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/gemini-3.1-flash-image-preview": { "input_cost_per_image": 0.00056, @@ -47710,9 +48073,11 @@ }, "vertex_ai/gemini-3.1-flash-lite-image": { "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_flex": 1.25e-08, "input_cost_per_image": 0.00028, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, "litellm_provider": "vertex_ai-language-models", "max_input_tokens": 65536, "max_output_tokens": 4096, @@ -47722,6 +48087,7 @@ "output_cost_per_image_token": 3e-05, "output_cost_per_token": 1.5e-06, "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -47796,6 +48162,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "vertex_ai/gemini-3.1-flash-lite": { + "cache_read_input_audio_token_cost": 5e-08, "deprecation_date": "2027-05-07", "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, @@ -47816,7 +48183,7 @@ "output_cost_per_token_flex": 7.5e-07, "output_cost_per_token_priority": 2.7e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47858,7 +48225,7 @@ "deprecation_date": "2027-07-21", "cache_read_input_token_cost": 3e-08, "cache_read_input_token_cost_flex": 1.5e-08, - "cache_read_input_token_cost_priority": 5e-08, + "cache_read_input_token_cost_priority": 5.4e-08, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, "input_cost_per_token_flex": 1.5e-07, @@ -47874,7 +48241,7 @@ "output_cost_per_token_flex": 1.25e-06, "output_cost_per_token_priority": 4.5e-06, "regional_endpoint_uplift_multiplier": 1.1, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -47913,6 +48280,7 @@ "google_maps_grounding_cost_per_query": 0.014 }, "vertex_ai/deep-research-pro-preview-12-2025": { + "cache_read_input_token_cost": 2e-07, "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -47925,7 +48293,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, - "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", @@ -49588,7 +49956,8 @@ "supported_endpoints": [ "/v1/audio/transcriptions" ], - "deprecation_date": "2027-02-26" + "deprecation_date": "2027-02-26", + "source": "https://developers.openai.com/api/docs/pricing" }, "xai/grok-3": { "cache_read_input_token_cost": 2e-07, @@ -52685,10 +53054,11 @@ "max_tokens": 40960, "max_input_tokens": 40960, "max_output_tokens": 40960, - "input_cost_per_token": 0.0, + "input_cost_per_token": 2e-07, "output_cost_per_token": 0.0, "litellm_provider": "fireworks_ai", - "mode": "rerank" + "mode": "rerank", + "source": "https://api.fireworks.ai/v1/serverless/models" }, "fireworks_ai/accounts/fireworks/models/qwen3-vl-235b-a22b-instruct": { "max_tokens": 262144, @@ -52753,7 +53123,7 @@ "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 1.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -54652,12 +55022,13 @@ }, "gpt-4o-mini-tts-2025-03-20": { "deprecation_date": "2026-07-23", - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -54670,12 +55041,13 @@ ] }, "gpt-4o-mini-tts-2025-12-15": { - "input_cost_per_token": 2.5e-06, + "input_cost_per_token": 6e-07, "litellm_provider": "openai", "mode": "audio_speech", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_second": 0.00025, "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/speech" ], @@ -54690,24 +55062,28 @@ "gpt-4o-mini-transcribe-2025-03-20": { "deprecation_date": "2027-01-20", "input_cost_per_audio_token": 1.25e-06, + "input_cost_per_second": 5e-05, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions" ] }, "gpt-4o-mini-transcribe-2025-12-15": { "input_cost_per_audio_token": 1.25e-06, + "input_cost_per_second": 5e-05, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -54726,6 +55102,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -54753,6 +55130,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "source": "https://developers.openai.com/api/docs/pricing", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, @@ -54780,6 +55158,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime" ], @@ -54831,13 +55210,14 @@ "supports_parallel_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "deprecation_date": "2027-01-20" + "deprecation_date": "2027-01-20", + "source": "https://developers.openai.com/api/docs/pricing" }, "gpt-realtime-whisper": { - "input_cost_per_second": 0.0002833333333333333, + "input_cost_per_second": 0.000283333333333, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-realtime-whisper", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime", "/v1/realtime/transcription_sessions" @@ -54855,7 +55235,7 @@ "litellm_provider": "openai", "mode": "video_generation", "output_cost_per_video_per_second": 0.1, - "source": "https://platform.openai.com/docs/api-reference/videos", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "image" @@ -54869,7 +55249,7 @@ "litellm_provider": "openai", "mode": "video_generation", "output_cost_per_video_per_second": 0.3, - "source": "https://platform.openai.com/docs/api-reference/videos", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "image" @@ -54894,11 +55274,15 @@ "chatgpt-image-latest": { "cache_read_input_token_cost": 1.25e-06, "deprecation_date": "2026-12-01", - "input_cost_per_image_token": 1e-05, + "input_cost_per_image_token": 8e-06, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "openai", "mode": "image_generation", - "output_cost_per_image_token": 4e-05, + "output_cost_per_image_token": 3.2e-05, + "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/images/generations", "/v1/images/edits" @@ -57373,7 +57757,7 @@ "input_cost_per_second": 7.5e-05, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-transcribe", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/audio/transcriptions", "/v1/realtime/transcription_sessions" @@ -57388,10 +57772,10 @@ "supports_audio_input": true }, "gpt-live-transcribe": { - "input_cost_per_second": 0.0002833333333333333, + "input_cost_per_second": 0.000283333333333, "litellm_provider": "openai", "mode": "audio_transcription", - "source": "https://developers.openai.com/api/docs/models/gpt-live-transcribe", + "source": "https://developers.openai.com/api/docs/pricing", "supported_endpoints": [ "/v1/realtime", "/v1/realtime/transcription_sessions" @@ -57406,10 +57790,10 @@ "supports_audio_input": true }, "gpt-live-1": { - "input_cost_per_second": 0.0008333333333333334, + "input_cost_per_second": 0.000833333333333, "litellm_provider": "openai", "mode": "realtime", - "source": "https://developers.openai.com/api/docs/models/gpt-live-1", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "text", "audio" @@ -57423,13 +57807,13 @@ "supports_function_calling": true }, "gpt-realtime-translate": { - "input_cost_per_second": 0.0005666666666666667, + "input_cost_per_second": 0.000566666666667, "litellm_provider": "openai", "max_input_tokens": 16000, "max_output_tokens": 2000, "max_tokens": 2000, "mode": "realtime", - "source": "https://developers.openai.com/api/docs/models/gpt-realtime-translate", + "source": "https://developers.openai.com/api/docs/pricing", "supported_modalities": [ "audio" ], @@ -57457,7 +57841,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, - "source": "https://platform.claude.com/docs/en/about-claude/models/overview", + "source": "https://platform.claude.com/docs/en/about-claude/pricing", "supports_adaptive_thinking": true, "thinking_always_on": true, "supports_mid_conversation_system": true, @@ -57518,7 +57902,7 @@ "supports_output_config": true, "prompt_cache_min_tokens": 512, "supports_native_structured_output": true, - "source": "https://platform.claude.com/docs/en/models/mythos-5-1/overview" + "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-mythos-preview": { "cache_creation_input_token_cost": 1.25e-05, @@ -57874,6 +58258,7 @@ }, "vertex_ai/gemini-3.5-live-translate-preview": { "input_cost_per_audio_token": 3.5e-06, + "input_cost_per_second": 8.83333333333e-05, "input_cost_per_token": 3.5e-06, "litellm_provider": "vertex_ai", "mode": "realtime", @@ -57976,14 +58361,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -57992,14 +58380,17 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4p1-flash": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://fireworks.ai/models/deepseek-ai/deepseek-v4p1-flash", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -58015,26 +58406,29 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, "supports_vision": true }, "fireworks_ai/accounts/fireworks/models/kimi-k3": { "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_priority": 3.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_priority": 1.875e-05, "reasoning_effort_levels": [ "low", "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58043,14 +58437,17 @@ }, "fireworks_ai/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58059,14 +58456,17 @@ }, "fireworks_ai/deepseek-v4p1-flash": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://fireworks.ai/models/deepseek-ai/deepseek-v4p1-flash", + "output_cost_per_token_priority": 8.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -58082,7 +58482,7 @@ "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, "supports_vision": true @@ -58121,19 +58521,22 @@ }, "fireworks_ai/kimi-k3": { "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_priority": 3.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_priority": 1.875e-05, "reasoning_effort_levels": [ "low", "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58184,12 +58587,15 @@ }, "fireworks_ai/qwen3p8-max": { "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 9e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58205,7 +58611,7 @@ "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58237,7 +58643,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 2.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58253,7 +58659,7 @@ "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58285,7 +58691,7 @@ "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 2.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58294,12 +58700,15 @@ }, "fireworks_ai/accounts/fireworks/models/qwen3p8-max": { "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_priority": 3.75e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_priority": 3e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 9e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58315,7 +58724,7 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 6.6e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -58352,7 +58761,7 @@ "high", "max" ], - "source": "https://docs.fireworks.ai/serverless/pricing", + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -61050,14 +61459,17 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p3": { "cache_read_input_token_cost": 2.6e-07, + "cache_read_input_token_cost_priority": 3.25e-07, "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -61066,13 +61478,16 @@ }, "fireworks_ai/accounts/fireworks/models/glm-5p3-flash": { "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_priority": 3.75e-08, "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.875e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 5e-07, - "source": "https://docs.fireworks.ai/serverless/pricing", + "output_cost_per_token_priority": 6.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -61100,7 +61515,7 @@ "max_output_tokens": 40960, "max_tokens": 40960, "mode": "embedding", - "source": "https://docs.fireworks.ai/serverless/pricing" + "source": "https://api.fireworks.ai/v1/serverless/models" }, "zai/glm-5.2": { "cache_creation_input_token_cost": 0, @@ -61124,7 +61539,7 @@ "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 4.7e-07, - "source": "https://docs.together.ai/docs/serverless-models" + "source": "https://api.together.ai/v1/models" }, "together_ai/moonshotai/Kimi-K2.6": { "deprecation_date": "2026-08-19", @@ -61134,7 +61549,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/moonshotai/Kimi-K2.5-fp4": { "input_cost_per_token": 5e-07, @@ -61142,7 +61557,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/MiniMaxAI/MiniMax-M2.7": { "input_cost_per_token": 3e-07, @@ -61151,7 +61566,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 196608, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5": { "deprecation_date": "2026-06-22", @@ -61160,7 +61575,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 202752, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/zai-org/GLM-5.1": { "deprecation_date": "2026-07-10", @@ -61170,7 +61585,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 202752, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/deepseek-ai/DeepSeek-R1-0528": { "input_cost_per_token": 3e-06, @@ -61178,7 +61593,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 163840, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-Coder-Next-FP8": { "deprecation_date": "2026-05-14", @@ -61187,7 +61602,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-32B-Instruct": { "deprecation_date": "2026-02-25", @@ -61196,7 +61611,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/Qwen3-VL-8B-Instruct": { "deprecation_date": "2026-04-16", @@ -61205,7 +61620,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Ministral-3-14B-Instruct-2512": { "input_cost_per_token": 2e-07, @@ -61213,7 +61628,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 262144, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": { "input_cost_per_token": 6e-08, @@ -61221,7 +61636,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 131072, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/mistralai/Mistral-7B-Instruct-v0.3": { "input_cost_per_token": 2e-07, @@ -61229,7 +61644,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 32768, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "together_ai/Qwen/QwQ-32B": { "deprecation_date": "2025-11-13", @@ -61238,7 +61653,7 @@ "litellm_provider": "together_ai", "max_input_tokens": 131072, "mode": "chat", - "source": "https://api.together.xyz/v1/models" + "source": "https://api.together.ai/v1/models" }, "cerebras/gemma-4-31b": { "input_cost_per_token": 9.9e-07, @@ -65365,5 +65780,263 @@ "supports_tool_choice": false, "supports_response_schema": true, "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": { + "cache_read_input_token_cost": 3.9e-07, + "input_cost_per_token": 2.1e-06, + "litellm_provider": "fireworks_ai", + "mode": "chat", + "output_cost_per_token": 6.6e-06, + "source": "https://api.fireworks.ai/v1/serverless/models" + }, + "together_ai/arcee-ai/trinity-mini": { + "input_cost_per_token": 4.5e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.5e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/deepseek-coder-33b-instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Llama-70B": { + "input_cost_per_token": 2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-Distill-Qwen-14B": { + "input_cost_per_token": 1.6e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-V4.1-Flash": { + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "vertex_ai/gemini-2.5-flash-native-audio": { + "input_cost_per_audio_token": 3e-06, + "input_cost_per_token": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "realtime", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_token": 2e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-2.5-flash-preview-tts": { + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "litellm_provider": "vertex_ai", + "mode": "audio_speech", + "output_cost_per_audio_token": 1e-05, + "output_cost_per_token": 1e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.1-flash-live-preview": { + "input_cost_per_audio_token": 3e-06, + "input_cost_per_second": 8.33333333333e-05, + "input_cost_per_token": 7.5e-07, + "litellm_provider": "vertex_ai", + "mode": "realtime", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_token": 4.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.1-flash-tts-preview": { + "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "audio_speech", + "output_cost_per_audio_token": 2e-05, + "output_cost_per_token": 2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.5-transcribe": { + "input_cost_per_audio_token": 2e-06, + "input_cost_per_second": 5e-05, + "litellm_provider": "vertex_ai", + "mode": "audio_transcription", + "output_cost_per_token": 1.2e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-3.5-transcribe-live": { + "input_cost_per_audio_token": 3.5e-06, + "input_cost_per_second": 8.33333333333e-05, + "litellm_provider": "vertex_ai", + "mode": "audio_transcription", + "output_cost_per_token": 2.1e-05, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-omni-1.1-flash": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 9e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemini-robotics-er-2": { + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "vertex_ai/gemma-4-26b-a4b-it": { + "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "vertex_ai", + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + }, + "together_ai/google/gemma-2-27b-it": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "gpt-5.5-cyber": { + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_token": 1.25e-05, + "litellm_provider": "openai", + "mode": "chat", + "output_cost_per_token": 7.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supports_reasoning": true + }, + "gpt-rosalind-research": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 5e-06, + "litellm_provider": "openai", + "mode": "chat", + "output_cost_per_token": 2.5e-05, + "source": "https://developers.openai.com/api/docs/pricing" + }, + "together_ai/meta-llama/Llama-3-8b-chat-hf": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.1-405B-Instruct": { + "input_cost_per_token": 3.5e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.2-1B-Instruct": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Llama-3.2-3B-Instruct": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Meta-Llama-3-70B-Instruct-Turbo": { + "input_cost_per_token": 8.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/meta-llama/Meta-Llama-3-8B-Instruct": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/NousResearch/Nous-Hermes-2-Mixtral-8x7B-DPO": { + "input_cost_per_token": 6e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/nvidia/Llama-3.1-Nemotron-70B-Instruct-HF": { + "input_cost_per_token": 8.8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8.8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-1.5B-Instruct": { + "input_cost_per_token": 2e-08, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 2e-08, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-72B-Instruct": { + "input_cost_per_token": 9e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 9e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2-VL-72B-Instruct": { + "input_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-14B-Instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-72B-Instruct": { + "input_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-Coder-32B-Instruct": { + "input_cost_per_token": 8e-07, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-07, + "source": "https://api.together.ai/v1/models" + }, + "together_ai/Qwen/Qwen2.5-VL-72B-Instruct": { + "input_cost_per_token": 1.95e-06, + "litellm_provider": "together_ai", + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://api.together.ai/v1/models" } } diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index d1ac3e67b2b..c2490041cf7 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -249,6 +249,10 @@ "type": "number", "minimum": 0 }, + "input_cost_per_audio_token_batches": { + "type": "number", + "minimum": 0 + }, "input_cost_per_audio_token_priority": { "type": "number", "minimum": 0, @@ -276,6 +280,10 @@ "type": "number", "minimum": 0 }, + "input_cost_per_image_token_batches": { + "type": "number", + "minimum": 0 + }, "input_cost_per_pixel": { "type": "number", "minimum": 0 @@ -375,6 +383,14 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "input_cost_per_video_token": { + "type": "number", + "minimum": 0 + }, + "input_cost_per_video_token_batches": { + "type": "number", + "minimum": 0 + }, "input_dbu_cost_per_token": { "type": "number", "minimum": 0 diff --git a/ruff.toml b/ruff.toml index 3ac4c1fc94d..fab3fe27aed 100644 --- a/ruff.toml +++ b/ruff.toml @@ -4,11 +4,12 @@ lint.ignore = ["F405", "E402", "F403"] # That gives editors and `ruff check --fix` the diagnostic, which the gate script cannot. lint.extend-select = [ "T20", "PGH004", "RUF008", "RUF009", "RUF100", - "B033", "FURB136", "FURB168", "FURB188", "I001", "PERF402", "PIE790", "PIE800", "PLC0208", - "PLR0402", "PLR1711", "PLR1730", "PLR2044", "PLW0133", "PYI030", "PYI041", "PYI064", "RET501", - "RUF010", "RUF022", "RUF023", "RUF051", "S113", "SIM114", "SIM118", "TC005", "UP006", "UP007", - "UP008", - "UP012", "UP018", "UP024", "UP032", "UP034", "UP035", "UP037", "UP045", + "B004", "B018", "B021", "B033", "FURB136", "FURB168", "FURB188", "I001", "PERF402", "PIE790", + "PIE800", "PLC0208", "PLR0124", "PLR0402", "PLR0206", "PLR1704", "PLR1711", "PLR1730", "PLR2044", + "PLW0133", "PYI030", "PYI041", "PYI064", "RET501", "RUF010", "RUF022", "RUF023", "RUF051", "S113", + "SIM114", "SIM118", "SIM201", "SIM211", "SIM222", "TC005", "UP006", "UP007", "UP008", + "UP012", "UP018", "UP024", "UP032", "UP034", "UP035", "UP036", "UP037", "UP045", + "C404", "C419", ] # RUF100 (unused-noqa) only knows the rules enabled in THIS config, so it would strip # `# noqa` directives that protect rules enforced elsewhere. List those codes as external diff --git a/tests/code_coverage_tests/test_e2e_changed_gate.py b/tests/code_coverage_tests/test_e2e_changed_gate.py index 588402e3996..101816c7f11 100644 --- a/tests/code_coverage_tests/test_e2e_changed_gate.py +++ b/tests/code_coverage_tests/test_e2e_changed_gate.py @@ -81,6 +81,25 @@ def test_missing_execution_evidence_fails(tmp_path: Path, contents: str) -> None assert result.returncode == 1 +@pytest.mark.parametrize("omitted_role", ("proxy_admin", "team_member", "internal_user_viewer")) +def test_one_passing_management_case_cannot_hide_a_missing_actor(tmp_path: Path, omitted_role: str) -> None: + suite: Final = ET.Element("testsuite") + path: Final = "tests/e2e/management/test_jwt_management_e2e.py" + case: Final = ET.SubElement(suite, "testcase", file=path) + properties: Final = ET.SubElement(case, "properties") + _ = ET.SubElement( + properties, + "property", + name="management_node", + value=f"{path}::TestJwtManagement::test_actor_subject_and_database_role[proxy_admin_viewer]", + ) + report: Final = tmp_path / "report.xml" + ET.ElementTree(suite).write(report) + result: Final = subprocess.run([sys.executable, str(GATE), str(report), path], capture_output=True, text=True) + assert result.returncode == 1 + assert f"test_actor_subject_and_database_role[{omitted_role}]" in result.stdout + + def test_short_values_are_written_without_masking_every_digit_in_the_log(tmp_path: Path) -> None: env_path: Final = tmp_path / ".env" @@ -141,6 +160,10 @@ def test_changed_suite_files_are_selected_unless_the_stack_cannot_run_them( ( "tests/e2e/proxy_client.py", "tests/e2e/conftest.py", + "tests/e2e/management/management_client.py", + "tests/e2e/management/jwt_actors.py", + "tests/e2e/management/conftest.py", + "tests/e2e/coverage_registry/management_cases.py", "tests/e2e/pytest.ini", "tests/e2e/gateway/stage_mirror_ci_config.yml", ".github/e2e-stack/up.sh", diff --git a/tests/code_coverage_tests/test_provider_replay_harness.py b/tests/code_coverage_tests/test_provider_replay_harness.py new file mode 100644 index 00000000000..e7c5c96b64b --- /dev/null +++ b/tests/code_coverage_tests/test_provider_replay_harness.py @@ -0,0 +1,329 @@ +from __future__ import annotations + +import json +import os +import subprocess +import sys +import threading +from pathlib import Path +from typing import Final + +import pytest +from fixture_bundle import BundleRecorder, LoadedBundle, load_bundle, prepare_bundle +from fixture_mode import current_test_key +from fixture_profile import MatchProfile +from provider_edge import REPLAY_MISS_STATUS, RecordEdge, ReplayEdge, ReplaySource +from test_provider_edge import ( + CHAT_PATH, + SSE_CHUNKS, + STREAM_BODY, + UPLOAD_PATH, + call_edge, + chunked_provider, + fake_provider, + json_object, + provider_url, + raw_stream_post, + running_edge, + this_tests_files, +) + + +class TestStrictIdentity: + @pytest.mark.parametrize("path", [CHAT_PATH, "/anthropic/v1/messages"]) + def test_roundtrip_rejects_semantic_changes(self, tmp_path: Path, path: str) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + original: Final = ( + b'{"model":"synthetic","messages":[{"role":"user",' + b'"content":"2031-04-05 00000000-0000-0000-0000-000000000001"}],"options":[1,2]}' + ) + headers: Final = { + "content-type": "application/json", + "accept": "application/json", + "anthropic-version": "2023-06-01", + "anthropic-beta": "feature-a", + "openai-beta": "feature-b", + "authorization": "Bearer synthetic-secret-one", + } + query: Final = "?part=one&part=two&blank=" + with fake_provider() as provider: + mounts: Final = {"openai": provider_url(provider), "anthropic": provider_url(provider)} + with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge: + captured: Final = call_edge(edge, "POST", path + query, body=original, headers=headers) + assert captured.status_code == 200 + assert json_object(captured.body)["echo"] == original.decode() + loaded: Final = load_bundle(recorder.root, profile="stateless_v1") + assert isinstance(loaded, LoadedBundle) + assert loaded.manifest.match_profile == "stateless_v1" + source: Final = ReplaySource(loaded) + with running_edge(ReplayEdge(source), mounts) as edge: + cases: Final = ( + (original.replace(b"2031-04-05", b"2032-06-07"), headers, query, "body"), + (original.replace(b"000000000001", b"000000000002"), headers, query, "body"), + (original.replace(b"[1,2]", b"[2,1]"), headers, query, "body"), + (original.replace(b"synthetic", b"other"), headers, query, "body"), + (original, headers, "?part=three&part=two&blank=", "query"), + (original, headers, "?part=two&part=one&blank=", "query"), + *( + (original, {k: v for k, v in headers.items() if k != name}, query, "headers") + for name in ("accept", "anthropic-version", "anthropic-beta", "openai-beta") + ), + *( + (original, {**headers, name: value}, query, "headers") + for name in ("accept", "anthropic-version", "anthropic-beta", "openai-beta") + for value in ("different", "") + ), + (original, {**headers, "authorization": "Basic synthetic-secret-two"}, query, "auth"), + (original, {k: v for k, v in headers.items() if k != "authorization"}, query, "auth"), + ) + for rejected, reason in ( + (call_edge(edge, "POST", path + changed_query, body=body, headers=changed_headers), reason) + for body, changed_headers, changed_query, reason in cases + ): + assert rejected.status_code == REPLAY_MISS_STATUS + assert reason in rejected.body.decode() + assert b"synthetic-secret" not in rejected.body + reordered: Final = json.dumps(dict(reversed(list(json_object(original).items())))).encode() + accepted: Final = call_edge( + edge, "POST", path + query, body=reordered, headers={k.upper(): v for k, v in headers.items()} + ) + assert accepted.status_code == 200 + assert accepted.body == captured.body + assert source.leftover_error(current_test_key()) is None + assert len(provider.hits) == 1 + assert "synthetic-secret" not in "".join(file.read_text() for file in recorder.root.rglob("*.json")) + + @pytest.mark.parametrize( + "body", + [ + b'{"value":null}', + b'{"value":""}', + b'{"value":false}', + b'{"value":0}', + b'{"value":[]}', + b'{"value":{}}', + b'{"value":0.123456789012345678901}', + b'{"value":0.123456789012345678902}', + b'{"value":1e400}', + b'{"value":1}', + b'{"value":1e0}', + b'{"value":-0}', + b'{"value":1e9999999999999999999}', + ], + ) + def test_json_values_remain_distinct(self, tmp_path: Path, body: bytes) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + with fake_provider() as provider: + mounts: Final = {"openai": provider_url(provider)} + with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge: + assert ( + call_edge( + edge, "POST", CHAT_PATH, body=body, headers={"content-type": "application/json"} + ).status_code + == 200 + ) + loaded: Final = load_bundle(recorder.root, profile="stateless_v1") + assert isinstance(loaded, LoadedBundle) + with running_edge(ReplayEdge(ReplaySource(loaded)), mounts) as edge: + values: Final = ( + b"{}", + b'{"value":null}', + b'{"value":""}', + b'{"value":false}', + b'{"value":0}', + b'{"value":[]}', + b'{"value":{}}', + b'{"value":0.123456789012345678901}', + b'{"value":0.123456789012345678902}', + b'{"value":1e400}', + b'{"value":1}', + b'{"value":1e0}', + b'{"value":-0}', + b'{"value":1e9999999999999999999}', + ) + for rejected in ( + call_edge(edge, "POST", CHAT_PATH, body=value, headers={"content-type": "application/json"}) + for value in values + if value != body + ): + assert rejected.status_code == REPLAY_MISS_STATUS + assert b"body" in rejected.body + assert ( + call_edge( + edge, "POST", CHAT_PATH, body=body, headers={"content-type": "application/json"} + ).status_code + == 200 + ) + assert len(provider.hits) == 1 + + @pytest.mark.parametrize( + "path,body,headers", + [ + (UPLOAD_PATH, b"{}", {"content-type": "application/json"}), + (CHAT_PATH + "?part=%FF", b"{}", {"content-type": "application/json"}), + (CHAT_PATH + "?part=%FE", b"{}", {"content-type": "application/json"}), + (CHAT_PATH, b"opaque", {"content-type": "application/octet-stream"}), + (CHAT_PATH, b"--boundary", {"content-type": "multipart/form-data; boundary=boundary"}), + (CHAT_PATH, b'{"x":1,"x":2}', {"content-type": "application/json"}), + (CHAT_PATH, b"{}", {"content-type": "application/json", "x-custom-behavior": "synthetic-private-value"}), + ], + ) + def test_ineligible_capture_never_calls_provider( + self, tmp_path: Path, path: str, body: bytes, headers: dict[str, str] + ) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + with fake_provider() as provider: + with running_edge(RecordEdge(recorder, threading.Lock()), {"openai": provider_url(provider)}) as edge: + result: Final = call_edge(edge, "POST", path, body=body, headers=headers) + assert result.status_code == REPLAY_MISS_STATUS + assert b"eligibility error" in result.body + assert b"synthetic-private-value" not in result.body + assert provider.hits == [] + assert this_tests_files(recorder.root) == [] + + def test_destination_is_part_of_actual_http_identity(self, tmp_path: Path) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + with fake_provider() as provider: + with running_edge(RecordEdge(recorder, threading.Lock()), {"openai": provider_url(provider)}) as edge: + assert ( + call_edge( + edge, "POST", CHAT_PATH, body=b"{}", headers={"content-type": "application/json"} + ).status_code + == 200 + ) + loaded: Final = load_bundle(recorder.root, profile="stateless_v1") + assert isinstance(loaded, LoadedBundle) + with running_edge(ReplayEdge(ReplaySource(loaded)), {"openai": provider_url(provider) + "/other"}) as edge: + result: Final = call_edge( + edge, "POST", CHAT_PATH, body=b"{}", headers={"content-type": "application/json"} + ) + assert result.status_code == REPLAY_MISS_STATUS + assert b"upstream" in result.body + assert len(provider.hits) == 1 + + def test_credentials_are_not_identity_and_fresh_process_replays(self, tmp_path: Path) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + headers: Final = { + "content-type": "application/json", + "authorization": "bEaReR synthetic-token", + "x-api-key": "synthetic-api-key", + "cookie": "synthetic-cookie", + } + path: Final = CHAT_PATH + "?api_key=synthetic-query-secret&part=one&part=two" + body: Final = b'{"model":"synthetic","messages":[]}' + with fake_provider(echo_request=False) as provider: + mounts: Final = {"openai": provider_url(provider)} + with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge: + captured: Final = call_edge(edge, "POST", path, body=body, headers=headers) + assert captured.status_code == 200 + seen_headers, seen_body = provider.requests[0] + assert {k.lower(): v for k, v in seen_headers.items()}.items() >= headers.items() + assert seen_body == body + assert provider.hits == ["POST " + path.removeprefix("/openai")] + artifacts: Final = "".join(file.read_text() for file in recorder.root.rglob("*.json")) + for secret in ("synthetic-token", "synthetic-api-key", "synthetic-cookie", "synthetic-query-secret"): + assert secret not in artifacts + child: Final = subprocess.run( + [ + sys.executable, + "-c", + """ +import json, sys +from pathlib import Path +from fixture_bundle import LoadedBundle, load_bundle +from provider_edge import ProviderRequestObservation, observed_provider_edge, replay_leftover_error +from test_provider_edge import call_edge +from fixture_profile import MatchProfile +from fixture_mode import current_test_key +loaded = load_bundle(Path(sys.argv[1]), profile="stateless_v1") +assert isinstance(loaded, LoadedBundle) +with observed_provider_edge(ProviderRequestObservation("synthetic"), mode_raw="replay", bundle_dir=Path(sys.argv[1]), bind_host="127.0.0.1", advertise_host="127.0.0.1", mounts={"openai": sys.argv[2]}) as edge: + response = call_edge(edge, "POST", sys.argv[3], body=sys.argv[4].encode(), headers=json.loads(sys.argv[5])) + assert response.status_code == 200 + print(response.body.decode()) +assert replay_leftover_error(mode_raw="replay", bundle_dir=Path(sys.argv[1]), test_key=current_test_key()) is None +""", + str(recorder.root), + provider_url(provider), + path.replace("synthetic-query-secret", "new-query-credential"), + body.decode(), + json.dumps({**headers, "authorization": "Bearer another-credential", "x-api-key": "another-key"}), + ], + env={ + **os.environ, + "PYTHONPATH": str(Path(__file__).resolve().parents[1] / "e2e"), + "E2E_REPLAY_MATCH_PROFILE": "stateless_v1", + }, + capture_output=True, + text=True, + timeout=30, + ) + assert child.returncode == 0, child.stderr + assert child.stdout.strip().encode() == captured.body + assert len(provider.hits) == 1 + + @pytest.mark.parametrize("profile,other", [("legacy", "stateless_v1"), ("stateless_v1", "legacy")]) + def test_profiles_cannot_load_each_others_bundles( + self, tmp_path: Path, profile: MatchProfile, other: MatchProfile + ) -> None: + from fixture_bundle import UnreadableBundle + + recorder: Final = prepare_bundle(tmp_path / profile, profile=profile) + assert isinstance(recorder, BundleRecorder) + mismatch: Final = load_bundle(recorder.root, profile=other) + assert isinstance(mismatch, UnreadableBundle) + assert "profile mismatch" in mismatch.reason + assert "re-record" in mismatch.reason + + @pytest.mark.parametrize("abort_after", [None, 2]) + def test_strict_stream_preserves_chunks_and_truncation(self, tmp_path: Path, abort_after: int | None) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + with chunked_provider(abort_after=abort_after) as provider: + mounts: Final = {"anthropic": provider_url(provider)} + with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge: + _, captured, captured_ending = raw_stream_post(edge.port, "/anthropic/v1/messages", STREAM_BODY) + loaded: Final = load_bundle(recorder.root, profile="stateless_v1") + assert isinstance(loaded, LoadedBundle) + source: Final = ReplaySource(loaded) + with running_edge(ReplayEdge(source), mounts) as edge: + _, replayed, ending = raw_stream_post(edge.port, "/anthropic/v1/messages", STREAM_BODY) + assert captured == replayed == list(SSE_CHUNKS[:abort_after]) + assert ending == captured_ending + assert (ending == "terminated") == (abort_after is None) + assert source.leftover_error(current_test_key()) is None + assert len(provider.hits) == 1 + + def test_auth_scheme_survives_missing_credentials(self, tmp_path: Path) -> None: + recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1") + assert isinstance(recorder, BundleRecorder) + headers: Final = {"content-type": "application/json", "authorization": "Bearer"} + with fake_provider() as provider: + mounts: Final = {"openai": provider_url(provider)} + with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge: + assert call_edge(edge, "POST", CHAT_PATH, body=b"{}", headers=headers).status_code == 200 + loaded: Final = load_bundle(recorder.root, profile="stateless_v1") + assert isinstance(loaded, LoadedBundle) + with running_edge(ReplayEdge(ReplaySource(loaded)), mounts) as edge: + for result in ( + call_edge(edge, "POST", CHAT_PATH, body=b"{}", headers={**headers, "authorization": scheme}) + for scheme in ("Basic", "Digest") + ): + assert result.status_code == REPLAY_MISS_STATUS + assert b"auth" in result.body + assert ( + call_edge( + edge, + "POST", + CHAT_PATH, + body=b"{}", + headers={**headers, "authorization": "bEaReR synthetic-token"}, + ).status_code + == 200 + ) + assert len(provider.hits) == 1 diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 44564a51e26..53a05931ca7 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -60,7 +60,11 @@ The suites run against a live proxy, so bring one up first by running the litell Keycloak's password grant is a test-only provisioning shortcut, not a production login recommendation. The `litellm-e2e-admin` client adds the proxy's admin scope; the normal client does not. Never reuse this permissive realm outside an isolated test stack. - Management tests can use the shared `idp` and `jwt_identity` fixtures. Each test gets a unique Keycloak group/user and a matching proxy user/team. Setup and fallback cleanup use the master key; the operations and read-backs being tested must explicitly use `caller_key=idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID)` (or a member token). See `management/test_jwt_management_e2e.py` for create/read/update/clear/delete and tenant-denial examples. A group claim alone is not database team membership: permission tests explicitly add the member and prove an allowed read before asserting the denied write. + Management tests can bind a credential once with `client.with_caller(Caller(...))`; direct calls, delegated helpers and replica read-backs then retain that caller. Explicit `caller_key` arguments override the binding. Keep the original master-backed client for bootstrap and cleanup. `actor_factory` lazily provisions database roles and tenant memberships, with `database_role` tokens carrying no groups and `group_scoped` actors retaining the existing team route gate. Token minting is explicit through `actor.mint_caller(idp)`. The factory runs requests without backend retries and reports cleanup failures. `coverage_registry/management_cases.py` records exact canary nodes and non-secret actor labels; the CI execution assertion rejects a missing or skipped actor row + + For the opt-in browser profile, start the existing IdP first, then run `.github/e2e-stack/oidc-profile.sh "$PROXY_BASE_URL" `. The wrapper creates a confidential client with an exact `/sso/callback` redirect and S256 PKCE, passes the client secret only through the child process environment, and removes the client on exit. It uses the existing generic OIDC handler with `GENERIC_USER_ID_ATTRIBUTE=sub`. Preserve the IdP's PostgreSQL data across restarts + + `tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping; browser journey specs under `ui/oidc/` are a separate coverage step Every successful IdP create immediately registers cleanup, including partial setup failures. Cleanup failures emit warnings. Tokens are minted on demand, and the expiration test waits relative to the token's actual `exp` with a bounded clock-drift check. To check first-attempt behavior locally, run both files with `--reruns 0`: @@ -232,3 +236,15 @@ Before you push 4. Capture screenshots of the test run and attach them to the PR as proof 5. If a test fails because it surfaced a real issue in the product, flag that explicitly in the PR rather than reworking the test until it passes + +### Strict stateless replay matching + +Set `E2E_REPLAY_MATCH_PROFILE=stateless_v1` for both recording and replay to bind OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` requests to their upstream destination, ordered query pairs, semantic headers and literal JSON content. The default remains `legacy`. Strict bundles use format 5 and cannot load as legacy bundles; select the matching profile or re-record with `E2E_FIXTURE_MODE=record`. Missing profile metadata never enrolls a legacy bundle in strict matching + +Strict matching preserves dates, UUIDs, hashes, model names, tool arguments, array order and omitted/null/empty/false/zero values. JSON object key order and header name casing may change. The strict body uses tagged JSON values so number precision and JSON types survive persistence, including exact numeric spelling and numbers larger than a floating-point value. Invalid UTF-8 query values fail eligibility. Duplicate JSON keys, unsupported endpoints, non-JSON bodies and unknown semantic headers fail eligibility before contacting a provider + +The semantic header set is `content-type`, `accept`, `anthropic-version`, `anthropic-beta` and `openai-beta`, including missing versus present values. Authorization records presence and the case-insensitive scheme; `x-api-key` records presence only. Credential values and cookies are excluded. Credential query values are redacted while their position and field name remain in the identity. Never use real customer inputs in fixture qualification + +Excluded transport and telemetry headers are `host`, `content-length`, `connection`, `accept-encoding`, `user-agent`, `traceparent`, `tracestate`, `x-request-id`, `x-client-request-id` and `x-stainless-*`. Inbound transfer-encoding is unsupported; send JSON with content-length framing. The destination represents host identity and the relay carries original body bytes. Replay does not verify credentials, SDK timeout/retry behavior, transport performance, model availability or stateful remote IDs. Live relay uses original request bytes and header values, never the stored identity + +Strict replay harness regression tests live in `tests/code_coverage_tests/test_provider_replay_harness.py`. The CircleCI `provider_replay_harness` job runs them alongside the existing legacy harness files with `--noconftest -o pythonpath=tests/e2e`; they need only synthetic HTTP providers and temporary fixture storage diff --git a/tests/e2e/coverage_registry/management_cases.py b/tests/e2e/coverage_registry/management_cases.py new file mode 100644 index 00000000000..812dbfe5b8d --- /dev/null +++ b/tests/e2e/coverage_registry/management_cases.py @@ -0,0 +1,151 @@ +from dataclasses import dataclass +from typing import Final, Literal + +CredentialKind = Literal["master", "idp_admin", "direct_jwt", "virtual_key", "dashboard_session"] +DependencyProfile = Literal["management_only", "real_oidc_browser", "external_provider_required"] + + +@dataclass(frozen=True, slots=True) +class ManagementCase: + node: str + credential_kind: CredentialKind + actor: str + profile: str + method: Literal["GET", "POST"] + path: str + operation_family: str + dependency_profile: DependencyProfile = "management_only" + + +JWT_FILE: Final = "tests/e2e/management/test_jwt_management_e2e.py" +JWT_CLASS: Final = f"{JWT_FILE}::TestJwtManagement" +ACTORS: Final = ( + "proxy_admin", + "proxy_admin_viewer", + "organization_admin", + "team_admin", + "team_member", + "internal_user", + "internal_user_viewer", + "unrelated_user", +) +MANAGEMENT_CASES: Final = tuple( + ManagementCase( + node=f"{JWT_CLASS}::test_actor_subject_and_database_role[{role}]", + credential_kind="direct_jwt", + actor=role, + profile="database_role", + method="GET", + path="/user/info", + operation_family="identity", + ) + for role in ACTORS +) + ( + ManagementCase( + node=f"{JWT_CLASS}::test_admin_viewer_reads_but_cannot_update", + credential_kind="direct_jwt", + actor="proxy_admin_viewer", + profile="database_role", + method="POST", + path="/key/update", + operation_family="denial", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_admin_creates_reads_updates_clears_and_deletes_a_key[direct_jwt]", + credential_kind="direct_jwt", + actor="proxy_admin", + profile="group_scoped", + method="POST", + path="/key/generate", + operation_family="key_lifecycle", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_admin_creates_reads_updates_clears_and_deletes_a_key[virtual_key]", + credential_kind="virtual_key", + actor="proxy_admin", + profile="group_scoped", + method="POST", + path="/key/generate", + operation_family="key_lifecycle", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_two_actor_sets_keep_tenants_and_keys_isolated", + credential_kind="direct_jwt", + actor="team_member", + profile="group_scoped", + method="GET", + path="/key/info", + operation_family="tenant_isolation", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_member_cannot_write_and_another_team_cannot_read_the_key", + credential_kind="direct_jwt", + actor="team_member", + profile="group_scoped", + method="POST", + path="/key/update", + operation_family="tenant_isolation", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_multi_group_actor_keeps_exact_memberships", + credential_kind="master", + actor="bootstrap", + profile="group_scoped", + method="GET", + path="/team/info", + operation_family="memberships", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_successful_actor_cleanup_removes_owned_state", + credential_kind="master", + actor="bootstrap", + profile="failure_cleanup", + method="GET", + path="/team/info", + operation_family="cleanup", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_partial_setup_removes_previously_created_identities[group]", + credential_kind="idp_admin", + actor="idp_admin", + profile="failure_cleanup", + method="POST", + path="/groups", + operation_family="cleanup", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_partial_setup_removes_previously_created_identities[user]", + credential_kind="idp_admin", + actor="idp_admin", + profile="failure_cleanup", + method="POST", + path="/users", + operation_family="cleanup", + ), + ManagementCase( + node=f"{JWT_CLASS}::test_oidc_browser_profile_identity_mapping", + credential_kind="direct_jwt", + actor="internal_user", + profile="oidc_configuration", + method="GET", + path="/protocol/openid-connect/userinfo", + operation_family="oidc_identity", + ), +) + + +def canonical_node(node: str) -> str: + return node if node.startswith("tests/e2e/") else f"tests/e2e/{node}" + + +def case_properties(node: str) -> tuple[tuple[str, str], ...]: + case: Final = next((case for case in MANAGEMENT_CASES if case.node == canonical_node(node)), None) + if case is None: + return () + return ( + ("management_node", case.node), + ("credential_kind", case.credential_kind), + ("actor", case.actor), + ("auth_profile", case.profile), + ("dependency_profile", case.dependency_profile), + ) diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index d1227fe7c0c..31ad61ba3e2 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -90,3 +90,11 @@ - {id: mgmt.mcp_toolset.update.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3098", rationale: "Narrowing the tools to one entry reads back exactly that entry"} - {id: mgmt.mcp_toolset.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "mcp_management_endpoints.py:3098", fail_before_fix: proven, rationale: "An explicit null clears the stored description; the update used to drop null and keep the old value"} - {id: mgmt.mcp_toolset.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3149", rationale: "A deleted toolset is gone by id and from the list on every replica"} + +- {id: mgmt.user.jwt.database_roles, module: mgmt, tier: P0, surface: api, assertions: [database_roles], source: "auth/handle_jwt.py", rationale: "User-only JWT subjects retain their seeded database roles and memberships"} +- {id: mgmt.key.jwt.viewer_denied, module: mgmt, tier: P0, surface: api, assertions: [viewer_denied], source: "auth/route_checks.py", rationale: "An admin viewer can read a key but cannot update it or change stored state"} +- {id: mgmt.user.oidc.identity_mapping, module: mgmt, tier: P0, surface: api, assertions: [identity_mapping], source: "tests/e2e/idp.py", rationale: "IdP configuration canary only: confidential-client token and userinfo subjects match the seeded user; application SSO is separate"} +- {id: mgmt.team.jwt.tenant_isolation, module: mgmt, tier: P0, surface: api, assertions: [tenant_isolation], source: "auth/handle_jwt.py", rationale: "Isolated team actors read their own key and receive 403 for the other tenant key"} +- {id: mgmt.team.jwt.multiple_memberships, module: mgmt, tier: P0, surface: api, assertions: [multiple_memberships], source: "auth/handle_jwt.py", rationale: "A multi-group actor has exactly the configured memberships without admin scope"} +- {id: mgmt.user.jwt.cleanup, module: mgmt, tier: P0, surface: api, assertions: [cleanup], source: "management_endpoints/internal_user_endpoints.py", rationale: "Owned users teams organizations keys and IdP objects disappear after successful cleanup"} +- {id: mgmt.user.jwt.partial_cleanup, module: mgmt, tier: P0, surface: api, assertions: [partial_cleanup], source: "auth/handle_jwt.py", rationale: "Partial identity setup removes the group and user created before failure"} diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index ce069720c6e..e0a20495964 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -16,9 +16,11 @@ requests itself imports. from __future__ import annotations import time -from collections.abc import Callable, Mapping +from collections.abc import Callable, Generator, Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar from dataclasses import dataclass -from typing import Final, Generator, Generic, Iterator, Literal, NewType, Protocol, TypeVar, cast +from typing import Final, Generic, Literal, NewType, Protocol, TypeVar, cast import pytest import requests @@ -36,8 +38,8 @@ class Headers(BaseModel): class AuthHeaders(Headers): # litellm accepts either; set whichever the call needs, leave the other None. - authorization: str | None = None - x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key") + authorization: str | None = Field(default=None, repr=False) + x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key", repr=False) class AnthropicHeaders(AuthHeaders): @@ -168,6 +170,7 @@ class StreamingResponse(BaseModel): # the consumed body is elided, so this is the only place they surface. stream_error: str | None = None stream_done: bool = False + stream_done_positions: tuple[int, ...] = () @property def ok(self) -> bool: @@ -292,6 +295,22 @@ def _params(params: BaseModel | None) -> dict[str, str]: TRANSIENT_STATUSES: frozenset[int] = frozenset({529}) RETRY_ATTEMPTS: int = 3 +_QUALIFICATION: Final[ContextVar[bool]] = ContextVar("e2e_qualification", default=False) + + +def retry_attempts(default: int) -> int: + return 1 if _QUALIFICATION.get() else default + + +@contextmanager +def without_retries() -> Generator[None]: + token: Final = _QUALIFICATION.set(True) + try: + yield + finally: + _QUALIFICATION.reset(token) + + RETRY_BACKOFF_SECONDS: float = 0.5 @@ -319,7 +338,7 @@ def request_with_retry[T: RetryableResponse]( hang should surface as a hang instead of doubling the wall clock. Every retry prints, so flakiness stays visible in the run log instead of vanishing into green.""" - for attempt in range(1, RETRY_ATTEMPTS): + for attempt in range(1, retry_attempts(RETRY_ATTEMPTS)): resp = issue() if resp.status_code not in TRANSIENT_STATUSES: return resp @@ -414,6 +433,7 @@ def get_external[R: BaseModel]( url: str, *, response_type: type[R], + headers: BaseModel | None = None, timeout: float = 30.0, ) -> Result[R]: """GET an absolute URL outside the proxy (e.g. a public /.well-known document). @@ -422,7 +442,7 @@ def get_external[R: BaseModel]( try: resp = requests.get( url, - headers={"Accept": "application/json"}, + headers={"Accept": "application/json", **(_headers(headers) if headers is not None else {})}, timeout=timeout, ) except requests.RequestException as exc: @@ -628,6 +648,7 @@ def streaming_outcome( stream_events=[payload for payload, _ in events], stream_event_arrivals=[arrived for _, arrived in events], stream_done=any(payload == _SSE_DONE for payload, _ in payloads), + stream_done_positions=tuple(index for index, (payload, _) in enumerate(payloads) if payload == _SSE_DONE), stream_error=next( (line.decode(errors="replace")[:300] for line, _ in stamped if _is_stream_error_line(line)), None, diff --git a/tests/e2e/fixture_bundle.py b/tests/e2e/fixture_bundle.py index 7c9dab1a687..4467d7e4ecc 100644 --- a/tests/e2e/fixture_bundle.py +++ b/tests/e2e/fixture_bundle.py @@ -30,9 +30,11 @@ from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Annotated, Final, Literal +from fixture_profile import MatchProfile, StrictIdentity from pydantic import BaseModel, Field, JsonValue BUNDLE_FORMAT_VERSION: Final = 4 +STRICT_BUNDLE_FORMAT_VERSION: Final = 5 MAX_BUNDLE_AGE: Final = timedelta(days=7) MANIFEST_FILENAME: Final = "manifest.json" @@ -41,6 +43,7 @@ class Manifest(BaseModel): format_version: int recorded_at: datetime harness_version: str + match_profile: MatchProfile = "legacy" class RecordedRequest(BaseModel): @@ -69,6 +72,7 @@ class RecordedRequest(BaseModel): file_name: str | None = None file_sha256: str | None = None file_bytes: int | None = None + strict_identity: StrictIdentity | None = None class RecordedHttpResponse(BaseModel): @@ -100,9 +104,7 @@ class RecordedStreamedResponse(BaseModel): truncated: str | None = None -type RecordedResponse = Annotated[ - RecordedHttpResponse | RecordedStreamedResponse, Field(discriminator="kind") -] +type RecordedResponse = Annotated[RecordedHttpResponse | RecordedStreamedResponse, Field(discriminator="kind")] class Interaction(BaseModel): @@ -152,6 +154,7 @@ class BundleRecorder: manifest, so record mode never reads (or merges into) an existing bundle.""" root: Path + profile: MatchProfile = "legacy" _ordinals: dict[str, int] = field(default_factory=dict) def record(self, *, test_key: str, request: RecordedRequest, response: RecordedResponse) -> None: @@ -162,7 +165,12 @@ class BundleRecorder: directory.mkdir(parents=True, exist_ok=True) interaction = Interaction(request=request, response=response) target = directory / interaction_filename(ordinal, request) - target.write_text(interaction.model_dump_json(indent=2), encoding="utf-8") + target.write_text( + interaction.model_dump_json( + indent=2, exclude={"request": {"strict_identity"}} if self.profile == "legacy" else None + ), + encoding="utf-8", + ) @dataclass(frozen=True, slots=True) @@ -171,7 +179,7 @@ class UnsafeBundleDir: reason: str -def prepare_bundle(root: Path) -> BundleRecorder | UnsafeBundleDir: +def prepare_bundle(root: Path, *, profile: MatchProfile = "legacy") -> BundleRecorder | UnsafeBundleDir: """Start a fresh bundle at ``root`` for record mode: wipe whatever bundle is there and write a new manifest. Refuses to wipe a directory that is neither empty nor a bundle (no manifest.json), so a mistyped E2E_FIXTURE_DIR can @@ -188,12 +196,15 @@ def prepare_bundle(root: Path) -> BundleRecorder | UnsafeBundleDir: shutil.rmtree(root) root.mkdir(parents=True) manifest = Manifest( - format_version=BUNDLE_FORMAT_VERSION, + format_version=BUNDLE_FORMAT_VERSION if profile == "legacy" else STRICT_BUNDLE_FORMAT_VERSION, + match_profile=profile, recorded_at=datetime.now(timezone.utc), harness_version=harness_version(), ) - (root / MANIFEST_FILENAME).write_text(manifest.model_dump_json(indent=2), encoding="utf-8") - return BundleRecorder(root=root) + (root / MANIFEST_FILENAME).write_text( + manifest.model_dump_json(indent=2, exclude={"match_profile"} if profile == "legacy" else None), encoding="utf-8" + ) + return BundleRecorder(root=root, profile=profile) @dataclass(frozen=True, slots=True) @@ -226,25 +237,30 @@ def _read_manifest(root: Path) -> Manifest | UnreadableBundle: return UnreadableBundle(reason=f"{MANIFEST_FILENAME} is invalid: {exc}") -def _supported_manifest(root: Path) -> Manifest | UnreadableBundle: +def _supported_manifest(root: Path, profile: MatchProfile = "legacy") -> Manifest | UnreadableBundle: """The manifest, refused when it was written under a different format version. A bundle is atomic (record wipes and rewrites the whole directory and never merges), so a foreign version is a hard reject rather than a partial read.""" manifest = _read_manifest(root) if isinstance(manifest, UnreadableBundle): return manifest - if manifest.format_version != BUNDLE_FORMAT_VERSION: + expected_version: Final = BUNDLE_FORMAT_VERSION if profile == "legacy" else STRICT_BUNDLE_FORMAT_VERSION + if manifest.match_profile != profile: + return UnreadableBundle( + reason="match profile mismatch; select the recorded E2E_REPLAY_MATCH_PROFILE or re-record" + ) + if manifest.format_version != expected_version: return UnreadableBundle( reason=( - f"format_version {manifest.format_version} != supported {BUNDLE_FORMAT_VERSION}; " + f"format_version {manifest.format_version} != supported {expected_version}; " "re-record with E2E_FIXTURE_MODE=record" ) ) return manifest -def check_freshness(root: Path, *, now: datetime) -> BundleFreshness: - manifest = _supported_manifest(root) +def check_freshness(root: Path, *, now: datetime, profile: MatchProfile = "legacy") -> BundleFreshness: + manifest = _supported_manifest(root, profile) if isinstance(manifest, UnreadableBundle): return manifest recorded_at = ( @@ -269,16 +285,27 @@ class LoadedBundle: interactions: dict[str, tuple[Interaction, ...]] -def load_bundle(root: Path) -> LoadedBundle | UnreadableBundle: - manifest = _supported_manifest(root) +def load_bundle(root: Path, *, profile: MatchProfile = "legacy") -> LoadedBundle | UnreadableBundle: + manifest = _supported_manifest(root, profile) if isinstance(manifest, UnreadableBundle): return manifest - interactions = { - directory.name: tuple( - Interaction.model_validate_json(file.read_text(encoding="utf-8")) - for file in sorted(directory.glob("*.json")) - ) - for directory in sorted(root.iterdir()) - if directory.is_dir() - } + try: + interactions = { + directory.name: tuple( + Interaction.model_validate_json(file.read_text(encoding="utf-8")) + for file in sorted(directory.glob("*.json")) + ) + for directory in sorted(root.iterdir()) + if directory.is_dir() + } + except (ValueError, OSError): + if profile == "legacy": + raise + return UnreadableBundle(reason="invalid stateless_v1 interaction; re-record with the selected profile") + if any( + (item.request.strict_identity is not None) != (profile == "stateless_v1") + for items in interactions.values() + for item in items + ): + return UnreadableBundle(reason="request identity/profile mismatch; re-record with the selected profile") return LoadedBundle(manifest=manifest, interactions=interactions) diff --git a/tests/e2e/fixture_canonical.py b/tests/e2e/fixture_canonical.py index c043951a108..e76d63ca33b 100644 --- a/tests/e2e/fixture_canonical.py +++ b/tests/e2e/fixture_canonical.py @@ -23,9 +23,8 @@ from dataclasses import dataclass from functools import reduce from typing import Final -from pydantic import JsonValue - from fixture_bundle import RecordedRequest +from pydantic import JsonValue VOLATILE_HEADER_NAMES: Final[frozenset[str]] = frozenset( { @@ -123,6 +122,12 @@ class CanonicalRequest: def canonicalize(request: RecordedRequest) -> CanonicalRequest: + if request.strict_identity is not None: + return CanonicalRequest( + method=request.method, + path=request.path, + content=json.dumps(request.strict_identity.model_dump(mode="json"), sort_keys=True, separators=(",", ":")), + ) file_identity: Final[JsonValue | None] = ( None if request.file_name is None and request.file_sha256 is None diff --git a/tests/e2e/fixture_mode.py b/tests/e2e/fixture_mode.py index 110f44380b4..9a7c1b6db12 100644 --- a/tests/e2e/fixture_mode.py +++ b/tests/e2e/fixture_mode.py @@ -26,6 +26,7 @@ from fixture_bundle import ( check_freshness, format_age, ) +from fixture_profile import match_profile type FixtureMode = Literal["live", "record", "replay"] @@ -82,6 +83,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet Called at collection time (conftest pytest_sessionstart) so a stale or missing bundle fails the whole run up front, naming the bundle age, instead of failing every test individually.""" + match_profile() mode = parse_fixture_mode(mode_raw) match mode: case InvalidFixtureMode(value=value): @@ -89,7 +91,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet case "live" | "record": return None case "replay": - freshness = check_freshness(bundle_dir, now=now) + freshness = check_freshness(bundle_dir, now=now, profile=match_profile()) match freshness: case FreshBundle(): return None @@ -110,6 +112,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> list[str]: """pytest report-header lines; empty in live mode so an unset E2E_FIXTURE_MODE keeps today's output byte-identical.""" + match_profile() mode = parse_fixture_mode(mode_raw) match mode: case InvalidFixtureMode() | "live": @@ -117,7 +120,7 @@ def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> l case "record": return [f"e2e fixture mode: record -> {bundle_dir}"] case "replay": - freshness = check_freshness(bundle_dir, now=now) + freshness = check_freshness(bundle_dir, now=now, profile=match_profile()) match freshness: case FreshBundle(manifest=manifest): return [ diff --git a/tests/e2e/fixture_profile.py b/tests/e2e/fixture_profile.py new file mode 100644 index 00000000000..f8405be746b --- /dev/null +++ b/tests/e2e/fixture_profile.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import json +import os +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Final, Literal +from urllib.parse import parse_qsl, urlsplit + +from pydantic import BaseModel, JsonValue, TypeAdapter + +type MatchProfile = Literal["legacy", "stateless_v1"] + + +@dataclass(frozen=True, slots=True) +class NumberToken: + literal: str + + +type ExactJson = dict[str, ExactJson] | list[ExactJson] | str | bool | NumberToken | None + +SEMANTIC_HEADERS: Final = frozenset({"content-type", "accept", "anthropic-version", "anthropic-beta", "openai-beta"}) +AUTH_HEADERS: Final = frozenset({"authorization", "x-api-key"}) +EXCLUDED_HEADERS: Final = frozenset( + { + "host", + "content-length", + "transfer-encoding", + "connection", + "accept-encoding", + "user-agent", + "traceparent", + "tracestate", + "x-request-id", + "x-client-request-id", + "cookie", + } +) +CREDENTIAL_QUERY: Final = frozenset( + { + "api_key", + "api-key", + "apikey", + "key", + "token", + "access_token", + "signature", + "password", + "secret", + "credentials", + "authorization", + "sig", + "client_secret", + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + } +) +JSON_VALUE: Final[TypeAdapter[ExactJson]] = TypeAdapter(ExactJson) + + +def match_profile() -> MatchProfile: + raw: Final = os.environ.get("E2E_REPLAY_MATCH_PROFILE", "legacy") + if raw in ("legacy", "stateless_v1"): + return raw + raise ValueError("E2E_REPLAY_MATCH_PROFILE must be legacy or stateless_v1") + + +class StrictIdentity(BaseModel): + upstream: str + mount: str + query: tuple[tuple[str, str], ...] + headers: dict[str, str] + auth: dict[str, str] + body_present: bool + body: JsonValue + + +@dataclass(frozen=True, slots=True) +class IneligibleRequest: + reason: str + + +def _unique_object(pairs: list[tuple[str, ExactJson]]) -> dict[str, ExactJson]: + if len({key for key, _ in pairs}) != len(pairs): + raise ValueError("duplicate JSON object keys") + return dict(pairs) + + +def _invalid_constant(value: str) -> ExactJson: + raise ValueError("nonfinite JSON number") + + +def _exact_value(value: ExactJson) -> JsonValue: + match value: + case dict(): + return {"object": {key: _exact_value(item) for key, item in value.items()}} + case list(): + return {"array": [_exact_value(item) for item in value]} + case bool(): + return {"boolean": value} + case NumberToken(literal=literal): + return {"number": literal} + case str(): + return {"string": value} + case None: + return None + + +def strict_identity( + *, + method: str, + path: str, + query: str, + headers: Mapping[str, str], + body: bytes | None, + mount: str, + upstream_base: str, +) -> StrictIdentity | IneligibleRequest: + if (mount, path, method.upper()) not in { + ("openai", "/openai/v1/chat/completions", "POST"), + ("anthropic", "/anthropic/v1/messages", "POST"), + }: + return IneligibleRequest("unsupported endpoint or method") + lowered: Final = {key.lower(): value for key, value in headers.items()} + if len(lowered) != len(headers): + return IneligibleRequest("duplicate header names") + if any( + key not in SEMANTIC_HEADERS | AUTH_HEADERS | EXCLUDED_HEADERS and not key.startswith("x-stainless-") + for key in lowered + ): + return IneligibleRequest("unsupported semantic header") + if "transfer-encoding" in lowered: + return IneligibleRequest("unsupported request transfer-encoding; send a content-length framed JSON body") + authorization: Final = lowered.get("authorization") + if authorization is not None and authorization.partition(" ")[0].lower() not in {"bearer", "basic", "digest"}: + return IneligibleRequest("unsupported authorization scheme") + destination: Final = urlsplit(upstream_base) + if destination.username or destination.password or destination.query or destination.fragment: + return IneligibleRequest("upstream destination contains credentials, query or fragment") + if destination.scheme not in ("http", "https") or not destination.netloc: + return IneligibleRequest("unsupported upstream destination") + if body and lowered.get("content-type", "").split(";", 1)[0].strip().lower() != "application/json": + return IneligibleRequest("unsupported body content-type; stateless_v1 requires JSON") + try: + parsed: Final = ( + JSON_VALUE.validate_python( + json.loads( + body, + object_pairs_hook=_unique_object, + parse_constant=_invalid_constant, + parse_float=NumberToken, + parse_int=NumberToken, + ) + ) + if body + else None + ) + except (ValueError, UnicodeError): + return IneligibleRequest("invalid JSON or duplicate JSON object keys") + if body and not isinstance(parsed, dict): + return IneligibleRequest("stateless inference requires a JSON object") + try: + query_pairs: Final = tuple(parse_qsl(query, keep_blank_values=True, errors="strict")) + except UnicodeError: + return IneligibleRequest("invalid UTF-8 query encoding") + return StrictIdentity( + upstream=upstream_base, + mount=mount, + query=tuple((key, "" if key.lower() in CREDENTIAL_QUERY else value) for key, value in query_pairs), + headers={key: value for key, value in lowered.items() if key in SEMANTIC_HEADERS}, + auth={ + key: (value.partition(" ")[0].lower() if key == "authorization" else "present") + for key, value in lowered.items() + if key in AUTH_HEADERS + }, + body_present=bool(body), + body=_exact_value(parsed), + ) diff --git a/tests/e2e/idp.py b/tests/e2e/idp.py index 6d2fc84eb27..2dc7c2ad71b 100644 --- a/tests/e2e/idp.py +++ b/tests/e2e/idp.py @@ -2,11 +2,18 @@ from __future__ import annotations +import base64 import os import secrets +import signal +import subprocess +import sys +import time import warnings from collections.abc import Callable -from dataclasses import dataclass, field +from contextlib import ExitStack +from dataclasses import dataclass, field, replace +from types import FrameType from typing import Final, Literal import pytest @@ -14,11 +21,15 @@ from e2e_http import ( AuthHeaders, ExternalWrite, NetworkError, + NoBody, Result, Success, + UnknownApiError, delete_external, + get_external, post_form_external, post_json_external, + unwrap, ) from pydantic import BaseModel, Field @@ -46,7 +57,9 @@ class TokenGrantForm(BaseModel): grant_type: Literal["password"] = "password" client_id: str username: str - password: str + password: str = Field(repr=False) + client_secret: str | None = Field(default=None, repr=False) + scope: str | None = None class TokenResponse(BaseModel): @@ -63,7 +76,7 @@ class GroupCreateBody(BaseModel): class PasswordCredential(BaseModel): type: Literal["password"] = "password" - value: str + value: str = Field(repr=False) temporary: bool = False @@ -101,8 +114,20 @@ class Identity: user_id: str username: str password: str = field(repr=False) - group: str - group_id: str + groups: tuple[str, ...] + group_ids: tuple[str, ...] + + @property + def group(self) -> str: + if len(self.groups) != 1: + raise ValueError("A single-group identity is required") + return self.groups[0] + + @property + def group_id(self) -> str: + if len(self.group_ids) != 1: + raise ValueError("A single-group identity is required") + return self.group_ids[0] @dataclass(frozen=True, slots=True) @@ -111,6 +136,10 @@ class Keycloak: realm: str admin_username: str admin_password: str = field(repr=False) + strict_cleanup: bool = False + + def with_strict_cleanup(self) -> Keycloak: + return replace(self, strict_cleanup=True) @property def issuer(self) -> str: @@ -150,7 +179,9 @@ class Keycloak: f"group {name}", ) - def create_user(self, *, username: str, email: str, password: str, group: str) -> str: + def create_user( + self, *, username: str, email: str, password: str, group: str | None = None, groups: tuple[str, ...] = () + ) -> str: return created_id( post_json_external( self._admin_url("/users"), @@ -158,7 +189,7 @@ class Keycloak: json=UserCreateBody( username=username, email=email, - groups=(group,), + groups=(group,) if group is not None else groups, credentials=(PasswordCredential(value=password),), ), ), @@ -171,14 +202,28 @@ class Keycloak: def delete_group(self, group_id: str) -> None: self._delete(f"/groups/{group_id}") + def assert_absent(self, kind: Literal["users", "groups", "clients"], resource_id: str) -> None: + result: Final = get_external( + self._admin_url(f"/{kind}/{resource_id}"), + headers=self._admin_headers(), + response_type=NoBody, + ) + assert isinstance(result, UnknownApiError) and result.status_code == 404, ( + f"Owned IdP {kind} still exists: {result}" + ) + def _delete(self, path: str) -> None: try: headers: Final = self._admin_headers() except pytest.fail.Exception as exc: + if self.strict_cleanup: + raise RuntimeError(f"Keycloak cleanup could not authenticate for {path}") from exc warnings.warn(f"Keycloak cleanup could not authenticate for {path}: {exc}", RuntimeWarning, stacklevel=2) return result: Final = delete_external(self._admin_url(path), headers=headers) if result.status_code not in (204, 404): + if self.strict_cleanup: + raise RuntimeError(f"Keycloak cleanup failed for {path}: HTTP {result.status_code}") warnings.warn( f"Keycloak cleanup failed for {path}: HTTP {result.status_code} {result.body[:300]}", RuntimeWarning, @@ -188,15 +233,34 @@ class Keycloak: def provision(self, *, marker: str, group: str, defer: Callable[[Callable[[], object]], None]) -> Identity: """Create `group` and a user in it, credentialed with a password generated for this test alone, and hand back the identity a token can be minted for.""" - group_id: Final = self.create_group(group) - defer(lambda: self.delete_group(group_id)) + return self.provision_groups(marker=marker, groups=(group,), defer=defer) + + def provision_groups( + self, *, marker: str, groups: tuple[str, ...], defer: Callable[[Callable[[], object]], None] + ) -> Identity: + def provision_group(name: str) -> str: + created: Final = self.create_group(name) + defer(lambda: self.delete_group(created)) + return created + + group_ids: Final = tuple(provision_group(group) for group in groups) + return self.provision_user(marker=marker, groups=groups, group_ids=group_ids, defer=defer) + + def provision_user( + self, + *, + marker: str, + groups: tuple[str, ...], + group_ids: tuple[str, ...], + defer: Callable[[Callable[[], object]], None], + ) -> Identity: username: Final = f"e2e-jwt-user-{marker}" password: Final = secrets.token_urlsafe(24) user_id: Final = self.create_user( - username=username, email=f"{username}@example.com", password=password, group=group + username=username, email=f"{username}@example.com", password=password, groups=groups ) defer(lambda: self.delete_user(user_id)) - return Identity(user_id=user_id, username=username, password=password, group=group, group_id=group_id) + return Identity(user_id=user_id, username=username, password=password, groups=groups, group_ids=group_ids) def access_token( self, identity: Identity, *, client_id: str = TESTS_CLIENT_ID, issuer_host: str | None = None @@ -211,6 +275,65 @@ class Keycloak: ) return self._token(result, f"a token for {identity.username}") + def discovery(self) -> Discovery: + return unwrap(get_external(f"{self.issuer}/.well-known/openid-configuration", response_type=Discovery)) + + def browser_client(self, *, callback_url: str, defer: Callable[[Callable[[], object]], None]) -> BrowserClient: + client: Final = BrowserClient( + client_id=f"e2e-browser-{secrets.token_hex(8)}", + secret=secrets.token_urlsafe(32), + callback_url=callback_url, + ) + resource_id: Final = created_id( + post_json_external( + self._admin_url("/clients"), + headers=self._admin_headers(), + json=BrowserClientBody( + clientId=client.client_id, + secret=client.secret, + redirectUris=(callback_url,), + ), + ), + "browser client", + ) + defer(lambda: self._delete(f"/clients/{resource_id}")) + configured: Final = unwrap( + get_external( + self._admin_url(f"/clients/{resource_id}"), + headers=self._admin_headers(), + response_type=BrowserClientBody, + ) + ) + assert configured.redirect_uris == (callback_url,) + assert configured.standard_flow_enabled and not configured.public_client + assert configured.attributes.pkce == "S256" + return client + + def browser_token(self, identity: Identity, client: BrowserClient) -> str: + return self._token( + post_form_external( + self.token_url(self.realm), + form=TokenGrantForm( + client_id=client.client_id, + client_secret=client.secret, + username=identity.username, + password=identity.password, + scope="openid email", + ), + response_type=TokenResponse, + ), + "browser-profile identity mapping", + ) + + def userinfo(self, token: str) -> UserInfo: + return unwrap( + get_external( + f"{self.issuer}/protocol/openid-connect/userinfo", + headers=AuthHeaders(authorization=f"Bearer {token}"), + response_type=UserInfo, + ) + ) + def keycloak_from_env() -> Keycloak: admin_username: Final = os.environ.get(KEYCLOAK_ADMIN_USER_ENV, "").strip() @@ -226,3 +349,137 @@ def keycloak_from_env() -> Keycloak: admin_username=admin_username, admin_password=admin_password, ) + + +class TokenClaims(BaseModel): + sub: str + iss: str + aud: str | tuple[str, ...] + exp: int + scope: str = "" + groups: tuple[str, ...] = () + + +class Discovery(BaseModel): + issuer: str + authorization_endpoint: str + token_endpoint: str + userinfo_endpoint: str + jwks_uri: str + + +class UserInfo(BaseModel): + sub: str + email: str + + +class BrowserAttributes(BaseModel): + pkce: str = Field(default="S256", alias="pkce.code.challenge.method") + + +class AudienceConfig(BaseModel): + audience: str = Field(default="litellm-e2e", alias="included.custom.audience") + access_token: str = Field(default="true", alias="access.token.claim") + id_token: str = Field(default="false", alias="id.token.claim") + + +class AudienceMapper(BaseModel): + name: str = "litellm-audience" + protocol: str = "openid-connect" + mapper: str = Field(default="oidc-audience-mapper", alias="protocolMapper") + config: AudienceConfig = Field(default_factory=AudienceConfig) + + +class BrowserClientBody(BaseModel): + client_id: str = Field(alias="clientId") + secret: str = Field(repr=False) + redirect_uris: tuple[str, ...] = Field(alias="redirectUris") + enabled: bool = True + public_client: bool = Field(default=False, alias="publicClient") + standard_flow_enabled: bool = Field(default=True, alias="standardFlowEnabled") + direct_access_grants_enabled: bool = Field(default=True, alias="directAccessGrantsEnabled") + default_client_scopes: tuple[str, ...] = Field(default=("email", "basic"), alias="defaultClientScopes") + attributes: BrowserAttributes = Field(default_factory=BrowserAttributes) + protocol_mappers: tuple[AudienceMapper, ...] = Field(default=(AudienceMapper(),), alias="protocolMappers") + + +@dataclass(frozen=True, slots=True) +class BrowserClient: + client_id: str + secret: str = field(repr=False) + callback_url: str + + def environment(self, discovery: Discovery) -> dict[str, str]: + return { + "GENERIC_CLIENT_ID": self.client_id, + "GENERIC_CLIENT_SECRET": self.secret, + "GENERIC_USER_ID_ATTRIBUTE": "sub", + "GENERIC_AUTHORIZATION_ENDPOINT": discovery.authorization_endpoint, + "GENERIC_TOKEN_ENDPOINT": discovery.token_endpoint, + "GENERIC_USERINFO_ENDPOINT": discovery.userinfo_endpoint, + "GENERIC_CLIENT_USE_PKCE": "true", + "GENERIC_SCOPE": "openid email", + } + + +def token_claims(token: str) -> TokenClaims: + payload: Final = token.split(".")[1] + return TokenClaims.model_validate_json(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))) + + +def _signal_process_group(process_id: int, signum: int) -> bool: + try: + os.killpg(process_id, signum) + except ProcessLookupError: + return False + return True + + +def _stop_process_group(child: subprocess.Popen[bytes]) -> None: + _signal_process_group(child.pid, signal.SIGTERM) + deadline: Final = time.monotonic() + 5 + while _process_group_exists(child.pid): + child.poll() + if time.monotonic() >= deadline: + _signal_process_group(child.pid, signal.SIGKILL) + break + time.sleep(0.05) + child.wait() + + +def _process_group_exists(process_id: int) -> bool: + try: + os.killpg(process_id, 0) + except ProcessLookupError: + return False + except PermissionError: + return True + return True + + +def run_oidc_profile(proxy_url: str, command: list[str]) -> int: + idp: Final = keycloak_from_env().with_strict_cleanup() + with ExitStack() as cleanup: + + def terminate(signum: int, frame: FrameType | None) -> None: + raise SystemExit(128 + signum) + + previous: Final = signal.signal(signal.SIGTERM, terminate) + cleanup.callback(signal.signal, signal.SIGTERM, previous) + + def defer(callback: Callable[[], object]) -> None: + cleanup.callback(callback) + + client: Final = idp.browser_client(callback_url=f"{proxy_url.rstrip('/')}/sso/callback", defer=defer) + environment: Final = {**os.environ, **client.environment(idp.discovery()), "PROXY_BASE_URL": proxy_url} + with subprocess.Popen(command, env=environment, start_new_session=True) as child: + try: + return child.wait() + finally: + _stop_process_group(child) + + +if __name__ == "__main__": + if len(sys.argv) < 3: + raise SystemExit("Usage: idp.py PROXY_URL COMMAND [ARG ...]; requires a running test IdP") + raise SystemExit(run_oidc_profile(sys.argv[1], sys.argv[2:])) diff --git a/tests/e2e/junit_properties.py b/tests/e2e/junit_properties.py index c5971c5362c..b9f5da871ae 100644 --- a/tests/e2e/junit_properties.py +++ b/tests/e2e/junit_properties.py @@ -19,6 +19,7 @@ from __future__ import annotations from collections.abc import Iterable import pytest +from coverage_registry.management_cases import case_properties # Hardcoded because the runner image copies tests/e2e/ to /app/e2e, so nothing # at runtime names this suite's place in the repo. test_junit_properties.py @@ -94,7 +95,7 @@ def result_properties(item: pytest.Item) -> tuple[tuple[str, str], ...]: ("package", package_from_nodeid(item.nodeid)), ("covers", ",".join(covers_from_item(item))), ("source", source_from_item(item)), - ) + ) + case_properties(item.nodeid) def attach_result_properties(item: pytest.Item) -> None: diff --git a/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py b/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py index 4db2fe004c5..fdb76df703d 100644 --- a/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py +++ b/tests/e2e/llm_translation/test_chat_stream_contract_e2e.py @@ -1,51 +1,90 @@ -"""Vendor §12.3: chat completions streaming SSE contract (LIT-4778). - -Asserts a streamed /chat/completions response is SSE, carries content chunks, -and terminates with the OpenAI [DONE] sentinel. -""" - from __future__ import annotations +from typing import Final + import pytest -from e2e_config import unique_marker +from e2e_config import provider_edge_base, unique_marker from e2e_http import require_successful_call from lifecycle import ResourceManager -from models import ChatBody, ChatMessage, LiteLLMParamsBody +from models import ChatBody, ChatMessage, ChatStreamOptions, LiteLLMParamsBody, Usage from proxy_client import ProxyClient +from pydantic import BaseModel -pytestmark = pytest.mark.e2e +pytestmark = [pytest.mark.e2e, pytest.mark.replayable] + + +class _Delta(BaseModel): + content: str | None = None + + +class _Choice(BaseModel): + index: int + delta: _Delta + finish_reason: str | None = None + + +class _Chunk(BaseModel): + choices: tuple[_Choice, ...] + usage: Usage | None = None class TestChatStreamContract: @pytest.mark.covers("llm.chat_completions.openai.basic.stream.works") def test_chat_stream_is_sse_and_ends_with_done(self, proxy: ProxyClient, resources: ResourceManager) -> None: - model = f"e2e-chat-stream-{unique_marker()}" - model_id = proxy.create_model( + model: Final = f"e2e-chat-stream-{unique_marker()}" + base: Final = provider_edge_base("openai") + model_id: Final = proxy.create_model( model, - LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY"), + LiteLLMParamsBody( + model="openai/gpt-5.6", + api_key="os.environ/OPENAI_API_KEY", + api_base=f"{base}/v1" if base else None, + ), ) resources.defer(lambda: proxy.delete_model(model_id)) - key = resources.key() - - result = proxy.chat_stream( + key: Final = resources.key() + expected: Final = "The amber kite crosses the quiet lake." + result: Final = proxy.chat_stream( key, ChatBody( model=model, messages=[ ChatMessage( - role="user", - content=f"Reply with the single word ok. {unique_marker()}", + role="user", content=f"Repeat exactly this sentence, with no additional text: {expected}" ) ], stream=True, - max_completion_tokens=32, - temperature=0.0, + stream_options=ChatStreamOptions(include_usage=True), + max_completion_tokens=256, + reasoning_effort="none", ), ) require_successful_call(result) assert result.is_streaming, f"expected SSE content-type, got {result.content_type!r}" assert result.stream_events, "stream returned no data events" - assert result.stream_done, ( - f"stream must terminate with [DONE]; " - f"chunks={result.chunks} done={result.stream_done} events={len(result.stream_events)}" + assert not result.stream_error, f"stream errored: {result.stream_error}" + assert result.stream_done, "stream must terminate with [DONE]" + assert result.stream_done_positions == (len(result.stream_events),), "[DONE] must occur once after all events" + chunks: Final = tuple(_Chunk.model_validate_json(event) for event in result.stream_events) + text_positions: Final = tuple( + i for i, chunk in enumerate(chunks) if any(c.delta.content for c in chunk.choices) ) + terminal_positions: Final = tuple( + i for i, chunk in enumerate(chunks) if any(c.finish_reason is not None for c in chunk.choices) + ) + assert text_positions, "stream completed without meaningful text" + assert len(terminal_positions) == 1, "expected exactly one terminal choice" + assert text_positions[0] < terminal_positions[0], "meaningful text must arrive before termination" + assert text_positions[-1] <= terminal_positions[0], "text arrived after termination" + assert all(c.index == 0 for chunk in chunks for c in chunk.choices) + assert tuple(c.finish_reason for c in chunks[terminal_positions[0]].choices) == ("stop",) + text: Final = "".join(c.delta.content or "" for chunk in chunks for c in chunk.choices) + assert text.strip() == expected, f"streamed answer was altered or incomplete: {text!r}" + usage_positions: Final = tuple(i for i, chunk in enumerate(chunks) if chunk.usage is not None) + assert usage_positions == (len(chunks) - 1,), "expected one final usage chunk" + assert terminal_positions[0] < usage_positions[0], "usage must follow the terminal choice" + usage: Final = chunks[-1].usage + assert usage is not None + assert usage.prompt_tokens is not None and usage.prompt_tokens > 0 + assert usage.completion_tokens is not None and usage.completion_tokens > 0 + assert usage.total_tokens == usage.prompt_tokens + usage.completion_tokens diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index c731b52acd8..ca58c30d40c 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -21,7 +21,12 @@ from e2e_http import assert_client_error, require_successful_call, unwrap from endpoints_client import EndpointsClient, MessagesResult from lifecycle import ResourceManager from models import ( + AnthropicAssistantTurn, + AnthropicContentBlock, AnthropicCustomTool, + AnthropicToolChoice, + AnthropicToolResultBlock, + AnthropicToolResultTurn, AnthropicMessagesBody, ChatMessage, JsonSchemaProperty, @@ -29,7 +34,7 @@ from models import ( SpendLogRow, ToolInputSchema, ) -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict pytestmark = [pytest.mark.e2e, pytest.mark.replayable] @@ -284,8 +289,139 @@ class TestAnthropicMessages: result = endpoints_client.proxy.transport.send( "/v1/messages", headers=endpoints_client.proxy.transport.bearer(key), - json=_OptionalMessagesBody( - messages=[ChatMessage(role="user", content="hi")], max_tokens=50 - ), + json=_OptionalMessagesBody(messages=[ChatMessage(role="user", content="hi")], max_tokens=50), ) assert_client_error(result, "messages missing model") + + +class _BridgeDelta(BaseModel): + type: str | None = None + partial_json: str | None = None + stop_reason: str | None = None + + +class _BridgeEvent(BaseModel): + type: str + index: int | None = None + content_block: AnthropicContentBlock | None = None + delta: _BridgeDelta | None = None + + +class _ParcelInput(BaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + parcel: str + shelf: int + + +def _tool_from_stream(events: tuple[_BridgeEvent, ...]) -> AnthropicContentBlock: + starts: Final = tuple( + event + for event in events + if event.type == "content_block_start" + and event.content_block is not None + and event.content_block.type == "tool_use" + ) + assert len(starts) == 1, "expected exactly one tool call" + start: Final = starts[0] + block: Final = start.content_block + assert block is not None and block.id and start.index is not None + fragments: Final = tuple( + event + for event in events + if event.type == "content_block_delta" and event.delta is not None and event.delta.type == "input_json_delta" + ) + assert fragments, "tool stream contained no argument fragments" + assert all(event.index == start.index for event in fragments), "tool fragments changed index" + positions: Final = tuple(i for i, event in enumerate(events) if event in fragments) + stops: Final = tuple( + i for i, event in enumerate(events) if event.type == "content_block_stop" and event.index == start.index + ) + assert len(stops) == 1 and events.index(start) < positions[0] <= positions[-1] < stops[0] + assert tuple( + event.delta.stop_reason for event in events if event.type == "message_delta" and event.delta is not None + ) == ("tool_use",) + terminal_positions: Final = tuple(i for i, event in enumerate(events) if event.type == "message_delta") + assert len(terminal_positions) == 1 and stops[0] < terminal_positions[0] < len(events) - 1 + assert tuple(i for i, event in enumerate(events) if event.type == "message_stop") == (len(events) - 1,), ( + "tool stream did not terminate exactly once" + ) + arguments: Final = _ParcelInput.model_validate_json( + "".join(event.delta.partial_json or "" for event in fragments if event.delta is not None) + ) + return AnthropicContentBlock(type="tool_use", id=block.id, name=block.name, input=arguments.model_dump()) + + +def _parcel_result(tool: AnthropicContentBlock, result: AnthropicToolResultBlock) -> AnthropicToolResultTurn: + assert tool.id and result.tool_use_id == tool.id, "tool result ID does not match the emitted call" + return AnthropicToolResultTurn(content=[result]) + + +def _request_tool( + client: EndpointsClient, key: str, request: AnthropicMessagesBody, stream: bool +) -> AnthropicContentBlock: + if stream: + response: Final = client.proxy.messages_stream(key, request) + require_successful_call(response) + assert response.is_streaming and not response.stream_error + return _tool_from_stream(tuple(_BridgeEvent.model_validate_json(event) for event in response.stream_events)) + response_body: Final = unwrap(client.proxy.messages(key, request)) + blocks: Final = tuple(block for block in response_body.content or () if block.type == "tool_use") + assert len(blocks) == 1 + return blocks[0] + + +class TestOpenAIMessagesToolContinuation: + @pytest.mark.parametrize("stream", [True, False], ids=["stream", "nonstream"]) + def test_required_tool_arguments_and_correlated_result( + self, endpoints_client: EndpointsClient, resources: ResourceManager, stream: bool + ) -> None: + model: Final = f"e2e-bridge-tool-{unique_marker()}" + base: Final = provider_edge_base("openai") + model_id: Final = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-5.6", api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key: Final = resources.key(models=[model]) + tool: Final = AnthropicCustomTool( + name="locate_parcel", + description="Look up the receipt for a parcel on a shelf. Return the receipt verbatim.", + input_schema=ToolInputSchema( + properties={"parcel": JsonSchemaProperty(type="string"), "shelf": JsonSchemaProperty(type="integer")}, + required=["parcel", "shelf"], + ), + ) + question: Final = ChatMessage( + role="user", + content="Call locate_parcel with parcel exactly amber-kite and shelf exactly 7. After the tool result, reply with only the receipt returned by the tool.", + ) + request: Final = AnthropicMessagesBody( + model=model, + max_tokens=2048, + messages=[question], + tools=[tool], + tool_choice=AnthropicToolChoice(type="tool", name=tool.name), + stream=stream, + ) + emitted: Final = _request_tool(endpoints_client, key, request, stream) + assert emitted.id and emitted.name == "locate_parcel" + assert emitted.input == {"parcel": "amber-kite", "shelf": 7}, "required tool arguments were lost or changed" + receipt: Final = f"receipt-{unique_marker()}" + result_turn: Final = _parcel_result(emitted, AnthropicToolResultBlock(tool_use_id=emitted.id, content=receipt)) + continuation: Final = unwrap( + endpoints_client.proxy.messages( + key, + AnthropicMessagesBody( + model=model, + max_tokens=2048, + tools=[tool], + tool_choice=AnthropicToolChoice(type="none"), + messages=[question, AnthropicAssistantTurn(content=[emitted]), result_turn], + ), + ) + ) + answer: Final = "".join(block.text or "" for block in continuation.content or ()) + assert answer.strip() == receipt, "continuation did not consume the correlated tool result" + assert all(block.type != "tool_use" for block in continuation.content or ()) diff --git a/tests/e2e/management/conftest.py b/tests/e2e/management/conftest.py index bd69c8c0ff3..5a11b634085 100644 --- a/tests/e2e/management/conftest.py +++ b/tests/e2e/management/conftest.py @@ -5,8 +5,14 @@ holds the shared ProxyClient so `resources` / `scoped_key` clean up keys, teams, users, and orgs this suite creates. """ -import pytest +from collections.abc import Generator +from typing import Final +import pytest +from e2e_http import without_retries +from idp import Keycloak +from lifecycle import ResourceManager +from management.jwt_actors import ActorFactory from management_client import ManagementClient, build_client from proxy_client import ProxyClient @@ -21,3 +27,14 @@ def pytest_configure(config: pytest.Config) -> None: @pytest.fixture(scope="session") def client(proxy: ProxyClient) -> ManagementClient: return build_client(proxy) + + +@pytest.fixture +def actor_factory(proxy: ProxyClient, idp: Keycloak) -> Generator[ActorFactory]: + bootstrap: Final = build_client(proxy) + resources: Final = ResourceManager(client=proxy, strict_cleanup=True) + with without_retries(): + try: + yield ActorFactory(bootstrap=bootstrap, idp=idp, resources=resources) + finally: + resources.teardown() diff --git a/tests/e2e/management/jwt_actors.py b/tests/e2e/management/jwt_actors.py new file mode 100644 index 00000000000..2d23549fe71 --- /dev/null +++ b/tests/e2e/management/jwt_actors.py @@ -0,0 +1,175 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Final, Literal + +from e2e_config import unique_marker +from e2e_http import NoBody, unwrap +from idp import ADMIN_CLIENT_ID, TESTS_CLIENT_ID, Identity, Keycloak +from lifecycle import ResourceManager +from management.management_client import ManagementClient +from models import ( + KeyGenerateBody, + KeyGenerateResponse, + OrgDeleteBody, + OrgDeleteResponse, + OrgMemberAddBody, + OrgMemberEntry, + OrgNewBody, + TeamDeleteBody, + TeamMemberAddBody, + TeamMemberEntry, + TeamNewBody, + UserNewBody, + UserRole, +) +from proxy_client import Caller + +ActorRole = Literal[ + "proxy_admin", + "proxy_admin_viewer", + "organization_admin", + "team_admin", + "team_member", + "internal_user", + "internal_user_viewer", + "unrelated_user", +] +ActorProfile = Literal["database_role", "group_scoped"] + + +@dataclass(frozen=True, slots=True) +class Tenant: + organization_id: str + team_id: str + group_id: str + + +@dataclass(frozen=True, slots=True) +class Actor: + identity: Identity + role: ActorRole + global_role: UserRole + profile: ActorProfile + tenants: tuple[Tenant, ...] + + def mint_caller(self, idp: Keycloak) -> Caller: + return Caller( + credential=idp.access_token( + self.identity, client_id=ADMIN_CLIENT_ID if self.role == "proxy_admin" else TESTS_CLIENT_ID + ), + kind="direct_jwt", + role=self.role, + tenant=self.tenants[0].organization_id if self.tenants else None, + ) + + +@dataclass(frozen=True, slots=True) +class ActorFactory: + bootstrap: ManagementClient + idp: Keycloak + resources: ResourceManager + + def __post_init__(self) -> None: + if self.bootstrap.proxy.caller is not None: + raise ValueError("Actor bootstrap requires a separately held master client") + + def key(self, tenant: Tenant | None = None, *, user_id: str | None = None) -> KeyGenerateResponse: + created: Final = unwrap( + self.bootstrap.generate_key( + KeyGenerateBody( + team_id=tenant.team_id if tenant is not None else None, + user_id=user_id, + key_alias=f"e2e-actor-key-{unique_marker()}", + ) + ) + ) + self.resources.defer(lambda: self.bootstrap.delete_key_strict(created.key, missing_ok=True)) + return created + + def tenant(self) -> Tenant: + marker: Final = unique_marker() + organization_id: Final = self.bootstrap.create_org(OrgNewBody(organization_alias=f"e2e-organization-{marker}")) + self.resources.defer( + lambda: unwrap( + self.bootstrap.proxy.transport.delete( + "/organization/delete", + headers=self.bootstrap.proxy.management_headers(), + json=OrgDeleteBody(organization_ids=[organization_id]), + response_type=OrgDeleteResponse, + ) + ) + ) + team_id: Final = self.bootstrap.proxy.create_team( + TeamNewBody(team_alias=f"e2e-team-{marker}", organization_id=organization_id) + ) + self.resources.defer( + lambda: unwrap( + self.bootstrap.proxy.transport.post( + "/team/delete", + headers=self.bootstrap.proxy.management_headers(), + json=TeamDeleteBody(team_ids=[team_id]), + response_type=NoBody, + ) + ) + ) + self.bootstrap.delete_team_member(team_id, self.bootstrap.user_info().user_id) + group_id: Final = self.idp.create_group(team_id) + self.resources.defer(lambda: self.idp.with_strict_cleanup().delete_group(group_id)) + return Tenant(organization_id=organization_id, team_id=team_id, group_id=group_id) + + def create( + self, role: ActorRole, *, tenants: tuple[Tenant, ...] = (), profile: ActorProfile = "database_role" + ) -> Actor: + if role in ("team_admin", "team_member", "organization_admin") and not tenants: + raise ValueError("A membership actor requires a tenant") + identity: Final = self.idp.with_strict_cleanup().provision_user( + marker=unique_marker(), + groups=tuple(tenant.team_id for tenant in tenants) if profile == "group_scoped" else (), + group_ids=tuple(tenant.group_id for tenant in tenants) if profile == "group_scoped" else (), + defer=self.resources.defer, + ) + global_role: Final[UserRole] = ( + role + if role in ("proxy_admin", "proxy_admin_viewer", "internal_user", "internal_user_viewer") + else "internal_user" + ) + self.bootstrap.create_user( + UserNewBody( + user_id=identity.user_id, + user_email=f"{identity.username}@example.com", + user_role=global_role, + auto_create_key=False, + ) + ) + self.resources.defer(lambda: self.bootstrap.delete_user_strict(identity.user_id)) + for tenant in tenants: + unwrap( + self.bootstrap.proxy.transport.post( + "/organization/member_add", + headers=self.bootstrap.proxy.management_headers(), + json=OrgMemberAddBody( + organization_id=tenant.organization_id, + member=OrgMemberEntry( + user_id=identity.user_id, + role="org_admin" if role == "organization_admin" else "internal_user", + ), + ), + response_type=NoBody, + ) + ) + unwrap( + self.bootstrap.proxy.transport.post( + "/team/member_add", + headers=self.bootstrap.proxy.management_headers(), + json=TeamMemberAddBody( + team_id=tenant.team_id, + member=TeamMemberEntry( + user_id=identity.user_id, + role="admin" if role == "team_admin" else "user", + ), + ), + response_type=NoBody, + ) + ) + return Actor(identity=identity, role=role, global_role=global_role, profile=profile, tenants=tenants) diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index e17b92a13ed..8470d318db8 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -7,7 +7,8 @@ llm-only key hitting a management route). from __future__ import annotations import time -from dataclasses import dataclass +import warnings +from dataclasses import dataclass, field, replace import jwt from e2e_config import MASTER_KEY @@ -20,6 +21,7 @@ from e2e_http import ( StreamingResponse, Success, UnknownApiError, + retry_attempts, unwrap, ) from models import ( @@ -81,7 +83,7 @@ from models import ( UserNewResponse, UserUpdateBody, ) -from proxy_client import ProxyClient +from proxy_client import Caller, ProxyClient MODEL_ACCESS_DENIED_MARKER = "key_model_access_denied" ROUTE_NOT_ALLOWED_MARKER = "not allowed to call this route" @@ -98,7 +100,7 @@ class DashboardSession: its bearer on every subsequent call, the claims it renders the signed-in user from, and where it lands the browser.""" - session_key: str + session_key: str = field(repr=False) claims: UiSessionClaims redirect_url: str @@ -106,7 +108,10 @@ class DashboardSession: @dataclass(frozen=True, slots=True) class ManagementClient: proxy: ProxyClient - master_key: str + master_key: str = field(repr=False) + + def with_caller(self, caller: Caller) -> ManagementClient: + return replace(self, proxy=self.proxy.with_caller(caller)) def llm_only_key(self) -> str: return self.proxy.generate_key(KeyGenerateBody(models=[], allowed_routes=["llm_api_routes"])) @@ -117,7 +122,7 @@ class ManagementClient: dashboard creates it under the session key their sign-in minted). Returns the outcome rather than unwrapping it, so a caller can poll a route that is only transiently refusing.""" - headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key) + headers = self.proxy.management_headers(caller_key) return self.proxy.transport.post( "/key/generate", headers=headers, @@ -131,9 +136,9 @@ class ManagementClient: sign-in minted, never the master key). Returns the outcome rather than unwrapping it, so a caller can poll a route that is only transiently refusing; `update_key_models` is the unwrapping shorthand.""" - headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key) + headers = self.proxy.management_headers(caller_key) last: Result[NoBody] = NetworkError(message="/key/update was never attempted") - for attempt in range(_KEY_WRITE_ATTEMPTS): + for attempt in range(retry_attempts(_KEY_WRITE_ATTEMPTS)): last = self.proxy.transport.post( "/key/update", headers=headers, @@ -144,6 +149,7 @@ class ManagementClient: case UnknownApiError(body=error_body) if any( marker in error_body.lower() for marker in _TRANSIENT_BACKEND_MARKERS ): + warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2) time.sleep(0.5 * (attempt + 1)) continue case _: @@ -153,25 +159,26 @@ class ManagementClient: def update_key_models(self, key: str, models: list[str]) -> None: _ = unwrap(self.update_key(KeyUpdateBody(key=key, models=models))) - def key_info_as(self, key: str, *, caller_key: str) -> Result[KeyInfoResponse]: + def key_info_as(self, key: str, *, caller_key: str | None = None) -> Result[KeyInfoResponse]: return self.proxy.transport.get( "/key/info", - headers=self.proxy.transport.bearer(caller_key), + headers=self.proxy.management_headers(caller_key), params=KeyInfoParams(key=key), response_type=KeyInfoResponse, ) - def delete_key_strict(self, key: str, *, caller_key: str | None = None) -> None: + def delete_key_strict(self, key: str, *, caller_key: str | None = None, missing_ok: bool = False) -> None: """Strict delete for the act phase of a test: a failed delete is a hard failure, unlike the warn-only ProxyClient.delete_key used at teardown.""" - _ = unwrap( - self.proxy.transport.post( - "/key/delete", - headers=self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key), - json=KeyDeleteBody(keys=[key]), - response_type=NoBody, - ) + result = self.proxy.transport.post( + "/key/delete", + headers=self.proxy.management_headers(caller_key), + json=KeyDeleteBody(keys=[key]), + response_type=NoBody, ) + if missing_ok and isinstance(result, UnknownApiError) and result.status_code == 404: + return + _ = unwrap(result) def delete_model_strict(self, model_id: str) -> None: """Strict delete for the act phase of a test: a failed delete is a hard @@ -179,7 +186,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/model/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=ModelDeleteBody(id=model_id), response_type=NoBody, ) @@ -190,7 +197,7 @@ class ManagementClient: Connection button, probing the live provider with the supplied params.""" return self.proxy.transport.post( "/health/test_connection", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=ConnectionTestResponse, timeout=120.0, @@ -200,7 +207,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/key/block", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=KeyBlockBody(key=key), response_type=NoBody, ) @@ -209,7 +216,7 @@ class ManagementClient: return unwrap( self.proxy.transport.post( "/key/regenerate", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=KeyRegenerateBody(key=key, grace_period=grace_period), response_type=KeyGenerateResponse, ) @@ -219,7 +226,7 @@ class ManagementClient: return unwrap( self.proxy.transport.post( f"/key/{key}/reset_spend", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=KeyResetSpendBody(reset_to=reset_to), response_type=KeyResetSpendResponse, ) @@ -228,7 +235,7 @@ class ManagementClient: def key_list(self, key_alias: str, *, caller_key: str | None = None) -> Result[KeyListResponse]: """GET /key/list, the Virtual Keys page's own inventory call. `caller_key` is who is asking: the master key by default, or a virtual key.""" - headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key) + headers = self.proxy.management_headers(caller_key) return self.proxy.transport.get( "/key/list", headers=headers, @@ -266,7 +273,7 @@ class ManagementClient: team_id = unwrap( self.proxy.transport.post( "/team/new", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=TeamNewResponse, ) @@ -276,10 +283,10 @@ class ManagementClient: def update_team(self, body: TeamUpdateBody) -> None: last: Result[NoBody] | None = None - for attempt in range(5): + for attempt in range(retry_attempts(5)): last = self.proxy.transport.post( "/team/update", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=NoBody, ) @@ -289,6 +296,7 @@ class ManagementClient: case UnknownApiError(body=body_text) if ( "connecting to redis" in body_text.lower() or "name resolution" in body_text.lower() ): + warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2) time.sleep(0.5 * (attempt + 1)) continue case _: @@ -299,7 +307,7 @@ class ManagementClient: def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=TeamDeleteBody(team_ids=[team_id]), response_type=NoBody, ) @@ -308,7 +316,7 @@ class ManagementClient: return unwrap( self.proxy.transport.get( "/team/info", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=TeamInfoParams(team_id=team_id), response_type=TeamInfoResponse, ) @@ -320,7 +328,7 @@ class ManagementClient: for entry in unwrap( self.proxy.transport.get( "/team/list", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=NoBody(), response_type=TeamListResponse, ) @@ -328,14 +336,16 @@ class ManagementClient: ) def team_info_status(self, team_id: str) -> ProbeResult: - return self.proxy.transport.probe("/team/info", params=TeamInfoParams(team_id=team_id)) + return self.proxy.transport.probe( + "/team/info", params=TeamInfoParams(team_id=team_id), headers=self.proxy.management_headers() + ) def _wait_for_team(self, team_id: str) -> None: last: Result[TeamInfoResponse] | None = None - for _ in range(_TEAM_READY_ATTEMPTS): + for _ in range(retry_attempts(_TEAM_READY_ATTEMPTS)): last = self.proxy.transport.get( "/team/info", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=TeamInfoParams(team_id=team_id), response_type=TeamInfoResponse, ) @@ -343,25 +353,29 @@ class ManagementClient: case Success(): return case _: + warnings.warn("Repeating team read while the team becomes available", RuntimeWarning, stacklevel=2) time.sleep(_TEAM_READY_SLEEP_SECONDS) assert last is not None raise AssertionError(last) def add_team_member(self, team_id: str, user_id: str) -> None: last: Result[NoBody] | None = None - for attempt in range(_TEAM_READY_ATTEMPTS): + for attempt in range(retry_attempts(_TEAM_READY_ATTEMPTS)): last = self.proxy.transport.post( "/team/member_add", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=TeamMemberAddBody(team_id=team_id, member=TeamMemberEntry(role="user", user_id=user_id)), response_type=NoBody, ) match last: case Success(): return - case UnknownApiError(body=body) if ( - "doesn't exist" in body and attempt + 1 < _TEAM_READY_ATTEMPTS + case UnknownApiError(body=body) if "doesn't exist" in body and attempt + 1 < retry_attempts( + _TEAM_READY_ATTEMPTS ): + warnings.warn( + "Retrying team membership while the team becomes available", RuntimeWarning, stacklevel=2 + ) time.sleep(_TEAM_READY_SLEEP_SECONDS) continue case _: @@ -373,7 +387,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/team/member_delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=TeamMemberDeleteBody(team_id=team_id, user_id=user_id), response_type=NoBody, ) @@ -383,7 +397,7 @@ class ManagementClient: return unwrap( self.proxy.transport.post( "/user/new", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=UserNewResponse, ) @@ -393,7 +407,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/customer/new", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=CustomerNewBody(user_id=user_id), response_type=CustomerResponse, ) @@ -404,7 +418,7 @@ class ManagementClient: return unwrap( self.proxy.transport.get( "/customer/info", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=CustomerInfoParams(end_user_id=end_user_id), response_type=CustomerResponse, ) @@ -413,7 +427,7 @@ class ManagementClient: def delete_customer(self, user_id: str) -> None: _ = self.proxy.transport.post( "/customer/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=CustomerDeleteBody(user_ids=[user_id]), response_type=NoBody, ) @@ -422,7 +436,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/user/update", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=NoBody, ) @@ -431,7 +445,7 @@ class ManagementClient: def delete_user(self, user_id: str) -> None: _ = self.proxy.transport.post( "/user/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=UserDeleteBody(user_ids=[user_id]), response_type=NoBody, ) @@ -442,17 +456,17 @@ class ManagementClient: _ = unwrap( self.proxy.transport.post( "/user/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=UserDeleteBody(user_ids=[user_id]), response_type=UserDeleteResponse, ) ) - def user_info(self, user_id: str) -> UserInfoResponse: + def user_info(self, user_id: str | None = None) -> UserInfoResponse: return unwrap( self.proxy.transport.get( "/user/info", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=UserInfoParams(user_id=user_id), response_type=UserInfoResponse, ) @@ -462,7 +476,7 @@ class ManagementClient: return unwrap( self.proxy.transport.get( "/user/list", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=UserListParams(user_ids=user_id), response_type=UserListResponse, ) @@ -472,7 +486,7 @@ class ManagementClient: listing = unwrap( self.proxy.transport.get( "/user/list", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=UserListParams(user_ids=user_id), response_type=UserListResponse, ) @@ -483,7 +497,7 @@ class ManagementClient: return unwrap( self.proxy.transport.post( "/organization/new", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=OrgNewResponse, ) @@ -493,7 +507,7 @@ class ManagementClient: _ = unwrap( self.proxy.transport.patch( "/organization/update", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=NoBody, ) @@ -502,7 +516,7 @@ class ManagementClient: def delete_org(self, organization_id: str) -> None: _ = self.proxy.transport.delete( "/organization/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=OrgDeleteBody(organization_ids=[organization_id]), response_type=NoBody, ) @@ -511,19 +525,24 @@ class ManagementClient: return unwrap( self.proxy.transport.get( "/organization/info", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=OrgInfoParams(organization_id=organization_id), response_type=OrgInfoResponse, ) ) def org_info_status(self, organization_id: str) -> ProbeResult: - return self.proxy.transport.probe("/organization/info", params=OrgInfoParams(organization_id=organization_id)) + return self.proxy.transport.probe( + "/organization/info", + params=OrgInfoParams(organization_id=organization_id), + headers=self.proxy.management_headers(), + ) + def create_tag(self, body: TagNewBody) -> None: _ = unwrap( self.proxy.transport.post( "/tag/new", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=NoBody, ) @@ -532,7 +551,7 @@ class ManagementClient: def delete_tag(self, name: str) -> None: _ = self.proxy.transport.post( "/tag/delete", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=TagDeleteBody(name=name), response_type=NoBody, ) @@ -542,7 +561,7 @@ class ManagementClient: unwrap( self.proxy.transport.get( "/tag/list", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), params=NoBody(), response_type=TagListResponse, ) @@ -553,7 +572,7 @@ class ManagementClient: return unwrap( self.proxy.transport.post( "/v1/mcp/server", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=McpServerRow, ) @@ -565,7 +584,7 @@ class ManagementClient: return unwrap( self.proxy.transport.put( "/v1/mcp/server", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=body, response_type=McpServerRow, ) @@ -576,7 +595,7 @@ class ManagementClient: unwrap it while a deferred teardown can ignore an already-deleted server.""" return self.proxy.transport.delete( f"/v1/mcp/server/{server_id}", - headers=self.proxy.transport.master, + headers=self.proxy.management_headers(), json=NoBody(), response_type=NoBody, ) diff --git a/tests/e2e/management/test_jwt_management_e2e.py b/tests/e2e/management/test_jwt_management_e2e.py index 22306da8eb8..5898073a4e6 100644 --- a/tests/e2e/management/test_jwt_management_e2e.py +++ b/tests/e2e/management/test_jwt_management_e2e.py @@ -2,60 +2,247 @@ from __future__ import annotations -from typing import Final +from typing import Final, Literal import pytest -from e2e_config import CHEAP_OPENAI_MODEL, unique_marker +from e2e_config import CHEAP_OPENAI_MODEL, PROXY_BASE_URL, unique_marker from e2e_http import UnauthorizedError, UnknownApiError, unwrap -from idp import ADMIN_CLIENT_ID, Identity, Keycloak +from idp import ADMIN_CLIENT_ID, Identity, Keycloak, token_claims from lifecycle import ResourceManager +from management.jwt_actors import ActorFactory, ActorRole from management_client import ManagementClient -from models import KeyGenerateBody, KeyUpdateBody, TeamNewBody, UserNewBody +from models import KeyGenerateBody, KeyUpdateBody, TeamNewBody, UserInfoParams, UserInfoResponse, UserNewBody +from proxy_client import Caller pytestmark = pytest.mark.e2e class TestJwtManagement: - @pytest.mark.covers("mgmt.key.jwt.lifecycle") - def test_admin_creates_reads_updates_clears_and_deletes_a_key( - self, client: ManagementClient, idp: Keycloak, jwt_identity: Identity, resources: ResourceManager - ) -> None: - admin: Final = idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID) - alias: Final = f"e2e-jwt-key-{unique_marker()}" - created: Final = unwrap( - client.generate_key( - KeyGenerateBody(key_alias=alias, team_id=jwt_identity.group, models=[CHEAP_OPENAI_MODEL]), - caller_key=admin, + @pytest.mark.parametrize( + "role", + ( + "proxy_admin", + "proxy_admin_viewer", + "organization_admin", + "team_admin", + "team_member", + "internal_user", + "internal_user_viewer", + "unrelated_user", + ), + ) + @pytest.mark.covers("mgmt.user.jwt.database_roles") + def test_actor_subject_and_database_role(self, actor_factory: ActorFactory, role: ActorRole) -> None: + tenants: Final = ( + (actor_factory.tenant(),) if role in ("organization_admin", "team_admin", "team_member") else () + ) + actor: Final = actor_factory.create(role, tenants=tenants) + caller: Final = actor.mint_caller(actor_factory.idp) + claims: Final = token_claims(caller.credential) + assert claims.sub == actor.identity.user_id + assert claims.iss == actor_factory.idp.issuer + assert claims.aud == "litellm-e2e" or "litellm-e2e" in claims.aud + assert actor.identity.groups == () + assert ("litellm_proxy_admin" in claims.scope.split()) == (role == "proxy_admin") + stored: Final = actor_factory.bootstrap.user_info(actor.identity.user_id) + assert stored.user_id == actor.identity.user_id + assert stored.user_info.user_role == actor.global_role + bound: Final = actor_factory.bootstrap.with_caller(caller) + own: Final = unwrap( + bound.proxy.transport.get( + "/user/info", + headers=bound.proxy.management_headers(), + params=UserInfoParams(), + response_type=UserInfoResponse, ) ) - resources.defer(lambda: client.proxy.delete_key(created.key)) + assert own.user_id == actor.identity.user_id + assert own.user_info.user_role == actor.global_role + for tenant in tenants: + info = actor_factory.bootstrap.team_info(tenant.team_id) + assert info.organization_id == tenant.organization_id + assert {(member.user_id, member.role) for member in info.members_with_roles} == { + (actor.identity.user_id, "admin" if role == "team_admin" else "user") + } + assert { + (member.user_id, member.user_role) + for member in actor_factory.bootstrap.org_info(tenant.organization_id).members + } == {(actor.identity.user_id, "org_admin" if role == "organization_admin" else "internal_user")} - original: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info - assert original.key_alias == alias and original.team_id == jwt_identity.group + @pytest.mark.covers("mgmt.key.jwt.viewer_denied") + def test_admin_viewer_reads_but_cannot_update(self, actor_factory: ActorFactory) -> None: + actor: Final = actor_factory.create("proxy_admin_viewer") + viewer: Final = actor_factory.bootstrap.with_caller(actor.mint_caller(actor_factory.idp)) + alias: Final = f"e2e-viewer-{unique_marker()}" + key: Final = actor_factory.key().key + unwrap(actor_factory.bootstrap.update_key(KeyUpdateBody(key=key, key_alias=alias))) + assert viewer.proxy.key_info(key).key_alias == alias + denied: Final = viewer.update_key(KeyUpdateBody(key=key, key_alias="forbidden")) + assert isinstance(denied, UnknownApiError) and denied.status_code == 403, f"viewer write was accepted: {denied}" + assert "proxy_admin_viewer" in denied.body and "/key/update" in denied.body + assert actor_factory.bootstrap.proxy.key_info(key).key_alias == alias + + @pytest.mark.covers("mgmt.user.oidc.identity_mapping") + def test_oidc_browser_profile_identity_mapping(self, actor_factory: ActorFactory) -> None: + actor: Final = actor_factory.create("internal_user") + idp: Final = actor_factory.idp.with_strict_cleanup() + discovery: Final = idp.discovery() + assert discovery.issuer == idp.issuer + assert discovery.jwks_uri == idp.jwks_url + callback: Final = f"{PROXY_BASE_URL}/sso/callback" + browser: Final = idp.browser_client(callback_url=callback, defer=actor_factory.resources.defer) + token: Final = idp.browser_token(actor.identity, browser) + assert token_claims(token).sub == actor.identity.user_id + userinfo: Final = idp.userinfo(token) + assert userinfo.sub == actor.identity.user_id + assert userinfo.email == f"{actor.identity.username}@example.com" + assert browser.environment(discovery)["GENERIC_USER_ID_ATTRIBUTE"] == "sub" + + @pytest.mark.covers("mgmt.key.jwt.lifecycle") + @pytest.mark.parametrize("credential_kind", ("direct_jwt", "virtual_key")) + def test_admin_creates_reads_updates_clears_and_deletes_a_key( + self, + actor_factory: ActorFactory, + credential_kind: Literal["direct_jwt", "virtual_key"], + ) -> None: + tenant: Final = actor_factory.tenant() + actor: Final = actor_factory.create("proxy_admin", tenants=(tenant,), profile="group_scoped") + virtual_key: Final = ( + actor_factory.key(user_id=actor.identity.user_id).key if credential_kind == "virtual_key" else None + ) + admin: Final = virtual_key if virtual_key is not None else actor.mint_caller(actor_factory.idp).credential + bound: Final = actor_factory.bootstrap.with_caller( + Caller(credential=admin, kind=credential_kind, role="proxy_admin") + ) + assert bound.user_info().user_id == actor.identity.user_id + alias: Final = f"e2e-jwt-key-{unique_marker()}" + created: Final = unwrap( + bound.generate_key( + KeyGenerateBody(key_alias=alias, team_id=tenant.team_id, models=[CHEAP_OPENAI_MODEL]), + ) + ) + actor_factory.resources.defer(lambda: actor_factory.bootstrap.delete_key_strict(created.key, missing_ok=True)) + + original: Final = unwrap(bound.key_info_as(created.key)).info + assert original.key_alias == alias and original.team_id == tenant.team_id assert original.models == [CHEAP_OPENAI_MODEL] updated_alias: Final = f"{alias}-updated" - unwrap( - client.update_key(KeyUpdateBody(key=created.key, key_alias=updated_alias, rpm_limit=120), caller_key=admin) - ) - updated: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info + unwrap(bound.update_key(KeyUpdateBody(key=created.key, key_alias=updated_alias, rpm_limit=120))) + updated: Final = unwrap(bound.key_info_as(created.key)).info assert updated.key_alias == updated_alias and updated.rpm_limit == 120 assert updated.models == [CHEAP_OPENAI_MODEL], "omitted models must preserve the restriction" - unwrap(client.update_key(KeyUpdateBody(key=created.key, models=[]), caller_key=admin)) - cleared: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info + unwrap(bound.update_key(KeyUpdateBody(key=created.key, models=[]))) + cleared: Final = unwrap(bound.key_info_as(created.key)).info assert cleared.models == [] and cleared.rpm_limit == 120 - assert unwrap(client.key_list(updated_alias, caller_key=admin)).total_count == 1 - client.delete_key_strict(created.key, caller_key=admin) - assert unwrap(client.key_list(updated_alias, caller_key=admin)).total_count == 0 + assert unwrap(bound.key_list(updated_alias)).total_count == 1 + bound.delete_key_strict(created.key) + assert unwrap(bound.key_list(updated_alias)).total_count == 0 + + @pytest.mark.covers("mgmt.team.jwt.tenant_isolation") + def test_two_actor_sets_keep_tenants_and_keys_isolated(self, actor_factory: ActorFactory) -> None: + first: Final = actor_factory.tenant() + second: Final = actor_factory.tenant() + assert first.organization_id != second.organization_id and first.team_id != second.team_id + actors: Final = tuple( + actor_factory.create("team_member", tenants=(tenant,), profile="group_scoped") for tenant in (first, second) + ) + assert actors[0].identity.user_id != actors[1].identity.user_id + callers: Final = tuple( + actor_factory.bootstrap.with_caller(actor.mint_caller(actor_factory.idp)) for actor in actors + ) + keys: Final = tuple(actor_factory.key(tenant) for tenant in (first, second)) + assert keys[0].key != keys[1].key + assert callers[0].proxy.key_info(keys[0].key).team_id == first.team_id + assert callers[1].proxy.key_info(keys[1].key).team_id == second.team_id + for caller, other_key in ((callers[0], keys[1].key), (callers[1], keys[0].key)): + hidden = caller.key_info_as(other_key) + assert isinstance(hidden, UnknownApiError) and hidden.status_code == 403 + assert tuple(actor.identity.groups for actor in actors) == ((first.team_id,), (second.team_id,)) + + @pytest.mark.covers("mgmt.team.jwt.multiple_memberships") + def test_multi_group_actor_keeps_exact_memberships(self, actor_factory: ActorFactory) -> None: + tenants: Final = (actor_factory.tenant(), actor_factory.tenant()) + actor: Final = actor_factory.create("team_member", tenants=tenants, profile="group_scoped") + claims: Final = token_claims(actor.mint_caller(actor_factory.idp).credential) + assert set(claims.groups) == {tenant.team_id for tenant in tenants} + assert "litellm_proxy_admin" not in claims.scope.split() + assert actor.identity.groups == tuple(tenant.team_id for tenant in tenants) + for tenant in tenants: + assert { + (entry.user_id, entry.role) + for entry in actor_factory.bootstrap.team_info(tenant.team_id).members_with_roles + } == {(actor.identity.user_id, "user")} + + @pytest.mark.covers("mgmt.user.jwt.cleanup") + def test_successful_actor_cleanup_removes_owned_state(self, actor_factory: ActorFactory) -> None: + resources: Final = ResourceManager(client=actor_factory.bootstrap.proxy, strict_cleanup=True) + factory: Final = ActorFactory(bootstrap=actor_factory.bootstrap, idp=actor_factory.idp, resources=resources) + try: + tenant: Final = factory.tenant() + actor: Final = factory.create("team_member", tenants=(tenant,), profile="group_scoped") + key: Final = factory.key(tenant) + alias: Final = factory.bootstrap.proxy.key_info(key.key).key_alias + assert alias is not None + finally: + resources.teardown() + assert factory.bootstrap.user_count(actor.identity.user_id) == 0 + assert factory.bootstrap.key_alias_count(alias) == 0 + assert factory.bootstrap.team_info_status(tenant.team_id).status_code == 404 + assert factory.bootstrap.org_info_status(tenant.organization_id).status_code == 404 + factory.idp.assert_absent("users", actor.identity.user_id) + factory.idp.assert_absent("groups", tenant.group_id) + + @pytest.mark.parametrize("stage", ("group", "user")) + @pytest.mark.covers("mgmt.user.jwt.partial_cleanup") + def test_partial_setup_removes_previously_created_identities( + self, + actor_factory: ActorFactory, + stage: Literal["group", "user"], + ) -> None: + idp: Final = actor_factory.idp.with_strict_cleanup() + resources: Final = ResourceManager(client=actor_factory.bootstrap.proxy, strict_cleanup=True) + marker: Final = unique_marker() + group_id: Final = idp.create_group(f"e2e-partial-{marker}") + resources.defer(lambda: idp.delete_group(group_id)) + try: + identity: Final = ( + idp.provision_user( + marker=marker, + groups=(f"e2e-partial-{marker}",), + group_ids=(group_id,), + defer=resources.defer, + ) + if stage == "user" + else None + ) + if identity is None: + with pytest.raises(pytest.fail.Exception, match="HTTP 409"): + idp.create_group(f"e2e-partial-{marker}") + else: + with pytest.raises(pytest.fail.Exception, match="HTTP 409"): + idp.create_user( + username=identity.username, + email=f"{identity.username}@example.com", + password=identity.password, + groups=identity.groups, + ) + finally: + resources.teardown() + idp.assert_absent("groups", group_id) + if identity is not None: + idp.assert_absent("users", identity.user_id) @pytest.mark.covers("mgmt.key.jwt.member_denied", "mgmt.key.jwt.other_team_denied") def test_member_cannot_write_and_another_team_cannot_read_the_key( self, client: ManagementClient, idp: Keycloak, jwt_identity: Identity, resources: ResourceManager ) -> None: admin: Final = idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID) + bound: Final = client.with_caller(Caller(credential=admin, kind="direct_jwt", role="proxy_admin")) member: Final = idp.access_token(jwt_identity) + member_client: Final = client.with_caller(Caller(credential=member, kind="direct_jwt", role="team_member")) alias: Final = f"e2e-jwt-owned-{unique_marker()}" created: Final = unwrap( client.generate_key(KeyGenerateBody(key_alias=alias, team_id=jwt_identity.group), caller_key=admin) @@ -63,14 +250,14 @@ class TestJwtManagement: resources.defer(lambda: client.proxy.delete_key(created.key)) client.add_team_member(jwt_identity.group, jwt_identity.user_id) - assert unwrap(client.key_info_as(created.key, caller_key=member)).info.key_alias == alias + assert unwrap(member_client.key_info_as(created.key)).info.key_alias == alias - refused: Final = client.update_key(KeyUpdateBody(key=created.key, key_alias="forbidden"), caller_key=member) + refused: Final = member_client.update_key(KeyUpdateBody(key=created.key, key_alias="forbidden")) assert isinstance(refused, UnauthorizedError), f"member write was accepted: {refused}" assert "does not have permissions for endpoint" in refused.body.lower(), ( f"expected a permission denial: {refused}" ) - assert unwrap(client.key_info_as(created.key, caller_key=admin)).info.key_alias == alias + assert unwrap(bound.key_info_as(created.key)).info.key_alias == alias marker: Final = unique_marker() outsider: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer) @@ -88,4 +275,4 @@ class TestJwtManagement: assert isinstance(hidden, UnknownApiError) and hidden.status_code == 403, ( f"another team must not read this key: {hidden}" ) - assert unwrap(client.key_info_as(created.key, caller_key=admin)).info.team_id == jwt_identity.group + assert unwrap(bound.key_info_as(created.key)).info.team_id == jwt_identity.group diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 3cab0334dea..ba6d0fe3334 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -283,10 +283,15 @@ class ChatToolResultTurn(BaseModel): type ChatTurn = ChatMessage | ChatAssistantTurn | ChatToolResultTurn +class ChatStreamOptions(BaseModel): + include_usage: bool + + class ChatBody(BaseModel): model: str messages: Sequence[ChatTurn] stream: bool = False + stream_options: ChatStreamOptions | None = None max_tokens: int | None = None max_completion_tokens: int | None = None temperature: float | None = None @@ -488,12 +493,18 @@ class AnthropicToolResultTurn(BaseModel): type AnthropicMessage = ChatMessage | AnthropicAssistantTurn | AnthropicToolResultTurn +class AnthropicToolChoice(BaseModel): + type: Literal["auto", "any", "tool", "none"] + name: str | None = None + + class AnthropicMessagesBody(BaseModel): model: str messages: list[AnthropicMessage] max_tokens: int stream: bool | None = None tools: list[AnthropicTool] | None = None + tool_choice: AnthropicToolChoice | None = None guardrails: list[str] | None = None cache: dict[str, bool] | None = {"no-cache": True} @@ -1091,13 +1102,13 @@ class UiLoginBody(BaseModel): class UiLoginResponse(BaseModel): - token: str + token: str = Field(repr=False) redirect_url: str class UiSessionClaims(BaseModel): user_id: str - key: str + key: str = Field(repr=False) user_role: str login_method: Literal["sso", "username_password"] exp: int @@ -1135,6 +1146,7 @@ class TeamInfoParams(BaseModel): class TeamData(BaseModel): + organization_id: str | None = None team_alias: str | None = None models: list[str] = [] members_with_roles: list[TeamMemberEntry] = [] @@ -1175,6 +1187,7 @@ class UserNewBody(BaseModel): user_email: str user_role: UserRole user_id: str | None = None + auto_create_key: bool | None = None class UserNewResponse(BaseModel): @@ -1187,7 +1200,7 @@ class UserUpdateBody(BaseModel): class UserInfoParams(BaseModel): - user_id: str + user_id: str | None = None class UserData(BaseModel): @@ -1240,16 +1253,36 @@ class OrgInfoParams(BaseModel): organization_id: str +class OrgMembership(BaseModel): + user_id: str + user_role: str + + class OrgInfoResponse(BaseModel): organization_id: str organization_alias: str | None = None models: list[str] = [] + members: tuple[OrgMembership, ...] = () + + +class OrgMemberEntry(BaseModel): + user_id: str + role: Literal["org_admin", "internal_user"] + + +class OrgMemberAddBody(BaseModel): + organization_id: str + member: OrgMemberEntry class OrgDeleteBody(BaseModel): organization_ids: list[str] +class OrgDeleteResponse(RootModel[tuple[OrgInfoResponse, ...]]): + pass + + # ---------- tags (management) ---------- diff --git a/tests/e2e/provider_edge.py b/tests/e2e/provider_edge.py index 6c87c7ef7ac..de36895ebb6 100644 --- a/tests/e2e/provider_edge.py +++ b/tests/e2e/provider_edge.py @@ -92,6 +92,7 @@ from fixture_mode import ( current_test_key, parse_fixture_mode, ) +from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity from pydantic import JsonValue, TypeAdapter EDGE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType( @@ -404,6 +405,12 @@ def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle: f"under {slug}; re-record with E2E_FIXTURE_MODE=record" ) closest, closest_file = _closest_recorded(canonical, recorded) + if bundle.manifest.match_profile == "stateless_v1": + expected: Final = _JSON.validate_json(closest.content) + actual: Final = _JSON.validate_json(canonical.content) + assert isinstance(expected, dict) and isinstance(actual, dict) + changed: Final = ", ".join(key for key in expected if expected[key] != actual.get(key)) + return f"stateless_v1 replay mismatch: {changed or 'method/path'}; re-record with E2E_FIXTURE_MODE=record" diff: Final = "\n".join( islice( difflib.unified_diff( @@ -785,11 +792,33 @@ def handle_edge_request( mount, _, upstream_path = split.path.lstrip("/").partition("/") upstream_base: Final = mounts.get(mount) if upstream_base is None: - return _text_reply( - 404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}" + return _text_reply(404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}") + profile: Final = ( + backend.recorder.profile + if isinstance(backend, RecordEdge) + else backend.source.bundle.manifest.match_profile + if isinstance(backend, ReplayEdge) + else "legacy" + ) + identity: Final = ( + strict_identity( + method=method, + path=split.path, + query=split.query, + headers=headers, + body=body, + mount=mount, + upstream_base=upstream_base, ) - request: Final = edge_request( - method, split.path, split.query, body, _header_value(headers, "content-type") + if profile == "stateless_v1" + else None + ) + if isinstance(identity, IneligibleRequest): + return _text_reply(REPLAY_MISS_STATUS, f"stateless_v1 eligibility error: {identity.reason}") + request: Final = ( + RecordedRequest(method=method.lower(), path=split.path, headers={}, strict_identity=identity) + if identity is not None + else edge_request(method, split.path, split.query, body, _header_value(headers, "content-type")) ) match backend: case LiveEdge(): @@ -837,6 +866,14 @@ class _EdgeHandler(BaseHTTPRequestHandler): body: Final = self.rfile.read(length) if length else None if edge_server.observation is not None: edge_server.observation.observe(body) + strict: Final = ( + isinstance(edge_server.backend, RecordEdge) and edge_server.backend.recorder.profile == "stateless_v1" + or isinstance(edge_server.backend, ReplayEdge) + and edge_server.backend.source.bundle.manifest.match_profile == "stateless_v1" + ) + if strict and len({name.lower() for name in self.headers.keys()}) != len(self.headers): + self._write_reply(_text_reply(REPLAY_MISS_STATUS, "stateless_v1 eligibility error: duplicate headers")) + return outcome: Final = handle_edge_request( edge_server.backend, edge_server.mounts, @@ -955,16 +992,16 @@ def start_provider_edge( @functools.lru_cache(maxsize=8) -def _shared_recorder(root: Path) -> BundleRecorder: - prepared = prepare_bundle(root) +def _shared_recorder(root: Path, profile: MatchProfile = "legacy") -> BundleRecorder: + prepared = prepare_bundle(root, profile=profile) if isinstance(prepared, UnsafeBundleDir): raise ValueError(f"E2E_FIXTURE_DIR {prepared.path} {prepared.reason}") return prepared @functools.lru_cache(maxsize=8) -def _shared_replay_source(root: Path) -> ReplaySource: - loaded = load_bundle(root) +def _shared_replay_source(root: Path, profile: MatchProfile = "legacy") -> ReplaySource: + loaded = load_bundle(root, profile=profile) if isinstance(loaded, UnreadableBundle): raise ValueError(f"cannot replay from {root}: {loaded.reason}") return ReplaySource(bundle=loaded) @@ -977,11 +1014,12 @@ def _shared_edge( bind_host: str, advertise_host: str, forward_timeout: float, + profile: MatchProfile, ) -> ProviderEdge: backend: Final[EdgeBackend] = ( - RecordEdge(recorder=_shared_recorder(bundle_dir), lock=threading.Lock()) + RecordEdge(recorder=_shared_recorder(bundle_dir, profile), lock=threading.Lock()) if mode == "record" - else ReplayEdge(source=_shared_replay_source(bundle_dir)) + else ReplayEdge(source=_shared_replay_source(bundle_dir, profile)) ) return start_provider_edge( backend, @@ -998,7 +1036,7 @@ def replay_leftover_error(*, mode_raw: str, bundle_dir: Path, test_key: str) -> recording it no longer matches. Inert in every other mode.""" if parse_fixture_mode(mode_raw) != "replay": return None - return _shared_replay_source(bundle_dir).leftover_error(test_key) + return _shared_replay_source(bundle_dir, match_profile()).leftover_error(test_key) def provider_edge_api_base( @@ -1021,10 +1059,10 @@ def provider_edge_api_base( return None case "record" | "replay": if mount not in EDGE_MOUNTS: - raise ValueError( - f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}" - ) - return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout).api_base(mount) + raise ValueError(f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}") + return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout, match_profile()).api_base( + mount + ) case _: assert_never(mode) @@ -1037,9 +1075,9 @@ def _observed_backend(mode_raw: str, bundle_dir: Path) -> EdgeBackend: case "live": return LiveEdge() case "record": - return RecordEdge(_shared_recorder(bundle_dir), threading.Lock()) + return RecordEdge(_shared_recorder(bundle_dir, match_profile()), threading.Lock()) case "replay": - return ReplayEdge(_shared_replay_source(bundle_dir)) + return ReplayEdge(_shared_replay_source(bundle_dir, match_profile())) case _: assert_never(mode) diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 1fe2ec905ef..48a6110dc0b 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -11,14 +11,23 @@ from __future__ import annotations import time import warnings from collections.abc import Callable, Mapping -from dataclasses import dataclass -from functools import reduce +from dataclasses import dataclass, field, replace from datetime import datetime +from functools import reduce from types import MappingProxyType -from typing import Final - -from pydantic import BaseModel +from typing import Final, Literal +from e2e_config import ( + CONTROL_PLANE_BASE_URL, + MASTER_KEY, + POLL_INTERVAL, + POLL_TIMEOUT, + PROXY_BASE_URL, + PROXY_REPLICA_URLS, + REQUEST_TIMEOUT, + SLOW_PROVIDER_TIMEOUT_SECONDS, + settle_propagation, +) from e2e_http import ( AnthropicHeaders, AuthHeaders, @@ -55,6 +64,7 @@ from models import ( KeyInfoParams, KeyInfoResponse, LiteLLMParamsBody, + MemorySummaryResponse, ModelDeleteBody, ModelInfoBody, ModelInfoEntry, @@ -63,7 +73,6 @@ from models import ( ModelNewBody, ModelNewResponse, ModelsListParams, - MemorySummaryResponse, ModelsListResponse, ModelUpdateBody, OcrBody, @@ -76,23 +85,13 @@ from models import ( TeamDeleteBody, TeamNewBody, TeamNewResponse, - UserDeleteBody, - UserDeleteResponse, ToolsetCreateBody, ToolsetRow, ToolsetUpdateBody, + UserDeleteBody, + UserDeleteResponse, ) -from e2e_config import ( - CONTROL_PLANE_BASE_URL, - MASTER_KEY, - POLL_INTERVAL, - POLL_TIMEOUT, - PROXY_BASE_URL, - PROXY_REPLICA_URLS, - REQUEST_TIMEOUT, - SLOW_PROVIDER_TIMEOUT_SECONDS, - settle_propagation, -) +from pydantic import BaseModel from transport import HttpTransport, SplitTransport, Transport, is_control_plane_path RowsPredicate = Callable[[list[SpendLogRow]], bool] @@ -421,11 +420,23 @@ def converge_timeout_message(*, what: str, replica: str, timeout: float, last_re ) +CredentialKind = Literal["master", "direct_jwt", "virtual_key", "dashboard_session"] + + +@dataclass(frozen=True, slots=True) +class Caller: + credential: str = field(repr=False) + kind: CredentialKind + role: str + tenant: str | None = None + + @dataclass(frozen=True, slots=True) class ProxyClient: transport: Transport replicas: Mapping[str, Transport] control_replicas: Mapping[str, Transport] + caller: Caller | None = None poll_timeout: float = 120.0 poll_interval: float = 5.0 model_servable_timeout: float = MODEL_SERVABLE_TIMEOUT @@ -433,13 +444,24 @@ class ProxyClient: model_servable_interval: float = MODEL_SERVABLE_INTERVAL model_servable_request_timeout: float = MODEL_SERVABLE_REQUEST_TIMEOUT + def with_caller(self, caller: Caller) -> ProxyClient: + return replace(self, caller=caller) + + def management_headers(self, caller_key: str | None = None, *, transport: Transport | None = None) -> AuthHeaders: + selected: Final = self.transport if transport is None else transport + if caller_key is not None: + return selected.bearer(caller_key) + if self.caller is not None: + return selected.bearer(self.caller.credential) + return selected.master + # ---- keys / customers (satisfies lifecycle.ResourceClient) ---------- def generate_key(self, body: KeyGenerateBody) -> str: return unwrap( self.transport.post( "/key/generate", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=KeyGenerateResponse, ) @@ -448,7 +470,7 @@ class ProxyClient: def delete_key(self, key: str) -> None: _ = self.transport.post( "/key/delete", - headers=self.transport.master, + headers=self.management_headers(), json=KeyDeleteBody(keys=[key]), response_type=NoBody, ) @@ -458,7 +480,7 @@ class ProxyClient: return _ = self.transport.post( "/customer/delete", - headers=self.transport.master, + headers=self.management_headers(), json=CustomerDeleteBody(user_ids=user_ids), response_type=NoBody, ) @@ -467,7 +489,7 @@ class ProxyClient: return unwrap( self.transport.get( "/key/info", - headers=self.transport.master, + headers=self.management_headers(), params=KeyInfoParams(key=key), response_type=KeyInfoResponse, ) @@ -477,7 +499,7 @@ class ProxyClient: return { url: transport.get( "/debug/memory/summary", - headers=transport.master, + headers=self.management_headers(transport=transport), params=NoBody(), response_type=MemorySummaryResponse, ) @@ -524,11 +546,12 @@ class ProxyClient: {replica: outcome.result for replica, outcome in outcomes.items() if isinstance(outcome, Converged)} ) - @staticmethod def _body_poller[R: BaseModel]( - transport: Transport, path: str, params: BaseModel, response_type: type[R] + self, transport: Transport, path: str, params: BaseModel, response_type: type[R] ) -> Poller[Result[R]]: - return lambda: transport.get(path, headers=transport.master, params=params, response_type=response_type) + return lambda: transport.get( + path, headers=self.management_headers(transport=transport), params=params, response_type=response_type + ) def model_info(self) -> list[ModelInfoEntry]: """Every configured deployment with the price the proxy resolved for it @@ -536,7 +559,7 @@ class ProxyClient: return unwrap( self.transport.get( "/model/info", - headers=self.transport.master, + headers=self.management_headers(), params=NoBody(), response_type=ModelInfoResponse, ) @@ -546,7 +569,7 @@ class ProxyClient: return unwrap( self.transport.get( "/public/litellm_model_cost_map", - headers=self.transport.master, + headers=self.management_headers(), params=NoBody(), response_type=CostMap, ) @@ -607,7 +630,7 @@ class ProxyClient: model_id = unwrap( self.transport.post( "/model/new", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=ModelNewResponse, ) @@ -623,7 +646,7 @@ class ProxyClient: def _await_model_servable(self, model_name: str, listed_for: str | None = None) -> None: """Block until every replica lists `model_name`, or fail at model_servable_timeout.""" - headers: Final = self.transport.master if listed_for is None else self.transport.bearer(listed_for) + headers: Final = self.management_headers(listed_for) outcome: Final = await_servable_everywhere( {url: self._models_poller(transport, headers) for url, transport in self.replicas.items()}, model_name=model_name, @@ -666,7 +689,7 @@ class ProxyClient: unwrap( self.transport.post( "/model/update", - headers=self.transport.master, + headers=self.management_headers(), json=ModelUpdateBody( litellm_params=litellm_params, model_info=ModelInfoBody(id=model_id), @@ -678,7 +701,7 @@ class ProxyClient: def delete_model(self, model_id: str) -> None: result = self.transport.post( "/model/delete", - headers=self.transport.master, + headers=self.management_headers(), json=ModelDeleteBody(id=model_id), response_type=NoBody, ) @@ -747,11 +770,10 @@ class ProxyClient: f"GET {path} on {replica} still answers {self.poll_timeout}s after the delete; last read: {last}" ) - @staticmethod - def _reader[R: BaseModel](transport: Transport, path: str, response_type: type[R]) -> ReplicaRead[Result[R]]: + def _reader[R: BaseModel](self, transport: Transport, path: str, response_type: type[R]) -> ReplicaRead[Result[R]]: return lambda request_timeout: transport.get( path, - headers=transport.master, + headers=self.management_headers(transport=transport), params=NoBody(), response_type=response_type, timeout=request_timeout, @@ -763,7 +785,7 @@ class ProxyClient: return unwrap( self.transport.post( "/v1/mcp/toolset", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=ToolsetRow, ) @@ -775,7 +797,7 @@ class ProxyClient: return unwrap( self.transport.put( "/v1/mcp/toolset", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=ToolsetRow, ) @@ -786,7 +808,7 @@ class ProxyClient: can unwrap it while a deferred teardown can ignore an already-deleted row.""" return self.transport.delete( f"/v1/mcp/toolset/{toolset_id}", - headers=self.transport.master, + headers=self.management_headers(), json=NoBody(), response_type=NoBody, ) @@ -795,7 +817,7 @@ class ProxyClient: unwrap( self.transport.post( "/credentials", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=CredentialCreateResponse, ) @@ -804,7 +826,7 @@ class ProxyClient: def delete_credential(self, credential_name: str) -> None: result = self.transport.delete( f"/credentials/{credential_name}", - headers=self.transport.master, + headers=self.management_headers(), json=NoBody(), response_type=NoBody, ) @@ -815,7 +837,7 @@ class ProxyClient: return unwrap( self.transport.post( "/team/new", - headers=self.transport.master, + headers=self.management_headers(), json=body, response_type=TeamNewResponse, ) @@ -824,7 +846,7 @@ class ProxyClient: def delete_team(self, team_id: str) -> None: result = self.transport.post( "/team/delete", - headers=self.transport.master, + headers=self.management_headers(), json=TeamDeleteBody(team_ids=[team_id]), response_type=NoBody, ) @@ -836,7 +858,7 @@ class ProxyClient: a user the proxy only upserts after a successful auth.""" result = self.transport.post( "/user/delete", - headers=self.transport.master, + headers=self.management_headers(), json=UserDeleteBody(user_ids=[user_id]), response_type=UserDeleteResponse, ) @@ -909,7 +931,7 @@ class ProxyClient: def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]: result = self.transport.get( "/spend/logs", - headers=self.transport.master, + headers=self.management_headers(), params=params, response_type=SpendLogs, ) @@ -924,7 +946,7 @@ class ProxyClient: return unwrap( self.transport.get( "/spend/logs/v2", - headers=self.transport.master, + headers=self.management_headers(), params=SpendLogsPageParams( start_date=start.strftime("%Y-%m-%d %H:%M:%S"), end_date=end.strftime("%Y-%m-%d %H:%M:%S"), @@ -977,7 +999,7 @@ class ProxyClient: # ---- route probe ---------------------------------------------------- def probe(self, path: str, *, params: NoBody) -> ProbeResult: - return self.transport.probe(path, params=params) + return self.transport.probe(path, params=params, headers=self.management_headers()) def build_proxy_client( diff --git a/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py b/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py new file mode 100644 index 00000000000..26809874aed --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/spend_reconciliation.py @@ -0,0 +1,113 @@ +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from math import isclose +from typing import Final + +from e2e_config import provider_edge_base, unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody +from spend_e2e_client import SpendClient + +INPUT_RATE: Final = 0.00004 +OUTPUT_RATE: Final = 0.00008 + + +@dataclass(frozen=True) +class TeamTraffic: + team_id: str + key: str + responses: tuple[ChatResponse, ...] + + @property + def prompt_tokens(self) -> int: + return sum(response.usage.prompt_tokens or 0 for response in self.responses if response.usage) + + @property + def completion_tokens(self) -> int: + return sum(response.usage.completion_tokens or 0 for response in self.responses if response.usage) + + @property + def spend(self) -> float: + return self.prompt_tokens * INPUT_RATE + self.completion_tokens * OUTPUT_RATE + + +def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[TeamTraffic, ...]: + base: Final = provider_edge_base("openai") + model: Final = f"e2e-reconciliation-{unique_marker()}" + model_id: Final = client.proxy.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-5.6-luna", + api_key="os.environ/OPENAI_API_KEY", + api_base=None if base is None else f"{base}/v1", + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + def team_traffic() -> TeamTraffic: + team: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-spend-{unique_marker()}")) + resources.defer(lambda: client.proxy.delete_team(team)) + key: Final = client.proxy.generate_key(KeyGenerateBody(team_id=team, models=[model])) + resources.defer(lambda: client.proxy.delete_key(key)) + + prompts: Final = tuple(f"Reply with one word. {index} {unique_marker()}" for index in range(7)) + + def call(index: int) -> ChatResponse: + response: Final = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=prompts[index])], + max_completion_tokens=128, + ), + ) + ) + assert response.id, "successful response must have an ID" + assert response.usage is not None, "successful response must have usage" + assert response.usage.prompt_tokens is not None and response.usage.prompt_tokens > 0 + assert response.usage.completion_tokens is not None and response.usage.completion_tokens > 0 + assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens + assert not response.usage.cache_creation_input_tokens + assert not response.usage.cache_read_input_tokens + assert not response.usage.prompt_tokens_details or not response.usage.prompt_tokens_details.cached_tokens + return response + + sequential: Final = call(0) + with ThreadPoolExecutor(max_workers=6) as pool: + concurrent: Final = tuple(pool.map(call, range(1, 7))) + return TeamTraffic(team, key, (sequential, *concurrent)) + + return tuple(team_traffic() for _ in range(2)) + + +def assert_logs_match(client: SpendClient, traffic: TeamTraffic) -> None: + expected_ids: Final = frozenset(response.id for response in traffic.responses) + assert len(expected_ids) == len(traffic.responses), "responses must have distinct IDs" + rows: Final = client.poll_logs_for_key( + traffic.key, + min_rows=len(traffic.responses), + predicate=lambda values: frozenset(row.request_id for row in values) == expected_ids, + ) + assert frozenset(row.request_id for row in rows) == expected_ids, "stored IDs must equal returned response IDs" + assert len(rows) == len(traffic.responses), "expected exactly one scoped spend row per response" + by_id: Final = {row.request_id: row for row in rows} + + def assert_response(response: ChatResponse) -> None: + row: Final = by_id[response.id] + usage: Final = response.usage + assert usage is not None and usage.prompt_tokens is not None and usage.completion_tokens is not None + assert row.team_id == traffic.team_id + assert row.status == "success" + assert row.cache_hit != "True" + assert row.prompt_tokens == usage.prompt_tokens + assert row.completion_tokens == usage.completion_tokens + assert row.total_tokens == usage.total_tokens + expected_cost: Final = usage.prompt_tokens * INPUT_RATE + usage.completion_tokens * OUTPUT_RATE + assert row.spend is not None and isclose(row.spend, expected_cost, rel_tol=1e-6, abs_tol=1e-9) + + for response in traffic.responses: + assert_response(response) 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 c5d76d44580..8a91e53e7d7 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 @@ -17,13 +17,13 @@ fails the test; a pricing or token-count drift does not. import time from collections.abc import Callable -from concurrent.futures import ThreadPoolExecutor +from math import isclose +from typing import Final import pytest - -from e2e_http import Result, Success +from e2e_http import Success from lifecycle import ResourceManager -from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams +from models import LiteLLMParamsBody, SpendLogs, SpendLogsParams from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap pytestmark = pytest.mark.e2e @@ -280,51 +280,22 @@ def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> N ), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}" +@pytest.mark.replayable @pytest.mark.covers("quota_management.spend_tracking.concurrent_burst.loses_no_spend") def test_burst_of_concurrent_calls_loses_no_spend( - client: SpendClient, scoped_key: str + client: SpendClient, resources: ResourceManager ) -> None: - """Six concurrent calls on one key: every call lands its own spend row under a - distinct request_id and the key aggregate equals the sum of the rows. - Sequential accuracy is covered by test_key_spend_equals_sum_of_logs; this pins - the concurrent increment path (parallel writers racing on one key's counter), - where a lost update can never be reproduced by sequential calls.""" - burst = 6 + from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic - def call(idx: int) -> Result[ChatResponse]: - return client.chat( - scoped_key, - "gemini-2.5-flash", - f"burst call {idx} {unique_marker()}", - max_tokens=16, - ) + traffic: Final = create_traffic(client, resources) - with ThreadPoolExecutor(max_workers=burst) as pool: - results = tuple(pool.map(call, range(burst))) - failed = [r for r in results if not is_ok(r)] - assert not failed, f"{len(failed)}/{burst} burst calls failed; first: {failed[0]}" + def assert_team(team: TeamTraffic) -> None: + assert_logs_match(client, team) + key_spend: Final = client.poll_key_spend(team.key, minimum=team.spend * 0.999999) + assert isclose(key_spend, team.spend, rel_tol=1e-6, abs_tol=1e-9) - rows = client.poll_logs_for_key( - scoped_key, - min_rows=burst, - predicate=lambda rs: len([r for r in rs if (r.spend or 0) > 0]) >= burst, - ) - costed = [r for r in rows if (r.spend or 0) > 0] - assert len(costed) >= burst, ( - f"only {len(costed)}/{burst} burst calls produced a costed row - " - f"rows lost under concurrency: {_summarize(rows)}" - ) - request_ids = [r.request_id for r in costed] - assert len(set(request_ids)) == len(request_ids), ( - f"concurrent rows collapsed onto shared request_ids: {_summarize(rows)}" - ) - - logs_total = sum((r.spend or 0) for r in rows) - key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999) - assert _approx_equal(key_spend, logs_total), ( - f"key aggregate {key_spend} != sum of {len(rows)} rows {logs_total} - " - f"spend increments lost under concurrency: {_summarize(rows)}" - ) + for team in traffic: + assert_team(team) @pytest.mark.covers("quota_management.spend_tracking.pagination.keeps_total") diff --git a/tests/e2e/quota_management/spend_tracking/test_team_daily_activity_e2e.py b/tests/e2e/quota_management/spend_tracking/test_team_daily_activity_e2e.py index ed0a6af4ec9..ef635e59743 100644 --- a/tests/e2e/quota_management/spend_tracking/test_team_daily_activity_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_team_daily_activity_e2e.py @@ -7,13 +7,18 @@ missing start/end dates are rejected. from __future__ import annotations +import time from datetime import datetime, timedelta, timezone +from math import isclose +from typing import Final import pytest from e2e_http import ProbeResult -from models import DateRangeParams +from lifecycle import ResourceManager +from proxy_client import Converged, await_converged from pydantic import BaseModel from spend_e2e_client import SpendClient +from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic pytestmark = pytest.mark.e2e @@ -24,22 +29,45 @@ class TeamDailyActivityParams(BaseModel): start_date: str | None = None end_date: str | None = None page: int = 1 + page_size: int = 1 + team_ids: str | None = None class TeamDailyActivityRow(BaseModel): date: str metrics: TeamDailyActivityMetrics + breakdown: TeamDailyActivityBreakdown class TeamDailyActivityMetrics(BaseModel): spend: float total_tokens: int + prompt_tokens: int + completion_tokens: int + api_requests: int + successful_requests: int + failed_requests: int + + +class TeamDailyActivityEntity(BaseModel): + metrics: TeamDailyActivityMetrics + + +class TeamDailyActivityBreakdown(BaseModel): + entities: dict[str, TeamDailyActivityEntity] class TeamDailyActivityMetadata(BaseModel): page: int total_pages: int has_more: bool + total_spend: float + total_prompt_tokens: int + total_completion_tokens: int + total_tokens: int + total_api_requests: int + total_successful_requests: int + total_failed_requests: int class TeamDailyActivityResponse(BaseModel): @@ -47,32 +75,128 @@ class TeamDailyActivityResponse(BaseModel): metadata: TeamDailyActivityMetadata -def _range_days(days: int) -> DateRangeParams: - end = datetime.now(timezone.utc).date() - start = end - timedelta(days=days) - return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat()) - - def _probe(client: SpendClient, params: BaseModel) -> ProbeResult: return client.proxy.transport.probe(ROUTE, params=params) class TestTeamDailyActivity: + @pytest.mark.replayable @pytest.mark.covers("mgmt.team.daily_activity.happy_path") - @pytest.mark.parametrize("days", [1, 7, 30]) - def test_valid_date_range_returns_results_and_metadata(self, client: SpendClient, days: int) -> None: - result = _probe(client, _range_days(days)) - assert result.status_code == 200, ( - f"{ROUTE} range={days}d must be 200, got {result.status_code}: {result.body[:600]}" + def test_valid_date_range_returns_results_and_metadata( + self, client: SpendClient, resources: ResourceManager + ) -> None: + started: Final = datetime.now(timezone.utc).date() + traffic: Final = create_traffic(client, resources) + for team in traffic: + assert_logs_match(client, team) + ended: Final = datetime.now(timezone.utc).date() + team_ids: Final = ",".join(team.team_id for team in traffic) + + def fetch( + page: int, start: str = (started - timedelta(days=1)).isoformat(), end: str = ended.isoformat() + ) -> TeamDailyActivityResponse: + result: Final = _probe( + client, + TeamDailyActivityParams( + start_date=start, + end_date=end, + page=page, + page_size=1, + team_ids=team_ids, + ), + ) + assert result.status_code == 200, f"daily activity failed: {result.status_code} {result.body[:300]}" + return TeamDailyActivityResponse.model_validate_json(result.body) + + def pages() -> tuple[TeamDailyActivityResponse, ...]: + first: Final = fetch(1) + assert first.metadata.total_pages <= len(traffic) * 2, "unexpected extra scoped daily groups" + return (first, *(fetch(page) for page in range(2, first.metadata.total_pages + 1))) + + outcome: Final = await_converged( + pages, + converged=lambda values: ( + sum(page.metadata.total_api_requests for page in values) >= sum(len(team.responses) for team in traffic) + ), + timeout=client.proxy.poll_timeout, + interval=client.proxy.poll_interval, + now=time.monotonic, + sleep=time.sleep, ) - parsed = TeamDailyActivityResponse.model_validate_json(result.body) - assert parsed.metadata.page == 1 - assert parsed.metadata.total_pages >= 1 - if parsed.results: - first = parsed.results[0] - assert first.date - assert first.metrics.spend >= 0 - assert first.metrics.total_tokens >= 0 + observed: Final = outcome.result if isinstance(outcome, Converged) else outcome.last_result + assert observed is not None, "daily aggregation must return a response before the deadline" + + assert len(observed) >= 2, "two teams must exercise a page boundary" + + def assert_page(index: int, page: TeamDailyActivityResponse) -> None: + assert page.metadata.page == index + assert page.metadata.total_pages == len(observed) + assert page.metadata.has_more == (index < len(observed)) + assert len(page.results) == 1, "each fetched daily group must appear in results" + row: Final = page.results[0] + assert started <= datetime.fromisoformat(row.date).date() <= ended + assert len(row.breakdown.entities) == 1 + assert row.metrics.total_tokens == page.metadata.total_tokens + assert row.metrics.prompt_tokens == page.metadata.total_prompt_tokens + assert row.metrics.completion_tokens == page.metadata.total_completion_tokens + assert row.metrics.api_requests == page.metadata.total_api_requests + assert row.metrics.successful_requests == page.metadata.total_successful_requests + assert row.metrics.failed_requests == page.metadata.total_failed_requests + assert isclose(row.metrics.spend, page.metadata.total_spend, rel_tol=1e-6, abs_tol=1e-9) + + for index, page in enumerate(observed, 1): + assert_page(index, page) + + entities: Final = tuple( + (team_id, entity.metrics) + for page in observed + for row in page.results + for team_id, entity in row.breakdown.entities.items() + ) + assert frozenset(team_id for team_id, _ in entities) == frozenset(team.team_id for team in traffic) + + def assert_team(team: TeamTraffic) -> None: + metrics: Final = tuple(metrics for team_id, metrics in entities if team_id == team.team_id) + assert sum(m.api_requests for m in metrics) == len(team.responses) + assert sum(m.successful_requests for m in metrics) == len(team.responses) + assert sum(m.failed_requests for m in metrics) == 0 + assert sum(m.prompt_tokens for m in metrics) == team.prompt_tokens + assert sum(m.completion_tokens for m in metrics) == team.completion_tokens + assert sum(m.total_tokens for m in metrics) == team.prompt_tokens + team.completion_tokens + assert isclose(sum(m.spend for m in metrics), team.spend, rel_tol=1e-6, abs_tol=1e-9) + + for team in traffic: + assert_team(team) + + assert isclose( + sum(page.metadata.total_spend for page in observed), + sum(team.spend for team in traffic), + rel_tol=1e-6, + abs_tol=1e-9, + ) + assert sum(page.metadata.total_tokens for page in observed) == sum( + team.prompt_tokens + team.completion_tokens for team in traffic + ) + + for days in (7, 30): + assert ( + tuple(fetch(page, (started - timedelta(days=days)).isoformat()) for page in range(1, len(observed) + 1)) + == observed + ), f"{days}-day activity must preserve the same isolated groups and totals" + + empty_date: Final = (started - timedelta(days=7)).isoformat() + empty: Final = fetch(1, empty_date, empty_date) + assert empty.results == [] + assert empty.metadata.total_pages == 0 + assert empty.metadata.page == 1 + assert not empty.metadata.has_more + assert empty.metadata.total_spend == 0 + assert empty.metadata.total_tokens == 0 + assert empty.metadata.total_api_requests == 0 + assert empty.metadata.total_prompt_tokens == 0 + assert empty.metadata.total_completion_tokens == 0 + assert empty.metadata.total_successful_requests == 0 + assert empty.metadata.total_failed_requests == 0 @pytest.mark.covers("mgmt.team.daily_activity.missing_start_date_rejected") def test_missing_start_date_is_rejected(self, client: SpendClient) -> None: diff --git a/tests/e2e/test_e2e_http.py b/tests/e2e/test_e2e_http.py index 81cd6c8d3d1..7201da84924 100644 --- a/tests/e2e/test_e2e_http.py +++ b/tests/e2e/test_e2e_http.py @@ -29,6 +29,7 @@ from e2e_http import ( request_with_retry, streaming_outcome, wire_body, + without_retries, ) from pydantic import BaseModel, TypeAdapter @@ -56,6 +57,15 @@ def _issue_from(responses: Sequence[FakeResponse]) -> Callable[[], FakeResponse] class TestTransientRetryPolicy: + def test_qualification_disables_retries_and_restores_the_default(self) -> None: + responses: Final = (FakeResponse(529), FakeResponse(200)) + sleep: Final = SleepRecorder() + with without_retries(): + assert request_with_retry(_issue_from(responses), sleep=sleep) is responses[0] + assert sleep.delays == () + assert request_with_retry(_issue_from(responses), sleep=sleep) is responses[1] + assert sleep.delays == (0.5,) + def test_transient_set_is_only_statuses_the_proxy_cannot_emit(self) -> None: assert TRANSIENT_STATUSES == frozenset({529}) assert 429 not in TRANSIENT_STATUSES diff --git a/tests/e2e/test_idp.py b/tests/e2e/test_idp.py index 33a09a0f13a..cf2d4f3118a 100644 --- a/tests/e2e/test_idp.py +++ b/tests/e2e/test_idp.py @@ -4,12 +4,20 @@ these carry no `e2e` marker and run everywhere.""" from __future__ import annotations +import os +import signal +import socket +import subprocess +import sys +import time +from builtins import ExceptionGroup from collections.abc import Callable, Generator from contextlib import ExitStack, contextmanager from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path from queue import SimpleQueue from threading import Thread -from typing import Final +from typing import Final, Literal import pytest from e2e_http import ExternalWrite @@ -18,6 +26,8 @@ from idp import ( KEYCLOAK_ADMIN_USER_ENV, KEYCLOAK_REALM_ENV, KEYCLOAK_URL_ENV, + BrowserClientBody, + Discovery, Keycloak, PasswordCredential, UserCreateBody, @@ -60,24 +70,48 @@ def _idp_server( ) -> Generator[tuple[Keycloak, SimpleQueue[str]]]: """Exercise provisioning failures through the same HTTP transport as live tests.""" deletions: SimpleQueue[str] = SimpleQueue() + clients: SimpleQueue[BrowserClientBody] = SimpleQueue() class Handler(BaseHTTPRequestHandler): def log_message(self, format: str, *args: object) -> None: pass def do_POST(self) -> None: - self.rfile.read(int(self.headers.get("Content-Length", "0"))) + body: Final = self.rfile.read(int(self.headers.get("Content-Length", "0"))) if self.path.endswith("/token"): self.send_response(admin_status) self.end_headers() self.wfile.write(b'{"access_token":"synthetic-harness-token"}') else: + if self.path.endswith("/clients"): + clients.put(BrowserClientBody.model_validate_json(body)) self.send_response(user_status if self.path.endswith("/users") else 201) self.send_header("Location", f"{self.path}/resource-1") self.end_headers() if user_status != 201 and self.path.endswith("/users"): self.wfile.write(b"injected create failure") + def do_GET(self) -> None: + self.send_response(200) + self.end_headers() + if "/clients/" in self.path: + client: Final = clients.get_nowait() + clients.put(client) + self.wfile.write(client.model_dump_json(by_alias=True).encode()) + else: + issuer: Final = f"http://127.0.0.1:{server.server_port}/realms/test" + self.wfile.write( + Discovery( + issuer=issuer, + authorization_endpoint=f"{issuer}/auth", + token_endpoint=f"{issuer}/token", + userinfo_endpoint=f"{issuer}/userinfo", + jwks_uri=f"{issuer}/certs", + ) + .model_dump_json() + .encode() + ) + def do_DELETE(self) -> None: deletions.put(self.path) self.send_response(delete_status) @@ -115,6 +149,68 @@ def test_partial_provisioning_removes_the_group_when_user_creation_fails() -> No assert deletions.empty() +@pytest.mark.parametrize( + ("exit_mode", "ignore_termination"), (("normal", False), ("parent", False), ("group", False), ("parent", True)) +) +def test_oidc_launcher_removes_client_on_exit_and_termination( + tmp_path: Path, exit_mode: Literal["normal", "parent", "group"], ignore_termination: bool +) -> None: + ready: Final = tmp_path / "ready" + descendant_command: Final = ( + "import signal,socket,time; from pathlib import Path; " + + ("signal.signal(signal.SIGTERM, signal.SIG_IGN); " if ignore_termination else "") + + "listener=socket.socket(); listener.bind(('127.0.0.1',0)); listener.listen(); " + f"Path({str(ready)!r}).write_text(str(listener.getsockname()[1])); time.sleep(120)" + ) + child_command: Final = ( + "import os,subprocess,sys,time; from pathlib import Path; " + 'assert os.environ["GENERIC_CLIENT_SECRET"]; ' + 'assert os.environ["GENERIC_CLIENT_USE_PKCE"] == "true"; ' + f"subprocess.Popen([sys.executable, '-c', {descendant_command!r}]); " + f"ready=Path({str(ready)!r})\n" + "while not ready.exists(): time.sleep(0.05)\n" + + ("raise SystemExit(7)" if exit_mode == "normal" else "time.sleep(120)") + ) + with _idp_server() as (idp, deletions): + with subprocess.Popen( + [ + sys.executable, + str(Path(__file__).with_name("idp.py")), + "http://127.0.0.1:9999", + sys.executable, + "-c", + child_command, + ], + env={ + **os.environ, + KEYCLOAK_URL_ENV: idp.base_url, + KEYCLOAK_REALM_ENV: idp.realm, + KEYCLOAK_ADMIN_USER_ENV: idp.admin_username, + KEYCLOAK_ADMIN_PASSWORD_ENV: idp.admin_password, + }, + start_new_session=True, + ) as process: + try: + deadline: Final = time.monotonic() + 15 + while not ready.exists() and time.monotonic() < deadline and process.poll() is None: + time.sleep(0.05) + assert ready.exists(), "OIDC child did not start" + if exit_mode == "parent": + process.terminate() + elif exit_mode == "group": + os.killpg(process.pid, signal.SIGTERM) + assert process.wait(timeout=15) == (7 if exit_mode == "normal" else 143) + with socket.socket() as connection: + connection.settimeout(1) + assert connection.connect_ex(("127.0.0.1", int(ready.read_text()))) != 0 + finally: + if process.poll() is None: + os.killpg(process.pid, signal.SIGKILL) + process.wait(timeout=5) + assert deletions.get(timeout=5) == "/admin/realms/test/clients/resource-1" + assert deletions.empty() + + def test_successful_provisioning_cleans_up_user_before_group() -> None: with _idp_server() as (idp, deletions): with ExitStack() as cleanup: @@ -134,6 +230,43 @@ def test_cleanup_failure_is_visible() -> None: idp.delete_group("group") +def test_strict_cleanup_reports_each_failure_and_continues() -> None: + from lifecycle import ResourceManager + from proxy_client import build_proxy_client + + with _idp_server(delete_status=500) as (idp, deletions): + resources: Final = ResourceManager(client=build_proxy_client(), strict_cleanup=True) + strict: Final = idp.with_strict_cleanup() + resources.defer(lambda: strict.delete_group("group")) + resources.defer(lambda: strict.delete_user("user")) + with pytest.raises(ExceptionGroup, match="Resource cleanup failed") as error: + resources.teardown() + assert len(error.value.exceptions) == 2 + assert deletions.get_nowait() == "/admin/realms/test/users/user" + assert deletions.get_nowait() == "/admin/realms/test/groups/group" + + +@pytest.mark.parametrize("groups", ((), ("one",), ("one", "two"))) +def test_provisioning_records_zero_one_or_multiple_groups(groups: tuple[str, ...]) -> None: + with _idp_server() as (idp, deletions): + with ExitStack() as cleanup: + + def defer(callback: Callable[[], object]) -> None: + cleanup.callback(callback) + + identity: Final = idp.provision_groups( + marker="memberships", + groups=groups, + defer=defer, + ) + assert identity.groups == groups + assert len(identity.group_ids) == len(groups) + assert deletions.get_nowait() == "/admin/realms/test/users/resource-1" + for _ in groups: + assert deletions.get_nowait() == "/admin/realms/test/groups/resource-1" + assert deletions.empty() + + def test_expired_admin_credentials_do_not_abort_remaining_cleanups() -> None: with _idp_server(admin_status=401) as (idp, _): cleanup: Final = ExitStack() diff --git a/tests/e2e/test_provider_edge.py b/tests/e2e/test_provider_edge.py index 18f72ac0e7a..81be81e7b59 100644 --- a/tests/e2e/test_provider_edge.py +++ b/tests/e2e/test_provider_edge.py @@ -31,6 +31,7 @@ from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path +from types import MappingProxyType from typing import Final import pytest @@ -81,9 +82,14 @@ def json_object(body: bytes) -> dict[str, object]: class _FakeProvider(ThreadingHTTPServer): daemon_threads = True - def __init__(self, bind: tuple[str, int]) -> None: + def __init__(self, bind: tuple[str, int], *, echo_request: bool = True) -> None: super().__init__(bind, _FakeProviderHandler) self.hits: list[str] = [] + self.echo_request = echo_request + self.requests: tuple[tuple[Mapping[str, str], bytes], ...] = () + + def capture_request(self, headers: Mapping[str, str], body: bytes) -> None: + self.requests = (*self.requests, (MappingProxyType(dict(headers)), body)) class _FakeProviderHandler(BaseHTTPRequestHandler): @@ -101,8 +107,11 @@ class _FakeProviderHandler(BaseHTTPRequestHandler): length = int(self.headers.get("content-length") or "0") body = self.rfile.read(length) if length else b"" provider.hits.append(f"{self.command} {self.path}") - payload = json.dumps( + provider.capture_request(dict(self.headers.items()), body) + payload: Final = json.dumps( {"echo": body.decode("utf-8"), "path": self.path, "hit": len(provider.hits)} + if provider.echo_request + else {"ok": True} ).encode() self.send_response(200) self.send_header("content-type", "application/json") @@ -117,8 +126,8 @@ class _FakeProviderHandler(BaseHTTPRequestHandler): @contextmanager -def fake_provider() -> Generator[_FakeProvider]: - server = _FakeProvider(("127.0.0.1", 0)) +def fake_provider(*, echo_request: bool = True) -> Generator[_FakeProvider]: + server = _FakeProvider(("127.0.0.1", 0), echo_request=echo_request) thread = threading.Thread(target=server.serve_forever, daemon=True) thread.start() try: diff --git a/tests/e2e/test_proxy_client.py b/tests/e2e/test_proxy_client.py index 1b0133f12cb..0c4aed5bd65 100644 --- a/tests/e2e/test_proxy_client.py +++ b/tests/e2e/test_proxy_client.py @@ -11,19 +11,53 @@ injected clock, so nothing here monkeypatches anything. from __future__ import annotations -from collections.abc import Iterable, Mapping +import json +from builtins import ExceptionGroup +from collections.abc import Callable, Generator, Iterable, Mapping +from contextlib import contextmanager from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from itertools import chain, repeat +from queue import SimpleQueue +from threading import Thread from types import MappingProxyType from typing import Final, cast import pytest from e2e_config import parse_replica_urls -from e2e_http import Result, Success -from models import KeyInfo, KeyInfoResponse, ModelListEntry, ModelsListResponse +from e2e_http import NoBody, Result, Success, without_retries +from idp import Keycloak +from lifecycle import ResourceManager +from management.jwt_actors import ActorFactory +from management.management_client import ManagementClient +from models import ( + ConnectionTestBody, + CredentialCreateBody, + KeyGenerateBody, + KeyInfo, + KeyInfoResponse, + KeyUpdateBody, + LiteLLMParamsBody, + McpServerCreateBody, + McpServerUpdateBody, + ModelListEntry, + ModelsListResponse, + OrgNewBody, + OrgUpdateBody, + SpendLogsParams, + TagNewBody, + TeamNewBody, + TeamUpdateBody, + ToolsetCreateBody, + ToolsetUpdateBody, + UserNewBody, + UserUpdateBody, +) from proxy_client import ( - ConvergeOutcome, + Caller, Converged, + ConvergeOutcome, + CredentialKind, EverywhereConverged, ModelsPoller, NeverConvergedOn, @@ -42,6 +76,115 @@ from proxy_client import ( ) from transport import Transport + +@contextmanager +def caller_boundary( + status: int = 200, bodies: SimpleQueue[bytes] | None = None, *, delete_status: int | None = None +) -> Generator[tuple[ManagementClient, SimpleQueue[str]]]: + received: Final[SimpleQueue[str]] = SimpleQueue() + + class Handler(BaseHTTPRequestHandler): + def log_message(self, format: str, *args: object) -> None: + pass + + def do_GET(self) -> None: + received.put(self.headers.get("Authorization", "")) + self.send_response(delete_status if self.path == "/key/delete" and delete_status is not None else status) + self.end_headers() + self.wfile.write( + b'{"key":"owned","info":{"key_alias":"owned"},"data":[{"id":"owned"}],"team_id":"owned","team_info":{},"model_id":"owned"}' + ) + + def do_POST(self) -> None: + body: Final = self.rfile.read(int(self.headers.get("Content-Length", "0"))) + if bodies is not None: + bodies.put(body) + self.do_GET() + + do_PATCH = do_POST + do_PUT = do_POST + do_DELETE = do_POST + + server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread: Final = Thread(target=lambda: server.serve_forever(poll_interval=0.01), daemon=True) + thread.start() + url: Final = f"http://127.0.0.1:{server.server_port}" + proxy: Final = build_proxy_client( + base_url=url, control_plane_base_url=url, replica_urls=(url,), master_key="bootstrap" + ) + try: + yield ManagementClient(proxy=proxy, master_key="bootstrap"), received + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) + + +class TestBoundManagementCaller: + def test_strict_key_cleanup_accepts_missing_only_when_requested(self) -> None: + with caller_boundary(delete_status=404) as (bootstrap, received), without_retries(): + with pytest.raises(AssertionError): + bootstrap.delete_key_strict("owned") + bootstrap.delete_key_strict("owned", missing_ok=True) + assert (received.get_nowait(), received.get_nowait()) == ("Bearer bootstrap", "Bearer bootstrap") + + def test_actor_key_cleanup_reports_failure_and_continues(self) -> None: + with caller_boundary(delete_status=500) as (bootstrap, received), without_retries(): + resources: Final = ResourceManager(client=bootstrap.proxy, strict_cleanup=True) + remaining: SimpleQueue[str] = SimpleQueue() + resources.defer(lambda: remaining.put("cleaned")) + factory: Final = ActorFactory( + bootstrap=bootstrap, + idp=Keycloak(base_url="http://unused.test", realm="test", admin_username="test", admin_password="test"), + resources=resources, + ) + assert factory.key().key == "owned" + with pytest.raises(ExceptionGroup, match="Resource cleanup failed") as failure: + resources.teardown() + assert len(failure.value.exceptions) == 1 + assert remaining.get_nowait() == "cleaned" + assert (received.get_nowait(), received.get_nowait()) == ("Bearer bootstrap", "Bearer bootstrap") + + @pytest.mark.parametrize("kind", ("direct_jwt", "virtual_key", "dashboard_session")) + def test_direct_delegated_and_replica_reads_keep_the_bound_caller(self, kind: CredentialKind) -> None: + with caller_boundary() as (bootstrap, received): + caller: Final = Caller(credential="synthetic-caller", kind=kind, role="internal_user", tenant="tenant-a") + bound: Final = bootstrap.with_caller(caller) + bound.update_key(KeyUpdateBody(key="owned", key_alias="updated")) + bound.proxy.key_info("owned") + bound.proxy.read_back_everywhere( + "/key/info", + params=KeyUpdateBody(key="owned"), + response_type=KeyInfoResponse, + converged=lambda result: isinstance(result, Success), + ) + bound.proxy.read_body_back_everywhere( + "/key/info", KeyInfoResponse, settled=lambda result: result.info.key_alias == "owned" + ) + assert tuple(received.get_nowait() for _ in range(4)) == ("Bearer synthetic-caller",) * 4 + assert received.empty() + bootstrap.proxy.key_info("owned") + assert received.get_nowait() == "Bearer bootstrap" + + def test_explicit_override_wins_without_rebinding_or_changing_master(self) -> None: + with caller_boundary() as (bootstrap, received): + bound: Final = bootstrap.with_caller(Caller(credential="bound", kind="direct_jwt", role="internal_user")) + bound.update_key(KeyUpdateBody(key="owned"), caller_key="override") + bound.proxy.key_info("owned") + assert received.get_nowait() == "Bearer override" + assert received.get_nowait() == "Bearer bound" + assert bound.master_key == "bootstrap" + + def test_credentials_are_absent_from_binding_and_header_diagnostics(self) -> None: + with caller_boundary() as (bootstrap, _): + caller: Final = Caller(credential="private-value", kind="direct_jwt", role="internal_user") + bound: Final = bootstrap.with_caller(caller) + assert "private-value" not in repr(caller) + assert "private-value" not in repr(bound) + assert "private-value" not in repr(bound.proxy.management_headers()) + assert "bootstrap" not in repr(bound) + + MODEL: Final = "gpt-under-test" _NO_TRANSPORTS: Final = cast(Transport, None) TIMEOUT: Final = 10.0 @@ -275,3 +418,166 @@ class TestReplicasFor: client: Final = ProxyClient(transport=_NO_TRANSPORTS, replicas={}, control_replicas={}) with pytest.raises(AssertionError, match="no replica is configured"): _ = client.replicas_for("/v1/models") + + +MANAGEMENT_OPERATIONS: Final[tuple[tuple[str, Callable[[ManagementClient], object]], ...]] = ( + ("generate_key", lambda c: c.generate_key(KeyGenerateBody())), + ("llm_only_key", lambda c: c.llm_only_key()), + ("update_key", lambda c: c.update_key(KeyUpdateBody(key="owned"))), + ("update_key_models", lambda c: c.update_key_models("owned", [])), + ("key_info", lambda c: c.key_info_as("owned")), + ("delete_key_strict", lambda c: c.delete_key_strict("owned")), + ("delete_model_strict", lambda c: c.delete_model_strict("owned")), + ( + "connection_test", + lambda c: c.connection_test( + ConnectionTestBody(litellm_params=LiteLLMParamsBody(model="synthetic"), mode="chat") + ), + ), + ("block_key", lambda c: c.block_key("owned")), + ("regenerate_key", lambda c: c.regenerate_key("owned")), + ("reset_key_spend", lambda c: c.reset_key_spend("owned", 0)), + ("key_list", lambda c: c.key_list("owned")), + ("key_alias_count", lambda c: c.key_alias_count("owned")), + ("create_team", lambda c: c.create_team(TeamNewBody(team_alias="owned"))), + ("update_team", lambda c: c.update_team(TeamUpdateBody(team_id="owned", team_alias="updated"))), + ("delete_team", lambda c: c.delete_team("owned")), + ("team_info", lambda c: c.team_info("owned")), + ("team_list_ids", lambda c: c.team_list_ids()), + ("team_info_status", lambda c: c.team_info_status("owned")), + ("add_team_member", lambda c: c.add_team_member("owned", "user")), + ("delete_team_member", lambda c: c.delete_team_member("owned", "user")), + ("create_user", lambda c: c.create_user(UserNewBody(user_email="actor@example.com", user_role="internal_user"))), + ("create_customer", lambda c: c.create_customer("owned")), + ("customer_info", lambda c: c.customer_info("owned")), + ("delete_customer", lambda c: c.delete_customer("owned")), + ("update_user", lambda c: c.update_user(UserUpdateBody(user_id="owned", user_role="internal_user"))), + ("delete_user", lambda c: c.delete_user("owned")), + ("delete_user_strict", lambda c: c.delete_user_strict("owned")), + ("user_info", lambda c: c.user_info("owned")), + ("user_count", lambda c: c.user_count("owned")), + ("user_list_ids", lambda c: c.user_list_ids("owned")), + ("create_org", lambda c: c.create_org(OrgNewBody(organization_alias="owned"))), + ("update_org", lambda c: c.update_org(OrgUpdateBody(organization_id="owned", organization_alias="updated"))), + ("delete_org", lambda c: c.delete_org("owned")), + ("org_info", lambda c: c.org_info("owned")), + ("org_info_status", lambda c: c.org_info_status("owned")), + ("create_tag", lambda c: c.create_tag(TagNewBody(name="owned"))), + ("delete_tag", lambda c: c.delete_tag("owned")), + ("tag_list", lambda c: c.tag_list()), + ("create_mcp_server", lambda c: c.create_mcp_server(McpServerCreateBody(alias="owned", url="http://example.test"))), + ("update_mcp_server", lambda c: c.update_mcp_server(McpServerUpdateBody(server_id="owned", alias=None))), + ("delete_mcp_server", lambda c: c.delete_mcp_server("owned")), + ("proxy.generate_key", lambda c: c.proxy.generate_key(KeyGenerateBody())), + ("proxy.delete_key", lambda c: c.proxy.delete_key("owned")), + ("proxy.delete_customers", lambda c: c.proxy.delete_customers(["owned"])), + ("proxy.key_info", lambda c: c.proxy.key_info("owned")), + ("proxy.memory_summary", lambda c: c.proxy.memory_summary_everywhere()), + ("proxy.model_info", lambda c: c.proxy.model_info()), + ("proxy.model_cost_map", lambda c: c.proxy.model_cost_map()), + ("proxy.create_model", lambda c: c.proxy.create_model("owned", LiteLLMParamsBody(model="synthetic"))), + ("proxy.update_model", lambda c: c.proxy.update_model("owned", LiteLLMParamsBody(model="synthetic"))), + ("proxy.delete_model", lambda c: c.proxy.delete_model("owned")), + ("proxy.create_toolset", lambda c: c.proxy.create_toolset(ToolsetCreateBody(toolset_name="owned", tools=[]))), + ("proxy.update_toolset", lambda c: c.proxy.update_toolset(ToolsetUpdateBody(toolset_id="owned", description=None))), + ("proxy.delete_toolset", lambda c: c.proxy.delete_toolset("owned")), + ( + "proxy.create_credential", + lambda c: c.proxy.create_credential(CredentialCreateBody(credential_name="owned", credential_values={})), + ), + ("proxy.delete_credential", lambda c: c.proxy.delete_credential("owned")), + ("proxy.create_team", lambda c: c.proxy.create_team(TeamNewBody(team_alias="owned"))), + ("proxy.delete_team", lambda c: c.proxy.delete_team("owned")), + ("proxy.delete_user", lambda c: c.proxy.delete_user("owned")), + ("proxy.spend_logs", lambda c: c.proxy.spend_logs(SpendLogsParams(api_key="owned"))), + ("proxy.probe", lambda c: c.proxy.probe("/user/info", params=NoBody())), +) + + +@pytest.mark.parametrize( + ("name", "operation"), MANAGEMENT_OPERATIONS, ids=tuple(name for name, _ in MANAGEMENT_OPERATIONS) +) +@pytest.mark.parametrize("kind", ("master", "direct_jwt", "virtual_key", "dashboard_session")) +def test_management_operations_send_the_selected_credential( + name: str, + operation: Callable[[ManagementClient], object], + kind: CredentialKind, +) -> None: + with caller_boundary(status=401) as (bootstrap, received), without_retries(): + client: Final = ( + bootstrap + if kind == "master" + else bootstrap.with_caller(Caller(credential=f"synthetic-{kind}", kind=kind, role="internal_user")) + ) + try: + operation(client) + except AssertionError: + pass + expected: Final = "Bearer bootstrap" if kind == "master" else f"Bearer synthetic-{kind}" + assert received.get_nowait() == expected, name + assert received.empty(), "an unauthorized request must not be retried" + + +class TestSplitCallerPropagation: + def test_control_and_data_replica_readers_keep_the_caller(self) -> None: + with caller_boundary() as (data, data_headers), caller_boundary() as (control, control_headers): + data_url: Final = next(iter(data.proxy.replicas)) + control_url: Final = next(iter(control.proxy.replicas)) + proxy: Final = build_proxy_client( + base_url=data_url, + control_plane_base_url=control_url, + replica_urls=(data_url,), + master_key="bootstrap", + ).with_caller(Caller(credential="tenant-token", kind="direct_jwt", role="team_member")) + proxy.key_info("owned") + proxy.read_body_back_everywhere( + "/key/info", KeyInfoResponse, settled=lambda info: info.info.key_alias == "owned" + ) + proxy.read_back_everywhere( + "/key/info", + params=NoBody(), + response_type=KeyInfoResponse, + converged=lambda result: isinstance(result, Success), + ) + assert control_headers.get_nowait() == "Bearer tenant-token" + assert control_headers.get_nowait() == "Bearer tenant-token" + assert data_headers.get_nowait() == "Bearer tenant-token" + assert control_headers.empty() and data_headers.empty() + + def test_successful_team_and_model_polling_uses_the_bound_caller(self) -> None: + with caller_boundary() as (bootstrap, received): + bound: Final = bootstrap.with_caller(Caller(credential="caller", kind="direct_jwt", role="proxy_admin")) + bound.create_team(TeamNewBody(team_alias="owned")) + bound.proxy.create_model("owned", LiteLLMParamsBody(model="synthetic")) + assert tuple(received.get_nowait() for _ in range(4)) == ("Bearer caller",) * 4 + assert received.empty() + + def test_expired_shaped_token_is_sent_once_without_renewal(self) -> None: + with caller_boundary(status=401) as (bootstrap, received): + bound: Final = bootstrap.with_caller( + Caller(credential="expired.payload.signature", kind="direct_jwt", role="internal_user") + ) + result: Final = bound.key_info_as("owned") + assert not isinstance(result, Success) + assert received.get_nowait() == "Bearer expired.payload.signature" + assert received.empty() + + +@pytest.mark.parametrize("operation", ("server", "toolset")) +def test_partial_updates_preserve_explicit_null_at_the_http_boundary(operation: str) -> None: + bodies: Final[SimpleQueue[bytes]] = SimpleQueue() + with caller_boundary(status=401, bodies=bodies) as (bootstrap, _): + try: + if operation == "server": + bootstrap.update_mcp_server(McpServerUpdateBody(server_id="owned", alias=None)) + else: + bootstrap.proxy.update_toolset(ToolsetUpdateBody(toolset_id="owned", description=None)) + except AssertionError: + pass + expected: Final = ( + {"server_id": "owned", "alias": None} + if operation == "server" + else {"toolset_id": "owned", "description": None} + ) + assert json.loads(bodies.get_nowait()) == expected + assert bodies.empty() diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index e8caa801467..037db0c340f 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -7,11 +7,9 @@ client touches requests.* or builds raw dicts; they pass pydantic models here. from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Protocol -from pydantic import BaseModel - import e2e_http from e2e_http import ( URL, @@ -21,6 +19,7 @@ from e2e_http import ( Result, StreamingResponse, ) +from pydantic import BaseModel class Transport(Protocol): @@ -85,7 +84,7 @@ class Transport(Protocol): self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] ) -> Result[R]: ... - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ... + def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult: ... def upload[R: BaseModel]( self, @@ -113,7 +112,7 @@ class Transport(Protocol): @dataclass(frozen=True, slots=True) class HttpTransport: base_url: str - master_key: str + master_key: str = field(repr=False) request_timeout: float = 60.0 def _url(self, path: str) -> URL: @@ -245,10 +244,10 @@ class HttpTransport: timeout=self.request_timeout, ) - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: + def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult: return e2e_http.probe( self._url(path), - headers=self.master, + headers=self.master if headers is None else headers, params=params, timeout=self.request_timeout, ) @@ -434,8 +433,8 @@ class SplitTransport: path, headers=headers, json=json, params=params, stream=stream ) - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: - return self._route(path).probe(path, params=params) + def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult: + return self._route(path).probe(path, params=params, headers=headers) def upload[R: BaseModel]( self, diff --git a/tests/e2e/ui/oidcSetup.ts b/tests/e2e/ui/oidcSetup.ts new file mode 100644 index 00000000000..943fa23cab3 --- /dev/null +++ b/tests/e2e/ui/oidcSetup.ts @@ -0,0 +1,30 @@ +import { chromium, expect } from "@playwright/test"; +import * as fs from "fs"; +import * as path from "path"; + +export default async function oidcSetup() { + const baseURL = process.env.E2E_OIDC_UI_URL; + const issuer = process.env.JWT_ISSUER; + const username = process.env.E2E_OIDC_USERNAME; + const password = process.env.E2E_OIDC_PASSWORD; + if (!baseURL || !issuer || !username || !password) { + throw new Error("The OIDC setup requires a running stack, issuer, and provisioned actor credentials"); + } + const artifactDir = process.env.E2E_UI_ARTIFACT_DIR || "."; + fs.mkdirSync(artifactDir, { recursive: true }); + const browser = await chromium.launch(); + try { + const page = await browser.newPage(); + await page.goto(`${baseURL.replace(/\/$/, "")}/sso/key/generate`); + await expect(page).toHaveURL(new RegExp(`^${issuer.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}/`)); + await page.getByLabel("Username or email").fill(username); + await page.getByLabel("Password", { exact: true }).fill(password); + await page.getByRole("button", { name: "Sign In", exact: true }).click(); + await page.waitForURL((url) => url.origin === new URL(baseURL).origin && url.pathname.startsWith("/ui")); + const statePath = path.join(artifactDir, "oidc.storageState.json"); + await page.context().storageState({ path: statePath }); + fs.chmodSync(statePath, 0o600); + } finally { + await browser.close(); + } +} diff --git a/tests/e2e/ui/playwright.oidc.config.ts b/tests/e2e/ui/playwright.oidc.config.ts new file mode 100644 index 00000000000..0fbe77e9bd2 --- /dev/null +++ b/tests/e2e/ui/playwright.oidc.config.ts @@ -0,0 +1,22 @@ +import { defineConfig, devices } from "@playwright/test"; +import * as path from "path"; + +const baseURL = process.env.E2E_OIDC_UI_URL; +if (!baseURL) throw new Error("E2E_OIDC_UI_URL must point to the running OIDC stack"); + +export default defineConfig({ + testDir: ".", + testMatch: "oidc/**/*.spec.ts", + retries: 0, + workers: 1, + outputDir: path.join(process.env.E2E_UI_ARTIFACT_DIR || ".", "oidc", "test-results"), + globalSetup: require.resolve("./oidcSetup"), + use: { + ...devices["Desktop Chrome"], + baseURL, + storageState: path.join(process.env.E2E_UI_ARTIFACT_DIR || ".", "oidc.storageState.json"), + trace: "off", + screenshot: "off", + video: "off", + }, +}); diff --git a/tests/integration/README.md b/tests/integration/README.md new file mode 100644 index 00000000000..5af4fdb9d06 --- /dev/null +++ b/tests/integration/README.md @@ -0,0 +1,19 @@ +# Integration contracts + +These tests exercise a running gateway, PostgreSQL and Redis with an owned local upstream. CircleCI owns this suite. Tests are grouped by behavior, with no automatic test retries or fallback to paid provider calls + +Use `tests/integration/run.py management`, `accounting` or `providers` to run a selected group. Set `INTEGRATION_PROXY_URL`, `INTEGRATION_UPSTREAM_URL`, `INTEGRATION_MASTER_KEY` and `DATABASE_URL` to an isolated test deployment. The runner selects the new domain directories explicitly; the legacy OCI and sandbox selections remain separate + +Management also requires `INTEGRATION_PEER_URL`, `REDIS_HOST` and `REDIS_PORT`. CircleCI starts two directly addressed proxy processes sharing only that job's stores. The test-only CLI wrapper supplies enterprise route entitlement, following the existing behavior suite's convention. It does not qualify license validation; run it with one worker and no reload + +The generated lifecycle models use 20 examples, eight steps, generation and shrinking, with isolated resources per example. HTTP operation caps include generation and shrinking and exempt cleanup. Local qualification defaults to seed 4106601; CircleCI derives its exploration seed from the checked-out revision. Use `--seed` to reproduce a run. Actual installed Hypothesis version, settings and seed are written beside the execution manifest + +Reuse the existing canned provider handlers through `_support/upstream.py`. It rejects internal request fields and exposes actual received requests for independent assertions. Register every created resource for cleanup immediately, keep expected values independent of production calculations, and assert readback plus the runtime effect of a change + +The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, skipped tests, failed cleanup or a selected test without a passed call fail qualification. Existing GitHub Actions jobs do not own these tests + +Define integration contract IDs and their canonical test nodes in `contracts.json`. Every node must declare the same IDs with `covers`. The runner checks exact collected and passed selections against that mapping. These IDs belong to this CircleCI suite and must not be added to the separate E2E coverage registry. A manifest declaration alone does not mean a test passed + +Provider sentinels currently use the controlled server, not live recordings. The provider shard also runs the existing strict replay controls for changed requests, exhausted interactions, leftover interactions and no provider connection. Future recorded scenarios must use that replay-only implementation; missing recordings cannot fall back to a real provider. The observation endpoint is destructive and the current selection runs serially against one owned upstream + +Fixtures must contain synthetic data only. Keep private incident records and source documents out of code, fixtures, logs and PR descriptions diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/_support/__init__.py b/tests/integration/_support/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/_support/client.py b/tests/integration/_support/client.py new file mode 100644 index 00000000000..8d6744c60a2 --- /dev/null +++ b/tests/integration/_support/client.py @@ -0,0 +1,170 @@ +from __future__ import annotations + +import os +import time +import uuid +from hashlib import sha256 +from collections.abc import Callable, Iterator, Mapping +from contextlib import ExitStack, contextmanager +from dataclasses import dataclass +from typing import Final, TypeVar + +import httpx +from pydantic import JsonValue, TypeAdapter + +from integration._support.database import read_rows + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +T = TypeVar("T") + + +def object_value(value: JsonValue) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_python(value) + + +def string_value(value: JsonValue) -> str: + assert isinstance(value, str), f"Expected a string, received {type(value).__name__}" + return value + + +def eventually(read: Callable[[], T], satisfied: Callable[[T], bool], seconds: float = 10) -> T: + deadline: Final = time.monotonic() + seconds + while True: + observed: Final = read() + if satisfied(observed): + return observed + assert time.monotonic() < deadline, f"State did not converge: {observed!r}" + time.sleep(0.1) + + +@dataclass(frozen=True, slots=True) +class Gateway: + client: httpx.Client + key: str + upstream_url: str + + def request( + self, + method: str, + path: str, + body: Mapping[str, JsonValue] | None = None, + *, + key: str | None = None, + params: Mapping[str, str] | None = None, + ) -> httpx.Response: + return self.client.request( + method, + path, + json=body, + params=params, + headers={"Authorization": f"Bearer {self.key if key is None else key}"}, + ) + + def post(self, path: str, body: Mapping[str, JsonValue], *, key: str | None = None) -> dict[str, JsonValue]: + response: Final = self.request("POST", path, body, key=key) + assert response.status_code == 200, f"POST {path}: {response.status_code} {response.text}" + return JSON_OBJECT.validate_json(response.content) + + def get(self, path: str, params: Mapping[str, str] | None = None) -> dict[str, JsonValue]: + response: Final = self.request("GET", path, params=params) + assert response.status_code == 200, f"GET {path}: {response.status_code} {response.text}" + return JSON_OBJECT.validate_json(response.content) + + def chat(self, model: str, *, key: str | None = None, text: str = "integration control") -> dict[str, JsonValue]: + return self.post( + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": text}]}, + key=key, + ) + + @contextmanager + def scenario(self) -> Iterator[Scenario]: + with ExitStack() as cleanups: + yield Scenario(self, cleanups) + + +@dataclass(frozen=True, slots=True) +class Scenario: + gateway: Gateway + cleanups: ExitStack + + def key(self, **fields: JsonValue) -> str: + created: Final = self.gateway.post("/key/generate", fields) + token: Final = string_value(created["key"]) + self.cleanups.callback(self.delete_key, token) + return token + + def team(self, **fields: JsonValue) -> str: + created: Final = self.gateway.post("/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}", **fields}) + identity: Final = string_value(created["team_id"]) + self.cleanups.callback(self.delete_team, identity) + return identity + + def delete_team(self, identity: str) -> None: + self.gateway.post("/team/delete", {"team_ids": [identity]}) + assert read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s', (identity,)) == [] + + def project(self, team_id: str, **fields: JsonValue) -> str: + created: Final = self.gateway.post( + "/project/new", {"team_id": team_id, "project_alias": f"integration-{uuid.uuid4().hex}", **fields} + ) + identity: Final = string_value(created["project_id"]) + self.cleanups.callback(self.delete_project, identity) + return identity + + def delete_project(self, identity: str) -> None: + response: Final = self.gateway.request("DELETE", "/project/delete", {"project_ids": [identity]}) + assert response.status_code == 200, response.text + assert read_rows('SELECT project_id FROM "LiteLLM_ProjectTable" WHERE project_id = %s', (identity,)) == [] + + def user(self, **fields: JsonValue) -> str: + created: Final = self.gateway.post( + "/user/new", {"user_id": f"integration-{uuid.uuid4().hex}", "auto_create_key": False, **fields} + ) + identity: Final = string_value(created["user_id"]) + self.cleanups.callback(self.delete_user, identity) + return identity + + def delete_user(self, identity: str) -> None: + response: Final = self.gateway.request("POST", "/user/delete", {"user_ids": [identity]}) + assert response.status_code == 200 and response.json() == 1, response.text + assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (identity,)) == [] + + def delete_key(self, token: str) -> None: + self.gateway.post("/key/delete", {"keys": [token]}) + response: Final = self.gateway.request("GET", "/key/info", params={"key": sha256(token.encode()).hexdigest()}) + assert response.status_code == 404, f"Deleted key remains readable: {response.status_code}" + + def delete_model(self, identity: str) -> None: + self.gateway.post("/model/delete", {"id": identity}) + entries: Final = self.gateway.get("/model/info")["data"] + assert isinstance(entries, list) + assert all(object_value(object_value(entry)["model_info"])["id"] != identity for entry in entries) + assert read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,)) == [] + + def model(self, **parameters: JsonValue) -> str: + name: Final = f"integration-{uuid.uuid4().hex}" + created: Final = self.gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": f"{self.gateway.upstream_url}/v1", + **parameters, + }, + "model_info": {}, + }, + ) + identity: Final = string_value(object_value(created["model_info"])["id"]) + self.cleanups.callback(self.delete_model, identity) + return name + + +@contextmanager +def gateway_from_environment() -> Iterator[Gateway]: + url: Final = os.environ["INTEGRATION_PROXY_URL"] + upstream: Final = os.environ["INTEGRATION_UPSTREAM_URL"] + with httpx.Client(base_url=url, timeout=15, trust_env=False) as client: + yield Gateway(client, os.environ["INTEGRATION_MASTER_KEY"], upstream) diff --git a/tests/integration/_support/database.py b/tests/integration/_support/database.py new file mode 100644 index 00000000000..283d26a632e --- /dev/null +++ b/tests/integration/_support/database.py @@ -0,0 +1,14 @@ +import os +from typing import Final + +import psycopg +from psycopg.rows import dict_row +from pydantic import JsonValue, TypeAdapter + +ROWS: Final = TypeAdapter(list[dict[str, JsonValue]]) + + +def read_rows(query: str, parameters: tuple[str, ...]) -> list[dict[str, JsonValue]]: + with psycopg.connect(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()) diff --git a/tests/integration/_support/generation.py b/tests/integration/_support/generation.py new file mode 100644 index 00000000000..afb3ec2e768 --- /dev/null +++ b/tests/integration/_support/generation.py @@ -0,0 +1,52 @@ +from typing import Final +from dataclasses import dataclass +from collections.abc import Iterator, Sequence +from contextlib import contextmanager + +import httpx +from hypothesis import Phase, settings + +from integration._support.client import Gateway + +LIFECYCLE_SETTINGS: Final = settings( + max_examples=20, + stateful_step_count=8, + deadline=None, + database=None, + phases=(Phase.generate, Phase.shrink), + print_blob=True, +) + + +@dataclass(slots=True) +class RequestBudget: + limit: int + requests: int = 0 + cleaning: bool = False + + def observe(self, _request: httpx.Request) -> None: + if self.cleaning: + return + self.requests += 1 + assert self.requests <= self.limit, f"Generated HTTP operation budget exceeded: {self.limit}" + + @contextmanager + def cleanup(self) -> Iterator[None]: + self.cleaning = True + try: + yield + finally: + self.cleaning = False + + +@contextmanager +def bounded_http_requests(gateways: Sequence[Gateway], limit: int) -> Iterator[RequestBudget]: + budget: Final = RequestBudget(limit) + for gateway in gateways: + gateway.client.event_hooks["request"].append(budget.observe) + try: + yield budget + finally: + for gateway in gateways: + gateway.client.event_hooks["request"].remove(budget.observe) + print(f"Generated HTTP operations: {budget.requests}/{budget.limit}; cleanup excluded") diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py new file mode 100644 index 00000000000..b3a82fa4cdd --- /dev/null +++ b/tests/integration/_support/manifest.py @@ -0,0 +1,31 @@ +import json +from pathlib import Path +from typing import Final + +from pydantic import TypeAdapter + +MAPPING: Final = TypeAdapter(dict[str, tuple[str, ...]]) +OWNED_DIRECTORIES: Final = frozenset( + { + "management", + "authorization", + "database", + "pricing", + "spend", + "routing", + "providers", + "streaming", + "configuration", + "mcp", + "observability", + "compatibility", + } +) + + +def contracts() -> dict[str, tuple[str, ...]]: + document: Final = json.loads((Path(__file__).resolve().parents[1] / "contracts.json").read_bytes()) + result: Final = MAPPING.validate_python(document["tests"]) + if not result or any(not values or any(not value.strip() for value in values) for values in result.values()): + raise ValueError("Integration manifest must contain nodes with contract IDs") + return result diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py new file mode 100644 index 00000000000..3139beaeb01 --- /dev/null +++ b/tests/integration/_support/proxy.py @@ -0,0 +1,16 @@ +"""Run the normal single-process CLI with the existing behavior-suite test entitlement.""" + +from unittest.mock import patch + +from litellm import run_server + + +def main() -> None: + with patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts + "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True + ): + run_server() + + +if __name__ == "__main__": + main() diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py new file mode 100644 index 00000000000..8bc4100abfd --- /dev/null +++ b/tests/integration/_support/upstream.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +import argparse +from dataclasses import dataclass, field +from collections import deque +from queue import SimpleQueue +from typing import Final + +import uvicorn +from pydantic import JsonValue, TypeAdapter +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse, Response +from starlette.routing import Route + +from _fake_openai_endpoint_server import chat_completions, completions, embeddings, health, moderations + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +INTERNAL_FIELDS: Final = frozenset( + { + "litellm_params", + "litellm_logging_obj", + "litellm_call_id", + "litellm_metadata", + "proxy_server_request", + "rpm", + "tpm", + "timeout", + "stream_chunk_size", + } +) + + +@dataclass(frozen=True, slots=True) +class Observation: + path: str + authorization: str + body: dict[str, JsonValue] + + +@dataclass(frozen=True, slots=True) +class Provider: + observations: SimpleQueue[Observation] = field(default_factory=SimpleQueue) + scripts: dict[str, deque[int]] = field(default_factory=dict) + + async def chat(self, request: Request) -> Response: + body: Final = JSON_OBJECT.validate_json(await request.body()) + self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body)) + leaked: Final = tuple(sorted(INTERNAL_FIELDS.intersection(body))) + if leaked: + return JSONResponse({"error": {"message": f"Unexpected provider fields: {leaked}"}}, status_code=400) + messages: Final = body.get("messages") + if not isinstance(body.get("model"), str) or not isinstance(messages, list) or not messages: + return JSONResponse({"error": {"message": "model and nonempty messages are required"}}, status_code=400) + if any( + not isinstance(message, dict) + or message.get("role") not in {"system", "developer", "user", "assistant", "tool"} + or "content" not in message + for message in messages + ): + return JSONResponse({"error": {"message": "Invalid selected message contract"}}, status_code=400) + script: Final = self.scripts.get(str(body["model"])) + if script is not None: + if not script: + return JSONResponse({"error": {"message": "Script exhausted", "type": "api_error"}}, status_code=500) + status: Final = script.popleft() + if status != 200: + return JSONResponse( + {"error": {"message": "Controlled provider failure", "type": "api_error", "code": str(status)}}, + status_code=status, + ) + return await chat_completions(request) + + async def script(self, request: Request) -> Response: + name: Final = request.path_params["model"] + if request.method in {"DELETE", "GET"} and name not in self.scripts: + return JSONResponse({"error": "Script not found"}, status_code=404) + if request.method == "GET": + return JSONResponse({"remaining": list(self.scripts[name])}) + if request.method == "DELETE": + remaining: Final = self.scripts.pop(name) + return JSONResponse({"remaining": list(remaining)}) + body: Final = JSON_OBJECT.validate_json(await request.body()) + statuses: Final = body.get("statuses") + if not isinstance(statuses, list) or not statuses or any(type(value) is not int for value in statuses): + return JSONResponse({"error": "A nonempty list of HTTP status codes is required"}, status_code=400) + self.scripts[name] = deque(int(str(value)) for value in statuses) + return JSONResponse({"configured": len(statuses)}) + + async def observed(self, _request: Request) -> Response: + values: Final = tuple(self.observations.get() for _ in range(self.observations.qsize())) + return JSONResponse( + { + "requests": [ + {"path": value.path, "authorization": value.authorization, "body": value.body} for value in values + ] + } + ) + + def app(self) -> Starlette: + return Starlette( + routes=[ + Route("/health", health), + Route("/__observations", self.observed), + Route("/__scripts/{model}", self.script, methods=["POST", "DELETE", "GET"]), + Route("/v1/chat/completions", self.chat, methods=["POST"]), + Route("/v1/completions", completions, methods=["POST"]), + Route("/v1/embeddings", embeddings, methods=["POST"]), + Route("/v1/moderations", moderations, methods=["POST"]), + ] + ) + + +def main() -> None: + parser: Final = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=8190) + arguments: Final = parser.parse_args() + uvicorn.run(Provider().app(), host="127.0.0.1", port=arguments.port, access_log=False) + + +if __name__ == "__main__": + main() diff --git a/tests/integration/authorization/test_warmed_policy.py b/tests/integration/authorization/test_warmed_policy.py new file mode 100644 index 00000000000..fd4271dbc41 --- /dev/null +++ b/tests/integration/authorization/test_warmed_policy.py @@ -0,0 +1,197 @@ +from contextlib import ExitStack +from hashlib import sha256 +from typing import Final +import os + +import psycopg +import pytest +from hypothesis import strategies as st +from hypothesis.stateful import RuleBasedStateMachine, invariant, rule, run_state_machine_as_test + +from integration._support.client import Gateway, eventually, object_value +from integration._support.database import read_rows +from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests + + +def assert_serving(gateway: Gateway, model: str, key: str, status: int, error_type: str = "auth_error") -> None: + response: Final = eventually( + lambda: gateway.request( + "POST", "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "warmed policy control"}]}, key=key, + ), + lambda value: value.status_code == status, + seconds=3, + ) + if status == 200: + assert response.json()["usage"]["total_tokens"] == 40 + assert response.json()["choices"][0]["message"]["content"] == ( + "Hello! This is a mock response from the fake OpenAI endpoint." + ) + else: + assert response.json()["error"]["type"] == error_type + + +@pytest.mark.covers("mgmt.key.update.two_workers_enforce_warmed_policy") +def test_generated_policy_changes_reach_both_warmed_workers(gateway: Gateway, peer: Gateway) -> None: + class Policies(RuleBasedStateMachine): + def __init__(self) -> None: + super().__init__() + self.resources = ExitStack() + try: + scenario = self.resources.enter_context(gateway.scenario()) + self.models = (scenario.model(), scenario.model()) + self.allowed = 0 + self.blocked = False + self.key = scenario.key(models=[self.models[0]], blocked=False) + self.control = scenario.key(models=list(self.models)) + for worker in (gateway, peer): + assert_serving(worker, self.models[0], self.key, 200) + assert_serving(worker, self.models[1], self.control, 200) + except BaseException: + with budget.cleanup(): + self.resources.close() + raise + + @rule(index=st.integers(min_value=0, max_value=1)) + def model_grant(self, index: int) -> None: + gateway.post("/key/update", {"key": self.key, "models": [self.models[index]]}) + self.allowed = index + + @rule(blocked=st.booleans()) + def block(self, blocked: bool) -> None: + gateway.post("/key/update", {"key": self.key, "blocked": blocked}) + self.blocked = blocked + + @invariant() + def both_workers_enforce_policy(self) -> None: + rows: Final = read_rows( + 'SELECT models, blocked FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(self.key.encode()).hexdigest(),), + ) + assert rows == [{"models": [self.models[self.allowed]], "blocked": self.blocked}] + for worker in (gateway, peer): + for index, model in enumerate(self.models): + status: Final = 401 if self.blocked else 200 if index == self.allowed else 403 + kind: Final = "auth_error" if self.blocked else "key_model_access_denied" + assert_serving(worker, model, self.key, status, kind) + assert_serving(worker, self.models[1], self.control, 200) + + def teardown(self) -> None: + with budget.cleanup(): + self.resources.close() + + with bounded_http_requests((gateway, peer), limit=3000) as budget: + run_state_machine_as_test(Policies, settings=LIFECYCLE_SETTINGS) + + +@pytest.mark.covers("mgmt.user.scim.deactivation_includes_nullable_blocked_keys") +def test_scim_deactivation_blocks_null_and_false_keys_but_preserves_other_owners(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + other: Final = scenario.user(user_role="internal_user") + null_key: Final = scenario.key(user_id=user, models=[model]) + false_key: Final = scenario.key(user_id=user, models=[model], blocked=False) + manual: Final = scenario.key(user_id=user, models=[model], blocked=True) + control: Final = scenario.key(user_id=other, models=[model]) + team: Final = scenario.team(models=[model]) + service: Final = gateway.post("/key/service-account/generate", {"team_id": team, "models": [model]}) + service_key: Final = service["key"] + assert isinstance(service_key, str) + scenario.cleanups.callback(scenario.delete_key, service_key) + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + 'UPDATE "LiteLLM_VerificationToken" SET blocked = NULL WHERE token = %s', + (sha256(null_key.encode()).hexdigest(),), + ) + assert read_rows( + 'SELECT user_id, blocked FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(null_key.encode()).hexdigest(),), + ) == [{"user_id": user, "blocked": None}] + assert read_rows( + 'SELECT user_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(service_key.encode()).hexdigest(),), + ) == [{"user_id": None}] + for token in (null_key, false_key, control, service_key): + assert_serving(gateway, model, token, 200) + for active in (False, True): + response: Final = gateway.request( + "PATCH", f"/scim/v2/Users/{user}", + {"schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "path": "active", "value": active}]}, + ) + assert response.status_code == 200, response.text + for token in (null_key, false_key): + rows: Final = read_rows( + 'SELECT blocked, metadata FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(token.encode()).hexdigest(),), + ) + assert rows[0]["blocked"] is not active + assert object_value(rows[0]["metadata"]).get("scim_blocked") is (None if active else True) + assert_serving(gateway, model, token, 200 if active else 401) + assert_serving(gateway, model, manual, 401) + for token in (control, service_key): + assert_serving(gateway, model, token, 200) + + +@pytest.mark.covers("mgmt.team.member_update.demoted_role_cannot_write") +def test_warmed_team_role_demotion_prevents_later_management_writes(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"user_id": user, "role": "admin"}]) + control_team: Final = scenario.team(models=[model]) + caller: Final = scenario.key( + user_id=user, team_id=team, models=[model], allowed_routes=["/team/update", "/v1/chat/completions"] + ) + gateway.chat(model, key=caller) + changed: Final = gateway.request("POST", "/team/update", {"team_id": team, "team_alias": "permitted"}, key=caller) + assert changed.status_code == 200, changed.text + unrelated_before: Final = read_rows( + 'SELECT team_alias FROM "LiteLLM_TeamTable" WHERE team_id = %s', (control_team,) + ) + unrelated: Final = gateway.request( + "POST", "/team/update", {"team_id": control_team, "team_alias": "must-not-persist"}, key=caller + ) + assert unrelated.status_code == 403, unrelated.text + assert read_rows( + 'SELECT team_alias FROM "LiteLLM_TeamTable" WHERE team_id = %s', (control_team,) + ) == unrelated_before + gateway.post("/team/member_update", {"team_id": team, "user_id": user, "role": "user"}) + for target in (team, control_team): + before: Final = read_rows('SELECT team_alias FROM "LiteLLM_TeamTable" WHERE team_id = %s', (target,)) + denied: Final = gateway.request( + "POST", "/team/update", {"team_id": target, "team_alias": "must-not-persist"}, key=caller + ) + assert denied.status_code == 403, denied.text + assert read_rows('SELECT team_alias FROM "LiteLLM_TeamTable" WHERE team_id = %s', (target,)) == before + roster: Final = read_rows('SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)) + members: Final = roster[0]["members_with_roles"] + assert isinstance(members, list) + assert next(object_value(member)["role"] for member in members if object_value(member)["user_id"] == user) == "user" + assert_serving(gateway, model, caller, 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: + model: Final = scenario.model() + key: Final = scenario.key(models=[model], duration="1h") + control: Final = scenario.key(models=[model]) + for worker in (gateway, peer): + assert_serving(worker, model, key, 200) + gateway.post("/key/update", {"key": key, "duration": "0s"}) + assert read_rows( + "SELECT expires <= timezone('UTC', now()) AS expired FROM \"LiteLLM_VerificationToken\" WHERE token = %s", + (sha256(key.encode()).hexdigest(),), + ) == [{"expired": True}] + for worker in (gateway, peer): + assert_serving(worker, model, key, 401, "expired_key") + assert_serving(worker, model, control, 200) + gateway.post("/key/update", {"key": key, "duration": None}) + assert read_rows( + 'SELECT expires IS NULL AS cleared FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) == [{"cleared": True}] + for worker in (gateway, peer): + assert_serving(worker, model, key, 200) diff --git a/tests/integration/configuration/test_effective_settings.py b/tests/integration/configuration/test_effective_settings.py new file mode 100644 index 00000000000..7fa440d1d8d --- /dev/null +++ b/tests/integration/configuration/test_effective_settings.py @@ -0,0 +1,123 @@ +import uuid +from typing import Final + +import httpx +import pytest + +from integration._support.client import Gateway, object_value, string_value +from integration._support.database import read_rows + + +def model_identity(gateway: Gateway, alias: str) -> str: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + entry: Final = next(object_value(value) for value in entries if object_value(value)["model_name"] == alias) + return string_value(object_value(entry["model_info"])["id"]) + + +@pytest.mark.covers("mgmt.model.block.changes_serving_and_preserves_control") +def test_model_block_changes_actual_route_and_leaves_other_route_working(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + other: Final = scenario.model() + identity: Final = model_identity(gateway, model) + gateway.chat(model) + gateway.chat(other) + gateway.post("/model/block", {"model_id": identity}) + assert read_rows( + 'SELECT blocked FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,) + ) == [{"blocked": True}] + response: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "blocked deployment"}]}, + ) + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "permission_error" + assert response.json()["error"]["message"] == "litellm.PermissionDeniedError: Model is blocked" + assert object_value(gateway.chat(other)["usage"])["total_tokens"] == 40 + gateway.post("/model/unblock", {"model_id": identity}) + assert read_rows( + 'SELECT blocked FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,) + ) == [{"blocked": False}] + assert object_value(gateway.chat(model)["usage"])["total_tokens"] == 40 + + +@pytest.mark.covers("mgmt.router_settings.update.changes_observed_attempt_count") +def test_saved_retry_setting_controls_real_attempts_and_restores(gateway: Gateway) -> None: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream, gateway.scenario() as scenario: + original: Final = object_value(gateway.get("/router/settings")["current_values"])["num_retries"] + provider_model: Final = f"retry-{uuid.uuid4().hex}" + model: Final = scenario.model(model=f"openai/{provider_model}", input_cost_per_token=0, output_cost_per_token=0) + + def remove_script() -> None: + response: Final = upstream.delete(f"/__scripts/{provider_model}") + assert response.status_code in (200, 404), response.text + assert upstream.get(f"/__scripts/{provider_model}").status_code == 404 + + scenario.cleanups.callback(remove_script) + try: + for generation, retries in enumerate((0, 1, original)): + gateway.post("/config/update", {"router_settings": {"num_retries": retries}}) + assert object_value(gateway.get("/router/settings")["current_values"])["num_retries"] == retries + configured: Final = upstream.post(f"/__scripts/{provider_model}", json={"statuses": [500, 200]}) + assert configured.status_code == 200, configured.text + upstream.get("/__observations").raise_for_status() + response: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"{provider_model} attempt {generation}"}]}, + ) + observed: Final = upstream.get("/__observations") + observed.raise_for_status() + requests: Final = observed.json()["requests"] + assert len(requests) == (1 if retries == 0 else 2), (response.status_code, response.text, requests) + assert all(value["body"]["model"] == provider_model for value in requests) + assert response.status_code == (500 if retries == 0 else 200), response.text + if retries != 0: + assert response.json()["usage"]["total_tokens"] == 40 + remaining: Final = upstream.delete(f"/__scripts/{provider_model}") + assert remaining.status_code == 200, remaining.text + assert remaining.json()["remaining"] == ([200] if retries == 0 else []) + finally: + gateway.post("/config/update", {"router_settings": {"num_retries": original}}) + assert object_value(gateway.get("/router/settings")["current_values"])["num_retries"] == original + + +@pytest.mark.covers("mgmt.credential.update.saved_value_reaches_wire") +def test_credential_value_update_and_model_reload_reach_provider(gateway: Gateway) -> None: + with gateway.scenario() as scenario, httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + name: Final = f"credential-{uuid.uuid4().hex}" + gateway.post("/credentials", { + "credential_name": name, "credential_values": {"api_key": "synthetic-credential-first"}, "credential_info": {} + }) + + def remove_credential() -> None: + response: Final = gateway.request("DELETE", f"/credentials/{name}") + assert response.status_code == 200, response.text + assert read_rows( + 'SELECT credential_name FROM "LiteLLM_CredentialsTable" WHERE credential_name = %s', (name,) + ) == [] + + scenario.cleanups.callback(remove_credential) + model: Final = scenario.model(api_key=None, litellm_credential_name=name) + identity: Final = model_identity(gateway, model) + for value in ("synthetic-credential-first", "synthetic-credential-second"): + patched: Final = gateway.request("PATCH", f"/credentials/{name}", { + "credential_name": name, "credential_values": {"api_key": value}, "credential_info": {} + }) + assert patched.status_code == 200, patched.text + rows: Final = read_rows( + 'SELECT credential_values FROM "LiteLLM_CredentialsTable" WHERE credential_name = %s', (name,) + ) + assert len(rows) == 1 + stored: Final = object_value(rows[0]["credential_values"]) + assert isinstance(stored["api_key"], str) and stored["api_key"] != value + for reload in (False, True): + if reload: + response: Final = gateway.request("PATCH", f"/model/{identity}/update", {"model_info": {"description": value}}) + assert response.status_code == 200, response.text + upstream.get("/__observations").raise_for_status() + assert object_value(gateway.chat(model, text=f"{name} {value} reload={reload}")["usage"])["total_tokens"] == 40 + observed: Final = upstream.get("/__observations") + observed.raise_for_status() + assert len(observed.json()["requests"]) == 1, (value, reload, observed.text) + assert observed.json()["requests"][0]["authorization"] == f"Bearer {value}" diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py new file mode 100644 index 00000000000..f5a018d305a --- /dev/null +++ b/tests/integration/conftest.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +import json +import os +from importlib.metadata import version +from collections.abc import Generator, Iterator +from pathlib import Path +from typing import Final + +import pytest +import httpx +from redis import Redis + +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.manifest import OWNED_DIRECTORIES, contracts +from integration._support.generation import LIFECYCLE_SETTINGS + +COLLECTED: Final = pytest.StashKey[tuple[str, ...]]() +REPORTS: Final = pytest.StashKey[list[pytest.TestReport]]() + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line("markers", "integration: owned real-service integration contracts") + config.addinivalue_line("markers", "covers(*ids): independently asserted behavior contracts") + config.stash[REPORTS] = [] + + +def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None: + manifest: Final = contracts() + root: Final = Path(__file__).parent + owned: Final = tuple( + item + for item in items + if item.path.is_relative_to(root) and item.path.relative_to(root).parts[0] in OWNED_DIRECTORIES + ) + if owned and os.environ.get("GITHUB_ACTIONS") == "true": + raise pytest.UsageError("Integration contracts are owned by CircleCI") + for item in owned: + if item.nodeid not in manifest: + raise pytest.UsageError(f"Integration node missing from manifest: {item.nodeid}") + item.add_marker(pytest.mark.integration) + declared: Final = tuple(value for mark in item.iter_markers("covers") for value in mark.args) + if set(declared) != set(manifest[item.nodeid]): + raise pytest.UsageError(f"Contract mapping differs for {item.nodeid}") + config.stash[COLLECTED] = tuple(item.nodeid for item in owned) + + +@pytest.hookimpl(wrapper=True) +def pytest_runtest_makereport( + item: pytest.Item, call: pytest.CallInfo[None] +) -> Generator[None, pytest.TestReport, pytest.TestReport]: + report: Final = yield + item.config.stash[REPORTS].append(report) + return report + + +def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: + destination: Final = os.environ.get("INTEGRATION_RESULTS_DIR") + if destination is None: + return + collected: Final = session.config.stash.get(COLLECTED, ()) + reports: Final = tuple(report for report in session.config.stash[REPORTS] if report.nodeid in collected) + passed: Final = tuple(report.nodeid for report in reports if report.when == "call" and report.passed) + complete: Final = ( + exitstatus == 0 + and bool(collected) + and sorted(collected) == sorted(passed) + and all(report.passed for report in reports) + ) + output: Final = Path(destination) + output.mkdir(parents=True, exist_ok=True) + (output / "execution.json").write_text( + json.dumps({ + "collected": collected, "passed": passed, "complete": complete, "exitstatus": exitstatus, + "hypothesis_version": version("hypothesis"), + "hypothesis_seed": session.config.getoption("hypothesis_seed"), + "generation": { + "max_examples": LIFECYCLE_SETTINGS.max_examples, + "stateful_step_count": LIFECYCLE_SETTINGS.stateful_step_count, + "database": str(LIFECYCLE_SETTINGS.database), + "phases": [phase.name for phase in LIFECYCLE_SETTINGS.phases], + }, + }, indent=2) + + "\n" + ) + if not complete and exitstatus == 0: + session.exitstatus = pytest.ExitCode.TESTS_FAILED + + +@pytest.fixture +def gateway() -> Iterator[Gateway]: + with gateway_from_environment() as value: + yield value + + +@pytest.fixture +def peer(gateway: Gateway) -> Iterator[Gateway]: + url: Final = os.environ["INTEGRATION_PEER_URL"] + assert url.rstrip("/") != str(gateway.client.base_url).rstrip("/") + with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache: + eventually(lambda: cache.pubsub_numsub("litellm_proxy.auth_cache_invalidation")[0][1], lambda count: count >= 2) + with httpx.Client(base_url=url, timeout=15, trust_env=False) as client: + yield Gateway(client, gateway.key, gateway.upstream_url) diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json new file mode 100644 index 00000000000..82cc64dd5c6 --- /dev/null +++ b/tests/integration/contracts.json @@ -0,0 +1,80 @@ +{ + "groups": { + "management": [ + "management", + "authorization", + "configuration" + ], + "accounting": [ + "pricing", + "spend" + ], + "database": [ + "database" + ], + "providers": [ + "providers", + "routing", + "streaming" + ], + "extensions": [ + "mcp", + "observability", + "compatibility" + ] + }, + "tests": { + "tests/integration/management/test_key_updates.py::test_update_preserves_independent_fields_and_serving": [ + "mgmt.key.update.preserves_independent_fields" + ], + "tests/integration/pricing/test_configured_prices.py::test_custom_price_is_reported_and_charged": [ + "quota_management.spend_tracking.custom_price.matches_input_rates" + ], + "tests/integration/providers/test_request_boundary.py::test_internal_request_state_does_not_reach_provider": [ + "other.provider_wire.internal_parameters_filtered" + ], + "tests/integration/pricing/test_configured_prices.py::test_default_prices_survive_nullable_sibling_and_reload": [ + "quota_management.spend_tracking.default_prices.survive_nullable_sibling_reload" + ], + "tests/integration/providers/test_request_boundary.py::test_upstream_rejects_corruption_and_accepts_supported_metadata": [ + "other.provider_wire.validator_rejects_corruption" + ], + "tests/integration/pricing/test_configured_prices.py::test_loaded_router_preserves_cached_defaults_during_real_requests": [ + "quota_management.spend_tracking.default_prices.loaded_router_preserves_cached_defaults" + ], + "tests/integration/management/test_partial_update_sequences.py::test_generated_partial_updates_preserve_persisted_and_effective_state": [ + "mgmt.key.update.generated_sequences_preserve_state" + ], + "tests/integration/management/test_partial_update_sequences.py::test_zero_false_and_empty_values_are_not_treated_as_omission": [ + "mgmt.key.update.false_zero_and_empty_values_affect_serving" + ], + "tests/integration/management/test_partial_update_sequences.py::test_project_omission_clear_and_invalid_update_have_distinct_effects": [ + "mgmt.key.update.project_clear_preserves_scope", + "mgmt.key.update.invalid_batch_is_atomic" + ], + "tests/integration/authorization/test_warmed_policy.py::test_generated_policy_changes_reach_both_warmed_workers": [ + "mgmt.key.update.two_workers_enforce_warmed_policy" + ], + "tests/integration/authorization/test_warmed_policy.py::test_scim_deactivation_blocks_null_and_false_keys_but_preserves_other_owners": [ + "mgmt.user.scim.deactivation_includes_nullable_blocked_keys" + ], + "tests/integration/authorization/test_warmed_policy.py::test_warmed_team_role_demotion_prevents_later_management_writes": [ + "mgmt.team.member_update.demoted_role_cannot_write" + ], + "tests/integration/configuration/test_effective_settings.py::test_model_block_changes_actual_route_and_leaves_other_route_working": [ + "mgmt.model.block.changes_serving_and_preserves_control" + ], + "tests/integration/configuration/test_effective_settings.py::test_saved_retry_setting_controls_real_attempts_and_restores": [ + "mgmt.router_settings.update.changes_observed_attempt_count" + ], + "tests/integration/configuration/test_effective_settings.py::test_credential_value_update_and_model_reload_reach_provider": [ + "mgmt.credential.update.saved_value_reaches_wire" + ], + "tests/integration/management/test_partial_update_sequences.py::test_denied_key_update_preserves_saved_grants_and_serving": [ + "mgmt.key.update.denied_request_preserves_effective_state" + ], + "tests/integration/authorization/test_warmed_policy.py::test_expiry_and_explicit_clear_reach_both_warmed_workers": [ + "mgmt.key.update.expiry_changes_reach_warmed_workers" + ] + } +} diff --git a/tests/integration/management/test_key_updates.py b/tests/integration/management/test_key_updates.py new file mode 100644 index 00000000000..6f2e850b17a --- /dev/null +++ b/tests/integration/management/test_key_updates.py @@ -0,0 +1,40 @@ +from typing import Final +from hashlib import sha256 + +import pytest + +from integration._support.client import Gateway, object_value +from integration._support.database import read_rows + + +@pytest.mark.covers("mgmt.key.update.preserves_independent_fields") +def test_update_preserves_independent_fields_and_serving(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=[model], key_alias="before", metadata={"retained": "value"}) + gateway.chat(model, key=key) + gateway.post("/key/update", {"key": key, "key_alias": "after"}) + info: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) + assert info["key_alias"] == "after" + assert info["models"] == [model] + assert object_value(info["metadata"])["retained"] == "value" + response: Final = gateway.chat(model, key=key) + assert object_value(response["usage"])["total_tokens"] == 40 + replacement: Final = scenario.model() + gateway.post("/key/update", {"key": key, "models": [replacement]}) + saved: Final = read_rows( + 'SELECT key_alias, models, metadata FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + assert len(saved) == 1 + assert saved[0]["models"] == [replacement] + assert saved[0]["key_alias"] == "after" + denied: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "old grant"}]}, + key=key, + ) + assert denied.status_code == 403, denied.text + assert object_value(object_value(denied.json())["error"])["type"] == "key_model_access_denied" + assert object_value(gateway.chat(replacement, key=key)["usage"])["total_tokens"] == 40 diff --git a/tests/integration/management/test_partial_update_sequences.py b/tests/integration/management/test_partial_update_sequences.py new file mode 100644 index 00000000000..c645b896448 --- /dev/null +++ b/tests/integration/management/test_partial_update_sequences.py @@ -0,0 +1,200 @@ +from contextlib import ExitStack +from hashlib import sha256 +from typing import Final + +import pytest +from hypothesis import strategies as st +from hypothesis.stateful import RuleBasedStateMachine, invariant, rule, run_state_machine_as_test +from pydantic import JsonValue + +from integration._support.client import Gateway, object_value +from integration._support.database import read_rows +from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests + + +@pytest.mark.covers("mgmt.key.update.generated_sequences_preserve_state") +def test_generated_partial_updates_preserve_persisted_and_effective_state(gateway: Gateway) -> None: + class KeyUpdates(RuleBasedStateMachine): + def __init__(self) -> None: + super().__init__() + self.resources = ExitStack() + try: + scenario = self.resources.enter_context(gateway.scenario()) + self.models = (scenario.model(), scenario.model()) + self.key = scenario.key(models=[self.models[0]], key_alias="initial", metadata={"revision": "initial"}) + self.expected: dict[str, JsonValue] = { + "models": [self.models[0]], "key_alias": "initial", "metadata": {"revision": "initial"} + } + gateway.chat(self.models[0], key=self.key) + except BaseException: + with budget.cleanup(): + self.resources.close() + raise + + @rule(alias=st.sampled_from(("first", "second", "", "unicode-λ"))) + def alias(self, alias: str) -> None: + gateway.post("/key/update", {"key": self.key, "key_alias": alias}) + self.expected["key_alias"] = alias + + @rule(index=st.integers(min_value=0, max_value=1), both=st.booleans()) + def grant(self, index: int, both: bool) -> None: + models: Final = list(self.models) if both else [self.models[index]] + gateway.post("/key/update", {"key": self.key, "models": models}) + self.expected["models"] = models + + @rule(value=st.sampled_from(("", "a", "different", "λ"))) + def metadata(self, value: str) -> None: + gateway.post("/key/update", {"key": self.key, "metadata": {"revision": value}}) + self.expected["metadata"] = {"revision": value} + + @invariant() + def persisted_state_and_serving_match(self) -> None: + rows: Final = read_rows( + 'SELECT models, key_alias, metadata FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(self.key.encode()).hexdigest(),), + ) + assert rows == [self.expected] + info: Final = object_value(gateway.get("/key/info", {"key": self.key})["info"]) + assert {field: info[field] for field in self.expected} == self.expected + for model in self.models: + response: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "generated update control"}]}, + key=self.key, + ) + if model in self.expected["models"]: + assert response.status_code == 200, response.text + assert response.json()["usage"]["total_tokens"] == 40 + else: + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "key_model_access_denied" + + def teardown(self) -> None: + with budget.cleanup(): + self.resources.close() + + with bounded_http_requests((gateway,), limit=2000) as budget: + run_state_machine_as_test(KeyUpdates, settings=LIFECYCLE_SETTINGS) + + +@pytest.mark.covers("mgmt.key.update.false_zero_and_empty_values_affect_serving") +def test_zero_false_and_empty_values_are_not_treated_as_omission(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + models: Final = (scenario.model(), scenario.model()) + key: Final = scenario.key(models=[models[0]], max_budget=0, metadata={"ordinary": "value"}) + denied: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": models[0], "messages": [{"role": "user", "content": "zero budget"}]}, key=key, + ) + assert denied.status_code == 429, denied.text + assert denied.json()["error"]["type"] == "budget_exceeded" + gateway.post("/key/update", {"key": key, "max_budget": 1, "models": [], "metadata": {}}) + info: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) + assert (info["models"], info["metadata"], info["max_budget"]) == ([], {}, 1) + for model in models: + assert object_value(gateway.chat(model, key=key)["usage"])["total_tokens"] == 40 + gateway.post("/key/update", {"key": key, "blocked": True}) + blocked: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": models[0], "messages": [{"role": "user", "content": "blocked control"}]}, key=key, + ) + assert blocked.status_code == 401, blocked.text + assert blocked.json()["error"]["type"] == "auth_error" + gateway.post("/key/update", {"key": key, "blocked": False}) + assert object_value(gateway.chat(models[0], key=key)["usage"])["total_tokens"] == 40 + rows: Final = read_rows( + 'SELECT blocked, models, metadata, max_budget FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + assert rows == [{"blocked": False, "models": [], "metadata": {}, "max_budget": 1.0}] + gateway.post("/key/update", {"key": key, "max_budget": 0}) + assert read_rows( + 'SELECT max_budget FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) == [{"max_budget": 0.0}] + zero_after_update: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": models[0], "messages": [{"role": "user", "content": "updated zero budget"}]}, key=key, + ) + assert zero_after_update.status_code == 429, zero_after_update.text + assert zero_after_update.json()["error"]["type"] == "budget_exceeded" + gateway.post("/key/update", {"key": key, "max_budget": None}) + assert read_rows( + 'SELECT max_budget FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) == [{"max_budget": None}] + assert object_value(gateway.chat(models[0], key=key)["usage"])["total_tokens"] == 40 + + +@pytest.mark.covers("mgmt.key.update.project_clear_preserves_scope", "mgmt.key.update.invalid_batch_is_atomic") +def test_project_omission_clear_and_invalid_update_have_distinct_effects(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + outside: Final = scenario.model() + team: Final = scenario.team(models=[model]) + project: Final = scenario.project(team, models=[model]) + other: Final = scenario.project(team, models=[model]) + key: Final = scenario.key(team_id=team, project_id=project, models=[model], key_alias="before", max_budget=5) + gateway.chat(model, key=key) + gateway.post("/key/update", {"key": key, "key_alias": "after"}) + digest: Final = sha256(key.encode()).hexdigest() + + def saved() -> list[dict[str, object]]: + return read_rows( + 'SELECT key_alias, project_id, team_id, models, max_budget FROM "LiteLLM_VerificationToken" ' + 'WHERE token = %s', (digest,), + ) + + before: Final = saved() + assert before == [{"key_alias": "after", "project_id": project, "team_id": team, "models": [model], "max_budget": 5}] + for invalid in (other, ""): + rejected: Final = gateway.request( + "POST", "/key/update", {"key": key, "project_id": invalid, "key_alias": "must-not-persist"} + ) + assert rejected.status_code == 400, rejected.text + assert saved() == before + gateway.chat(model, key=key) + gateway.post("/project/update", {"project_id": project, "blocked": True}) + denied: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "blocked project"}]}, + key=key, + ) + assert denied.status_code == 401, denied.text + assert denied.json()["error"]["type"] == "auth_error" + for _ in range(2): + gateway.post("/key/update", {"key": key, "project_id": None}) + assert saved() == [{**before[0], "project_id": None}] + assert object_value(gateway.chat(model, key=key)["usage"])["total_tokens"] == 40 + outside_request: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": outside, "messages": [{"role": "user", "content": "detached scope control"}]}, key=key, + ) + assert outside_request.status_code == 403, outside_request.text + assert outside_request.json()["error"]["type"] == "key_model_access_denied" + + +@pytest.mark.covers("mgmt.key.update.denied_request_preserves_effective_state") +def test_denied_key_update_preserves_saved_grants_and_serving(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + outside: Final = scenario.model() + owner: Final = scenario.user(user_role="internal_user") + other: Final = scenario.user(user_role="internal_user") + key: Final = scenario.key(user_id=owner, models=[model], key_alias="unchanged", max_budget=2) + caller: Final = scenario.key(user_id=other, models=[model], allowed_routes=["/key/update", "/v1/chat/completions"]) + gateway.chat(model, key=key) + denied: Final = gateway.request( + "POST", "/key/update", {"key": key, "key_alias": "wrong", "models": [outside], "max_budget": 0}, key=caller + ) + assert denied.status_code == 403, denied.text + assert read_rows( + 'SELECT user_id, models, key_alias, max_budget FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) == [{"user_id": owner, "models": [model], "key_alias": "unchanged", "max_budget": 2.0}] + assert object_value(gateway.chat(model, key=key)["usage"])["total_tokens"] == 40 + rejected: Final = gateway.request( + "POST", "/v1/chat/completions", + {"model": outside, "messages": [{"role": "user", "content": "unchanged scope"}]}, key=key, + ) + assert rejected.status_code == 403, rejected.text + assert rejected.json()["error"]["type"] == "key_model_access_denied" diff --git a/tests/integration/pricing/test_configured_prices.py b/tests/integration/pricing/test_configured_prices.py new file mode 100644 index 00000000000..151103f6df5 --- /dev/null +++ b/tests/integration/pricing/test_configured_prices.py @@ -0,0 +1,148 @@ +from collections.abc import Iterator, Mapping +from typing import Final +from pathlib import Path +import uuid + +import pytest +import yaml + +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows + + +@pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates") +def test_custom_price_is_reported_and_charged(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "price control"}]} + ) + assert response.status_code == 200, response.text + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(20 * 0.001 + 20 * 0.002) + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + matching: Final = tuple(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model) + assert len(matching) == 1 + params: Final = object_value(matching[0]["litellm_params"]) + assert params["input_cost_per_token"] == 0.001 + assert params["output_cost_per_token"] == 0.002 + + +@pytest.mark.covers("quota_management.spend_tracking.default_prices.survive_nullable_sibling_reload") +def test_default_prices_survive_nullable_sibling_and_reload(gateway: Gateway) -> None: + for registration_order in (("custom", "omitted", "nullable"), ("nullable", "omitted", "custom")): + with gateway.scenario() as scenario: + configured: Final = { + "custom": {"input_cost_per_token": 0.001, "output_cost_per_token": 0.002}, + "omitted": {}, + "nullable": {"input_cost_per_token": None, "output_cost_per_token": None}, + } + rates: Final = { + "custom": (0.001, 0.002), + "omitted": (0.00000015, 0.0000006), + "nullable": (0.00000015, 0.0000006), + } + models: Final = {kind: scenario.model(**configured[kind]) for kind in registration_order} + + def observe_requests( + registration_order: tuple[str, ...], + models: Mapping[str, str], + rates: Mapping[str, tuple[float, float]], + ) -> Iterator[tuple[str, float]]: + for generation in range(2): + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + kinds: Final = tuple(reversed(registration_order)) if generation else registration_order + for index, kind in enumerate(kinds): + model: Final = models[kind] + target: Final = next( + object_value(entry) for entry in entries if object_value(entry)["model_name"] == model + ) + info: Final = object_value(target["model_info"]) + assert info["input_cost_per_token"] == rates[kind][0] + assert info["output_cost_per_token"] == rates[kind][1] + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [ + {"role": "user", "content": f"price {generation * len(registration_order) + index}"} + ], + }, + ) + assert response.status_code == 200, response.text + expected: Final = 20 * rates[kind][0] + 20 * rates[kind][1] + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6) + request_id: Final = string_value(object_value(response.json())["id"]) + yield request_id, expected + target: Final = next( + object_value(entry) + for entry in entries + if object_value(entry)["model_name"] == models["nullable"] + ) + identity: Final = string_value(object_value(target["model_info"])["id"]) + updated: Final = gateway.request( + "PATCH", f"/model/{identity}/update", {"model_info": {"description": "reload price contract"}} + ) + assert updated.status_code == 200, updated.text + + observations: Final = tuple(observe_requests(registration_order, models, rates)) + for request_id, expected in observations: + rows: Final = eventually( + lambda request_id=request_id: read_rows( + 'SELECT request_id, spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" ' + "WHERE request_id = %s", + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["prompt_tokens"] == 20 + assert rows[0]["completion_tokens"] == 20 + assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6) + + +@pytest.mark.covers("quota_management.spend_tracking.default_prices.loaded_router_preserves_cached_defaults") +def test_loaded_router_preserves_cached_defaults_during_real_requests(gateway: Gateway, tmp_path: Path) -> None: + from litellm import Router + + aliases: Final = (f"pricing-{uuid.uuid4().hex}", f"pricing-{uuid.uuid4().hex}") + path: Final = tmp_path / "models.yaml" + path.write_text( + yaml.safe_dump( + { + "model_list": [ + { + "model_name": alias, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": f"{gateway.upstream_url}/v1", + }, + "model_info": {"id": alias, **pricing}, + } + for alias, pricing in zip( + aliases, ({}, {"input_cost_per_token": None, "output_cost_per_token": None}), strict=True + ) + ] + } + ) + ) + for reverse in (False, True): + configured: Final = yaml.safe_load(path.read_text())["model_list"] + router: Final = Router(model_list=list(reversed(configured)) if reverse else configured, num_retries=0) + try: + for alias in (*aliases, *reversed(aliases)): + result: Final = router.completion( + model=alias, messages=[{"role": "user", "content": "router price control"}] + ) + assert result.usage.prompt_tokens == 20 + assert result.usage.completion_tokens == 20 + deployment: Final = router.get_deployment(model_id=alias) + assert deployment is not None + info: Final = router.get_router_model_info(deployment=deployment, received_model_name=alias) + assert info["input_cost_per_token"] == 0.00000015 + assert info["output_cost_per_token"] == 0.0000006 + finally: + router.reset() diff --git a/tests/integration/providers/test_request_boundary.py b/tests/integration/providers/test_request_boundary.py new file mode 100644 index 00000000000..aad10843642 --- /dev/null +++ b/tests/integration/providers/test_request_boundary.py @@ -0,0 +1,57 @@ +from typing import Final + +import httpx +import pytest + +from integration._support.client import Gateway, JSON_OBJECT, object_value + + +@pytest.mark.covers("other.provider_wire.internal_parameters_filtered") +def test_internal_request_state_does_not_reach_provider(gateway: Gateway) -> None: + with gateway.scenario() as scenario, httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream: + upstream.get("/__observations").raise_for_status() + model: Final = scenario.model() + key: Final = scenario.key(models=[model], tpm_limit=10000, rpm_limit=100) + result: Final = gateway.post( + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "wire contract"}], + "temperature": 0.4, + "max_tokens": 20, + "timeout": 12, + }, + key=key, + ) + assert object_value(result["usage"])["total_tokens"] == 40 + observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"] + assert isinstance(observations, list) + assert len(observations) == 1 + observed: Final = object_value(observations[0]) + body: Final = object_value(observed["body"]) + assert body["model"] == "gpt-4o-mini" + assert body["messages"] == [{"role": "user", "content": "wire contract"}] + assert body["temperature"] == 0.4 + assert body["max_tokens"] == 20 + assert observed["authorization"] == "Bearer integration-provider-key" + assert "litellm_metadata" not in body + assert "litellm_params" not in body + assert "timeout" not in body + assert "tpm" not in body + + +@pytest.mark.covers("other.provider_wire.validator_rejects_corruption") +def test_upstream_rejects_corruption_and_accepts_supported_metadata(gateway: Gateway) -> None: + with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream: + missing: Final = upstream.post("/v1/chat/completions", json={"model": "gpt-4o-mini"}) + assert missing.status_code == 400 + assert ( + object_value(object_value(missing.json())["error"])["message"] == "model and nonempty messages are required" + ) + body: Final = {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "strict control"}]} + leaked: Final = upstream.post("/v1/chat/completions", json={**body, "litellm_metadata": {"hidden": "value"}}) + assert leaked.status_code == 400 + assert "litellm_metadata" in str(object_value(object_value(leaked.json())["error"])["message"]) + valid: Final = upstream.post("/v1/chat/completions", json={**body, "metadata": {"purpose": "synthetic"}}) + assert valid.status_code == 200, valid.text + assert object_value(object_value(valid.json())["usage"])["total_tokens"] == 40 diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml new file mode 100644 index 00000000000..b6f9767c210 --- /dev/null +++ b/tests/integration/proxy_config.yaml @@ -0,0 +1,17 @@ +model_list: [] +general_settings: + master_key: os.environ/LITELLM_MASTER_KEY + database_url: os.environ/DATABASE_URL + store_model_in_db: true + disable_spend_logs: false + proxy_batch_write_at: 1 +litellm_settings: + enable_redis_auth_cache: true + cache: true + cache_params: + type: redis + host: os.environ/REDIS_HOST + port: os.environ/REDIS_PORT +router_settings: + num_retries: 0 + disable_cooldowns: true diff --git a/tests/integration/run.py b/tests/integration/run.py new file mode 100644 index 00000000000..a48798475a2 --- /dev/null +++ b/tests/integration/run.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import argparse +import json +import os +import subprocess +import sys +from pathlib import Path +from types import MappingProxyType +from typing import Final + +GROUPS: Final = MappingProxyType(json.loads(Path(__file__).with_name("contracts.json").read_text())["groups"]) + + +def main() -> int: + parser: Final = argparse.ArgumentParser() + parser.add_argument("group", choices=tuple(GROUPS)) + parser.add_argument("--results", type=Path, default=Path("test-results/integration")) + parser.add_argument("--seed", type=int, default=int(os.environ.get("INTEGRATION_SEED", "4106601"))) + options: Final = parser.parse_args() + root: Final = Path(__file__).resolve().parents[2] + selected: Final = tuple( + str(path.relative_to(root)) + for folder in GROUPS[options.group] + for path in sorted((root / "tests/integration" / folder).glob("test_*.py")) + ) + if not selected: + parser.error(f"No integration contracts selected for {options.group}") + output: Final = options.results.resolve() + output.mkdir(parents=True, exist_ok=True) + manifest: Final = json.loads((root / "tests/integration/contracts.json").read_text())["tests"] + expected: Final = sorted(node for node in manifest if node.split("::", 1)[0] in selected) + if not expected or set(selected) != {node.split("::", 1)[0] for node in expected}: + parser.error("Every selected file must have canonical manifest nodes") + environment: Final = { + **os.environ, + "PYTHONPATH": os.pathsep.join((str(root), str(root / "tests"), str(root / "tests/e2e"))), + "INTEGRATION_RESULTS_DIR": str(output), + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + } + result: Final = subprocess.call( + [ + sys.executable, + "-m", + "pytest", + *selected, + "-vv", + "--strict-markers", + "-p", + "no:pytest-retry", + "-p", + "no:rerunfailures", + "--timeout=90", + "--durations=15", + f"--hypothesis-seed={options.seed}", + f"--junitxml={output / 'junit.xml'}", + ], + cwd=root, + env=environment, + ) + if result != 0: + return result + evidence: Final = json.loads((output / "execution.json").read_text()) + if not evidence["complete"] or sorted(evidence["passed"]) != expected or sorted(evidence["collected"]) != expected: + print("Executed integration nodes differ from the canonical manifest", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py index 05bb9113835..c7712d96969 100644 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ b/tests/llm_responses_api_testing/test_openai_responses_api.py @@ -1627,9 +1627,10 @@ async def test_openai_responses_api_token_limit_error(): Parsing the in-stream ErrorEvent must not raise "pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent". - The iterator now surfaces the event as litellm.APIError with status 400 - (invalid_request_error is a non-retriable client error, so no - MidStreamFallbackError wrapping) carrying the provider's message. + The iterator routes the event through litellm.exception_type, so it surfaces as + the typed 400 client error the non-streaming path raises (litellm.BadRequestError) + carrying the provider's message. invalid_request_error is a non-retriable client + error, so there is no MidStreamFallbackError wrapping. """ litellm._turn_on_debug() @@ -1644,7 +1645,7 @@ async def test_openai_responses_api_token_limit_error(): async for event in response: print(event) - with pytest.raises(litellm.APIError) as exc_info: + with pytest.raises(litellm.BadRequestError) as exc_info: await _drain() assert exc_info.value.status_code == 400 diff --git a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py index e2eb6d0b68b..8b3dc436b8f 100644 --- a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py +++ b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py @@ -6,12 +6,15 @@ including the logging handler, cost tracking, and WebSocket message processing. """ import json +from collections.abc import Sequence from datetime import datetime from unittest.mock import AsyncMock, Mock, patch, MagicMock from typing import Dict, List, Any, Optional import pytest import httpx +import litellm +from typing_extensions import NotRequired, ReadOnly, TypedDict # Add the parent directory to the system path @@ -22,10 +25,16 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.types.utils import LlmProviders +from litellm.types.utils import CostBreakdown, LlmProviders, Usage from litellm.proxy._types import UserAPIKeyAuth +class _LiveTurn(TypedDict): + prompt: ReadOnly[tuple[int, int]] + candidates: ReadOnly[tuple[int, int]] + candidate_audio_token_count_missing: NotRequired[ReadOnly[bool]] + + class TestVertexAILivePassthroughLoggingHandler: """Test the Vertex AI Live Passthrough Logging Handler""" @@ -39,6 +48,7 @@ class TestVertexAILivePassthroughLoggingHandler: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} + mock._response_cost_calculator.return_value = None return mock @pytest.fixture @@ -201,88 +211,490 @@ class TestVertexAILivePassthroughLoggingHandler: assert text_prompt["tokenCount"] == 10 assert audio_prompt["tokenCount"] == 10 - @patch( - "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info" - ) - def test_calculate_cost_basic(self, mock_get_model_info, handler): - """Test basic cost calculation""" - mock_get_model_info.return_value = { - "input_cost_per_token": 0.000001, - "output_cost_per_token": 0.000002, - } + def test_usage_carries_every_modality(self, handler): + """Regression: the Usage object reported only TEXT, so audio and image billed as nothing. + prompt_tokens must be the full count and the details must name each modality, + because the cost calculator prices audio and image from *_tokens_details. + """ usage_metadata = { - "promptTokenCount": 100, - "candidatesTokenCount": 50, - "totalTokenCount": 150, - } - - cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata) - - # The cost calculation may include additional factors, so we check it's reasonable - expected_min_cost = (100 * 0.000001) + (50 * 0.000002) - assert cost >= expected_min_cost - assert cost > 0 - - @patch( - "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info" - ) - def test_calculate_cost_with_audio(self, mock_get_model_info, handler): - """Test cost calculation with audio tokens""" - mock_get_model_info.return_value = { - "input_cost_per_token": 0.000001, - "output_cost_per_token": 0.000002, - "input_cost_per_audio_token": 0.0001, - "output_cost_per_audio_token": 0.0002, - } - - usage_metadata = { - "promptTokenCount": 100, - "candidatesTokenCount": 50, - "totalTokenCount": 150, + "promptTokenCount": 1300, + "candidatesTokenCount": 124, + "totalTokenCount": 1424, "promptTokensDetails": [ - {"modality": "TEXT", "tokenCount": 80}, - {"modality": "AUDIO", "tokenCount": 20}, + {"modality": "TEXT", "tokenCount": 13}, + {"modality": "AUDIO", "tokenCount": 127}, + {"modality": "IMAGE", "tokenCount": 1160}, ], "candidatesTokensDetails": [ - {"modality": "TEXT", "tokenCount": 30}, - {"modality": "AUDIO", "tokenCount": 20}, + {"modality": "TEXT", "tokenCount": 29}, + {"modality": "AUDIO", "tokenCount": 95}, ], } - cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata) + usage = handler._create_usage_object_from_metadata( + usage_metadata=usage_metadata, model="gemini-live-2.5-flash" + ) - # Should include both text and audio costs - assert cost > 0 - assert cost > (100 * 0.000001) + ( - 50 * 0.000002 - ) # Should be higher due to audio + assert usage.prompt_tokens == 1300, "the full prompt count must survive, not just its text share" + assert usage.completion_tokens == 124 + assert usage.prompt_tokens_details.text_tokens == 13 + assert usage.prompt_tokens_details.audio_tokens == 127 + assert usage.prompt_tokens_details.image_tokens == 1160 + assert usage.completion_tokens_details.text_tokens == 29 + assert usage.completion_tokens_details.audio_tokens == 95 - @patch( - "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info" + def test_usage_sums_repeated_modality_entries(self, handler): + """A modality can appear more than once across aggregated turns; sum, don't overwrite.""" + usage = handler._create_usage_object_from_metadata( + usage_metadata={ + "promptTokenCount": 40, + "candidatesTokenCount": 0, + "promptTokensDetails": [ + {"modality": "IMAGE", "tokenCount": 10}, + {"modality": "IMAGE", "tokenCount": 25}, + {"modality": "TEXT", "tokenCount": 5}, + ], + }, + model="gemini-live-2.5-flash", + ) + assert usage.prompt_tokens_details.image_tokens == 35 + assert usage.prompt_tokens_details.text_tokens == 5 + + NATIVE_AUDIO_MODEL = "gemini-live-2.5-flash-preview-native-audio-09-2025" + + # A four-turn native-audio session. Google charges per turn for the whole session context + # window, so the prompt side repeats the accumulated audio while the candidates side reports + # only that turn's own response. The last turn names AUDIO and omits its tokenCount, which is + # the shape Live really emits at the end of a spoken answer. + AUDIO_SESSION: tuple[_LiveTurn, ...] = ( + {"prompt": (14, 122), "candidates": (8, 20)}, + {"prompt": (21, 182), "candidates": (5, 50)}, + {"prompt": (24, 203), "candidates": (13, 27)}, + {"prompt": (24, 203), "candidates": (0, 3), "candidate_audio_token_count_missing": True}, ) - def test_calculate_cost_with_web_search(self, mock_get_model_info, handler): - """Test cost calculation with web search (tool use)""" - mock_get_model_info.return_value = { - "input_cost_per_token": 0.000001, - "output_cost_per_token": 0.000002, - "web_search_cost_per_request": 0.01, - } - usage_metadata = { - "promptTokenCount": 100, - "candidatesTokenCount": 50, - "totalTokenCount": 150, - "toolUsePromptTokenCount": 10, - } + @staticmethod + def _live_messages(turns: Sequence[_LiveTurn]) -> list[dict[str, object]]: + """Wrap (text, audio) prompt/candidate pairs as the server messages a Live session emits.""" + return [{"type": "session.created", "session": {"id": "s"}}] + [ + { + "type": "response.done", + "usageMetadata": { + "promptTokenCount": sum(turn["prompt"]), + "candidatesTokenCount": sum(turn["candidates"]), + "totalTokenCount": sum(turn["prompt"]) + sum(turn["candidates"]), + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": turn["prompt"][0]}, + {"modality": "AUDIO", "tokenCount": turn["prompt"][1]}, + ], + "candidatesTokensDetails": ( + [{"modality": "AUDIO"}] + if turn.get("candidate_audio_token_count_missing") + else [ + {"modality": "TEXT", "tokenCount": turn["candidates"][0]}, + {"modality": "AUDIO", "tokenCount": turn["candidates"][1]}, + ] + ), + }, + } + for turn in turns + ] - cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata) + @staticmethod + def _session_usage( + handler: VertexAILivePassthroughLoggingHandler, + mock_logging_obj: MagicMock, + messages: list[dict[str, object]], + model: str, + ) -> Usage: + result = handler.vertex_ai_live_passthrough_handler( + websocket_messages=messages, + logging_obj=mock_logging_obj, + url_route="/vertex_ai/live", + start_time=datetime.now(), + end_time=datetime.now(), + request_body={}, + model=model, + ) + assert result["result"] is not None, "the handler must produce a usage-bearing response to bill" + return result["result"].usage - # Should include web search cost - expected_base_cost = (100 * 0.000001) + (50 * 0.000002) - # The web search cost might be handled differently, so just check it's reasonable - assert cost >= expected_base_cost - assert cost > 0 + @classmethod + def _session_cost( + cls, + handler: VertexAILivePassthroughLoggingHandler, + mock_logging_obj: MagicMock, + messages: list[dict[str, object]], + model: str, + ) -> float: + from litellm.cost_calculator import completion_cost + from litellm.types.utils import ModelResponse + + usage = cls._session_usage(handler, mock_logging_obj, messages, model) + return completion_cost( + completion_response=ModelResponse( + id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[] + ), + model=f"vertex_ai/{model}", + custom_llm_provider="vertex_ai", + call_type="acompletion", + ) + + @classmethod + def _expected_session_cost(cls, turns: Sequence[_LiveTurn]) -> float: + from litellm.utils import get_model_info + + info = get_model_info(model=cls.NATIVE_AUDIO_MODEL, custom_llm_provider="vertex_ai") + return ( + sum(turn["prompt"][0] for turn in turns) * info["input_cost_per_token"] + + sum(turn["prompt"][1] for turn in turns) * info["input_cost_per_audio_token"] + + sum(turn["candidates"][0] for turn in turns) * info["output_cost_per_token"] + + sum(turn["candidates"][1] for turn in turns) * info["output_cost_per_audio_token"] + ) + + def test_every_turn_of_a_session_is_billed(self, handler, mock_logging_obj): + """Google charges per turn for the whole context window, so every turn adds to the bill. + + Billing one snapshot instead gives away all the other turns: on this session the + largest single turn is well under the session total, and its share of the audio is + priced 6x the text rate, so the gap is money rather than rounding. + """ + turns = self.AUDIO_SESSION[:3] + cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL) + + assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9) + widest_single_turn = max(self._expected_session_cost([turn]) for turn in turns) + assert cost > widest_single_turn, "billing one snapshot drops every other turn of the session" + + def test_audio_named_without_a_token_count_bills_at_the_audio_rate(self, handler, mock_logging_obj): + """Live can name the modality carrying the rest of a turn and omit its tokenCount. + + Reading the absent key as zero left those tokens inside candidatesTokenCount but outside + the breakdown, so the calculator charged real speech at the text output rate. At this + entry's rates the last turn's 3 audio tokens are $0.0000360 rather than $0.0000060. + """ + turns = self.AUDIO_SESSION + usage = self._session_usage(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL) + + assert usage.completion_tokens_details.audio_tokens == 100, "the unpriced entry takes the turn's residual" + assert usage.completion_tokens_details.text_tokens == 26 + assert usage.completion_tokens == 126 + + cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL) + assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9) + + TOOL_USE_PER_TURN = (100, 250, 400) + + def _grounded_messages(self): + """The three-turn session again, with each turn's own toolUsePromptTokenCount attached.""" + messages = self._live_messages(self.AUDIO_SESSION[:3]) + head, turns = messages[0], messages[1:] + return [head] + [ + {**message, "usageMetadata": {**message["usageMetadata"], "toolUsePromptTokenCount": tool_use}} + for message, tool_use in zip(turns, self.TOOL_USE_PER_TURN) + ] + + def test_server_side_tool_use_prompt_tokens_are_summed_over_the_session(self, handler, mock_logging_obj): + """toolUsePromptTokenCount rode the unknown-key pass-through, so it took the first turn only. + + Every other total beside it is summed across the session, and the first turn is the + smallest number in the series, so a grounded session logged far fewer tool-use tokens + than it used. This session's turns are deliberately distinct, so 750 can only come from + summing: first-turn selection gives 100, last-turn or max gives 400. + """ + grounded = self._grounded_messages() + + usage = self._session_usage(handler, mock_logging_obj, grounded, self.NATIVE_AUDIO_MODEL) + assert usage.prompt_tokens_details.tool_use_tokens == sum(self.TOOL_USE_PER_TURN) + + @staticmethod + def _grounding_frame(metadata: dict[str, object]) -> dict[str, object]: + """One server frame carrying grounding metadata, the way Live reports it.""" + return {"type": "response.done", "serverContent": {"groundingMetadata": metadata}} + + def test_web_grounding_is_counted_so_it_can_be_billed(self, handler, mock_logging_obj): + """Live reports grounding in the server frames and never in usageMetadata. + + Nothing read those frames, so web_search_requests stayed unset and the cost path's only + trigger for the per-query grounding charge never fired. Google bills a grounded Live + prompt on top of its tokens, so the whole fee was missing from the bill. + """ + messages = [ + self._grounding_frame( + { + "webSearchQueries": ["who won the 2026 world cup final"], + "groundingChunks": [{"web": {"uri": "https://example.com"}}], + } + ), + *self._live_messages(self.AUDIO_SESSION[:1]), + ] + + usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL) + + assert usage.prompt_tokens_details.web_search_requests == 1, "a grounded turn must report its query" + assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None + + def test_maps_grounding_is_counted_under_its_own_sku(self, handler, mock_logging_obj): + """Maps grounding is a separate SKU from web search, so it needs its own counter. + + A maps-only turn carries grounding chunks but no webSearchQueries, so counting queries + alone would report nothing and bill nothing. + """ + messages = [ + self._grounding_frame({"groundingChunks": [{"maps": {"placeId": "abc123"}}]}), + *self._live_messages(self.AUDIO_SESSION[:1]), + ] + + usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL) + + assert usage.prompt_tokens_details.google_maps_grounding_requests == 1 + assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None + + def test_an_ungrounded_session_reports_no_grounding(self, handler, mock_logging_obj): + """The counters must stay absent when no tool ran, or every session pays a grounding fee.""" + usage = self._session_usage( + handler, mock_logging_obj, self._live_messages(self.AUDIO_SESSION[:1]), self.NATIVE_AUDIO_MODEL + ) + + assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None + assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None + + def test_grounding_adds_its_query_fee_to_the_session_bill(self, handler, mock_logging_obj): + """The counter only matters if it reaches the bill, so assert against the cost, not the field. + + Same tokens either way: the difference between the two sessions is the grounding fee alone. + """ + turns = self.AUDIO_SESSION[:1] + plain = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL) + grounded = self._session_cost( + handler, + mock_logging_obj, + [self._grounding_frame({"webSearchQueries": ["q"]}), *self._live_messages(turns)], + self.NATIVE_AUDIO_MODEL, + ) + + assert grounded > plain, "a grounded session must cost more than the same tokens ungrounded" + + def _priced_logging_obj(self) -> LiteLLMLoggingObj: + """A real logging object, since the session's price is handed to it turn by turn.""" + logging_obj = LiteLLMLoggingObj( + model=self.NATIVE_AUDIO_MODEL, + messages=[], + stream=True, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id="live-session", + function_id="live", + ) + logging_obj.update_environment_variables( + model=self.NATIVE_AUDIO_MODEL, + user="u", + optional_params={}, + litellm_params={}, + call_type="pass_through_endpoint", + ) + logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai" + return logging_obj + + def _billed_session( + self, handler: VertexAILivePassthroughLoggingHandler, messages: list[dict[str, object]] + ) -> tuple[float, CostBreakdown]: + logging_obj = self._priced_logging_obj() + result = handler.vertex_ai_live_passthrough_handler( + websocket_messages=messages, + logging_obj=logging_obj, + url_route="/vertex_ai/live", + start_time=datetime.now(), + end_time=datetime.now(), + request_body={}, + model=self.NATIVE_AUDIO_MODEL, + custom_llm_provider="vertex_ai", + ) + assert result["result"] is not None, "the handler must produce a usage-bearing response to bill" + assert logging_obj.cost_breakdown is not None, "the session's price must reach the logging object" + return result["result"]._hidden_params["response_cost"], logging_obj.cost_breakdown + + def test_each_grounded_turn_pays_its_own_query_fee(self, handler): + """Google charges the grounding fee per grounded prompt, not per session. + + Summing the session into one usage collapsed two grounded turns into one query, so the + second question was answered for free. The bill now grows by one fee per grounded turn. + """ + head, turn = self._live_messages(self.AUDIO_SESSION[:1]) + grounding = self._grounding_frame({"webSearchQueries": ["q"]}) + + plain_cost, _ = self._billed_session(handler, [head, turn, turn]) + one_cost, one_breakdown = self._billed_session(handler, [head, grounding, turn, turn]) + two_cost, two_breakdown = self._billed_session(handler, [head, grounding, turn, grounding, turn]) + + fee = one_cost - plain_cost + assert fee > 0, "a grounded turn must cost more than the same tokens ungrounded" + assert two_cost - plain_cost == pytest.approx(2 * fee), "two grounded turns must pay the fee twice" + assert two_breakdown["total_cost"] == pytest.approx(two_cost) + assert two_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"]) + + def test_a_query_repeated_across_turns_is_reported_once_per_turn(self, handler): + """The reported query count must agree with the bill, which charges every grounded turn. + + The session usage collapsed duplicate query strings across turns while the price was + per turn, so two turns asking the same question paid two fees yet reported one query. + Duplicates within one turn still collapse, since that turn ran one search. + """ + head, turn = self._live_messages(self.AUDIO_SESSION[:1]) + grounding = self._grounding_frame({"webSearchQueries": ["q"]}) + logging_obj = self._priced_logging_obj() + + result = handler.vertex_ai_live_passthrough_handler( + websocket_messages=[head, grounding, turn, grounding, turn], + logging_obj=logging_obj, + url_route="/vertex_ai/live", + start_time=datetime.now(), + end_time=datetime.now(), + request_body={}, + model=self.NATIVE_AUDIO_MODEL, + custom_llm_provider="vertex_ai", + ) + _, one_breakdown = self._billed_session(handler, [head, grounding, turn]) + repeated_within_turn = handler._session_usage( + [head, self._grounding_frame({"webSearchQueries": ["q", "q"]}), turn], self.NATIVE_AUDIO_MODEL + ) + + assert result["result"].usage.prompt_tokens_details.web_search_requests == 2 + assert logging_obj.cost_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"]) + assert repeated_within_turn.prompt_tokens_details.web_search_requests == 1 + + def test_the_fixed_cost_margin_is_charged_once_per_session(self, handler): + """A fixed cost margin is a flat per-request fee, and a Live session is one spend row. + + Pricing each turn on its own applied the fixed margin per turn, so a two-turn session paid it + twice. The session now carries the fixed margin once no matter how many turns it billed. + """ + head, turn = self._live_messages(self.AUDIO_SESSION[:1]) + grounding = self._grounding_frame({"webSearchQueries": ["q"]}) + messages = [head, grounding, turn, grounding, turn] + + plain_cost, _ = self._billed_session(handler, messages) + + fixed_amount = 0.01 + with patch.object(litellm, "cost_margin_config", {"vertex_ai": {"fixed_amount": fixed_amount}}): + margined_cost, breakdown = self._billed_session(handler, messages) + + assert margined_cost - plain_cost == pytest.approx( + fixed_amount + ), "a two-turn session must add the fixed margin once, not once per billed turn" + assert breakdown["margin_fixed_amount"] == pytest.approx(fixed_amount) + assert breakdown["margin_total_amount"] == pytest.approx(fixed_amount) + + def test_reporting_tool_use_tokens_does_not_move_the_bill(self, handler, mock_logging_obj): + """Deliberate boundary: these tokens are reported here, and priced nowhere. + + generic_cost_per_token reads the input bill out of prompt_tokens_details, and falls + back to prompt_tokens only when the details carry no text or a cache hit overlaps them, + so adding tool-use tokens to prompt_tokens is worth nothing on an ordinary Live turn and + over-charges against the cache-overlap correction when it is not. Pricing them belongs + in the shared input-cost path, beside the modality terms that already read the details. + """ + turns = self.AUDIO_SESSION[:3] + plain_cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL) + grounded_cost = self._session_cost( + handler, mock_logging_obj, self._grounded_messages(), self.NATIVE_AUDIO_MODEL + ) + + assert plain_cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9) + assert grounded_cost == pytest.approx(plain_cost, rel=1e-9), "reporting tool use must not move the bill" + + def test_a_malformed_details_entry_does_not_cost_the_whole_session(self, handler, mock_logging_obj): + """A ``*TokensDetails`` value that is not a list of objects must not take the session down. + + The handler's only error path returns no result at all, so one odd frame used to throw + while reading it and the whole session billed nothing. The good turns still bill. + """ + turns = self.AUDIO_SESSION[:3] + messages = self._live_messages(turns) + mangled = [dict(message) for message in messages] + mangled[1]["usageMetadata"] = {**mangled[1]["usageMetadata"], "promptTokensDetails": "TEXT"} + + usage = self._session_usage(handler, mock_logging_obj, mangled, self.NATIVE_AUDIO_MODEL) + + surviving = turns[1:] + assert usage.prompt_tokens_details.audio_tokens == sum(turn["prompt"][1] for turn in surviving) + assert usage.prompt_tokens_details.text_tokens == sum(turn["prompt"][0] for turn in surviving) + assert usage.prompt_tokens == sum(sum(turn["prompt"]) for turn in turns), "the totals still cover every turn" + + direct = handler._create_usage_object_from_metadata( + usage_metadata={ + "promptTokenCount": 40, + "candidatesTokenCount": 12, + "promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 40}, "AUDIO"], + "candidatesTokensDetails": {"modality": "TEXT", "tokenCount": 12}, + }, + model=self.NATIVE_AUDIO_MODEL, + ) + assert direct.prompt_tokens_details.audio_tokens == 40, "the well-formed entry beside a bad one still counts" + assert direct.completion_tokens == 12 + + @pytest.mark.parametrize( + "label,prompt_details,candidate_details", + [ + ("text only", [("TEXT", 6)], [("TEXT", 2)]), + ("audio in", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 18)]), + ("image in", [("TEXT", 10), ("IMAGE", 258)], [("TEXT", 24)]), + ("frames in", [("TEXT", 11), ("IMAGE", 1032)], [("TEXT", 26)]), + ("audio both ways", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 29), ("AUDIO", 95)]), + ], + ) + def test_live_session_bills_each_modality_at_its_own_rate(self, handler, label, prompt_details, candidate_details): + """Every payload here is a real Vertex Live session's usageMetadata. + + Before the fix these billed the text share only, from 1x (text) to 55x under. + The expected amount is derived from the entry's own rates rather than hardcoded, + so this stays correct as prices move, and it is asserted exactly, so dropping a + modality and double-charging one both fail. + """ + from litellm.cost_calculator import completion_cost + from litellm.types.utils import ModelResponse + from litellm.utils import get_model_info + + model = self.NATIVE_AUDIO_MODEL + info = get_model_info(model=model, custom_llm_provider="vertex_ai") + + text_in = info["input_cost_per_token"] + audio_in = info.get("input_cost_per_audio_token") or text_in + image_in = info.get("input_cost_per_image_token") or text_in + text_out = info["output_cost_per_token"] + audio_out = info.get("output_cost_per_audio_token") or text_out + rate_in = {"TEXT": text_in, "AUDIO": audio_in, "IMAGE": image_in} + rate_out = {"TEXT": text_out, "AUDIO": audio_out} + + expected = sum(c * rate_in[m] for m, c in prompt_details) + sum(c * rate_out[m] for m, c in candidate_details) + + usage = handler._create_usage_object_from_metadata( + usage_metadata={ + "promptTokenCount": sum(c for _, c in prompt_details), + "candidatesTokenCount": sum(c for _, c in candidate_details), + "promptTokensDetails": [{"modality": m, "tokenCount": c} for m, c in prompt_details], + "candidatesTokensDetails": [{"modality": m, "tokenCount": c} for m, c in candidate_details], + }, + model=model, + ) + + cost = completion_cost( + completion_response=ModelResponse( + id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[] + ), + model=f"vertex_ai/{model}", + custom_llm_provider="vertex_ai", + call_type="acompletion", + ) + + assert cost == pytest.approx(expected, rel=1e-9), label + + text_only = sum(c for m, c in prompt_details if m == "TEXT") * text_in + sum( + c for m, c in candidate_details if m == "TEXT" + ) * text_out + if any(m != "TEXT" for m, _ in prompt_details + candidate_details) and audio_in != text_in: + assert cost > text_only, f"{label}: non-text modalities must add cost" def test_vertex_ai_live_passthrough_handler_integration( self, handler, mock_logging_obj, sample_websocket_messages @@ -376,6 +788,7 @@ class TestVertexAILivePassthroughIntegration: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} + mock._response_cost_calculator.return_value = None return mock @patch( @@ -509,6 +922,7 @@ class TestVertexAILivePassthroughErrorHandling: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} + mock._response_cost_calculator.return_value = None return mock def test_invalid_websocket_messages_format(self): @@ -540,25 +954,24 @@ class TestVertexAILivePassthroughErrorHandling: result = handler._extract_usage_metadata_from_websocket_messages(messages) assert result is None - @patch( - "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info" - ) - def test_cost_calculation_with_missing_model_info(self, mock_get_model_info): - """Test cost calculation when model info is missing""" + def test_usage_without_modality_details(self): + """Older payloads carry only the totals; fall back to them rather than reporting zero.""" handler = VertexAILivePassthroughLoggingHandler() - # Mock missing model info - mock_get_model_info.return_value = {} + usage = handler._create_usage_object_from_metadata( + usage_metadata={ + "promptTokenCount": 100, + "candidatesTokenCount": 50, + "totalTokenCount": 150, + }, + model="unknown-model", + ) - usage_metadata = { - "promptTokenCount": 100, - "candidatesTokenCount": 50, - "totalTokenCount": 150, - } - - # Should not raise an exception, should return 0 or handle gracefully - cost = handler._calculate_live_api_cost("unknown-model", usage_metadata) - assert cost == 0.0 + assert usage.prompt_tokens == 100 + assert usage.completion_tokens == 50 + assert usage.total_tokens == 150 + assert usage.prompt_tokens_details.audio_tokens is None + assert usage.prompt_tokens_details.image_tokens is None def test_handler_with_none_websocket_messages(self, mock_logging_obj): """Test handler with None websocket messages""" diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py index ee4750e9db8..0d33435cf7a 100644 --- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py +++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py @@ -372,7 +372,7 @@ async def test_aresponses_fallback_on_in_stream_error_event(): raised = mock_fallback.await_args.kwargs["e"] assert isinstance(raised, MidStreamFallbackError) assert raised.status_code == 429 - assert isinstance(raised.original_exception, litellm.APIError) + assert isinstance(raised.original_exception, litellm.RateLimitError) assert raised.original_exception.status_code == 429 assert mock_fallback.await_args.kwargs["kwargs"]["input"] == "original question" diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index 5b06c5fdb01..b18bf9351c8 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -11,7 +11,7 @@ import litellm from unittest.mock import patch, MagicMock, AsyncMock from create_mock_standard_logging_payload import create_standard_logging_payload from litellm.types.utils import StandardLoggingPayload -from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo +from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS @@ -630,10 +630,12 @@ def test_deployment_callback_respects_cooldown_time(model_list): @pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) -def test_log_retry(model_list, metadata_key): - """log_retry appends one flat record per failed attempt and copies neither the request kwargs nor - the request metadata into it""" +def test_log_retry(model_list: list[DeploymentTypedDict], metadata_key: str) -> None: + """log_retry appends one flat record per failed attempt, copies neither the request kwargs nor the + request metadata into it, counts every failed attempt of the request independently of the + per-hop attempted_retries, and never trusts a negative count planted before the first failure""" router = Router(model_list=model_list) + rate_limit_error = litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo") new_kwargs = router.log_retry( kwargs={ "model": "gpt-3.5-turbo", @@ -641,7 +643,7 @@ def test_log_retry(model_list, metadata_key): "messages": [{"role": "user", "content": "hi"}], metadata_key: {"model_info": {"id": "deployment-1"}, "attempted_retries": 2, "user_api_key": "sk-proxy"}, }, - e=litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo"), + e=rate_limit_error, ) assert json.loads(json.dumps(new_kwargs[metadata_key]["previous_models"])) == [ { @@ -652,6 +654,10 @@ def test_log_retry(model_list, metadata_key): "attempted_retries": 2, } ] + assert new_kwargs[metadata_key]["request_retry_count"] == 1 + assert router.log_retry(kwargs=new_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 2 + planted_kwargs = {"model": "gpt-3.5-turbo", metadata_key: {"request_retry_count": -100}} + assert router.log_retry(kwargs=planted_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 1 def test_update_usage(model_list): diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 768ea332677..8b04d7af70a 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -695,6 +695,38 @@ def test_vertex_cost_and_usage_aggregation(monkeypatch): assert result.failed_requests == 0 +def test_vertex_batch_usage_preserves_modality_token_details(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "vertex_ai/gemini-embedding-2", + { + "input_cost_per_token_batches": 1e-7, + "input_cost_per_audio_token_batches": 3.25e-6, + "input_cost_per_image_token_batches": 2.25e-7, + "input_cost_per_video_token_batches": 6e-6, + }, + ) + responses = [ + { + "response": { + "usageMetadata": { + "promptTokenCount": 84, + "candidatesTokenCount": 0, + "totalTokenCount": 84, + "promptTokensDetails": [ + {"modality": "AUDIO", "tokenCount": 64}, + {"modality": "TEXT", "tokenCount": 20}, + ], + } + } + } + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-embedding-2") + + assert result.prompt_cost == pytest.approx(64 * 3.25e-6 + 20 * 1e-7) + + def test_vertex_cost_skips_none_response_body(monkeypatch): import litellm.cost_calculator as cc diff --git a/tests/test_litellm/compression/test_compress.py b/tests/test_litellm/compression/test_compress.py index 6827c37dfd5..6e908bcbdcd 100644 --- a/tests/test_litellm/compression/test_compress.py +++ b/tests/test_litellm/compression/test_compress.py @@ -6,7 +6,8 @@ never rewrite. It is consumed by compress() and by the Headroom guardrail, so the two agree on what "never compress this" means. """ -from litellm.compression.compress import get_protected_indices +from litellm.compression.compress import compress, get_protected_indices +from litellm.types.utils import CallTypes def test_protects_system_last_user_and_last_assistant(): @@ -53,3 +54,94 @@ def test_every_system_row_is_protected(): def test_no_user_or_assistant_rows(): assert sorted(get_protected_indices([{"role": "system", "content": "sys"}])) == [0] assert get_protected_indices([]) == () + + +def test_mid_history_cache_control_part_is_protected(): + # A large cached tool result from a few turns back, not the last user or + # last assistant row -- exactly the row a provider prompt-cache pins to + # exact bytes. Rewriting it (even leaving the marker on) changes those + # bytes and turns the next request's cache read into a cache write. + messages = [ + {"role": "user", "content": "old question"}, + {"role": "assistant", "content": "old answer"}, + { + "role": "user", + "content": [ + {"type": "text", "text": "a large cached tool result", "cache_control": {"type": "ephemeral"}}, + ], + }, + {"role": "assistant", "content": "ack"}, + {"role": "user", "content": "live instruction"}, + ] + + # index 3 = last assistant, index 4 = last user (both protected by role + # regardless), index 2 = the cache_control-marked row itself. + assert sorted(get_protected_indices(messages)) == [2, 3, 4] + + +def test_cache_control_directly_on_message_is_protected(): + messages = [ + {"role": "user", "content": "old question", "cache_control": {"type": "ephemeral"}}, + {"role": "assistant", "content": "old answer"}, + {"role": "user", "content": "live instruction"}, + ] + + assert sorted(get_protected_indices(messages)) == [0, 1, 2] + + +def test_cache_control_protection_does_not_duplicate_already_protected_rows(): + # The last user row is already protected by role; marking it too must not + # produce a duplicate index. + messages = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "live", "cache_control": {"type": "ephemeral"}}, + ] + + protected = get_protected_indices(messages) + + assert sorted(protected) == [0, 1] + assert len(protected) == len(set(protected)) + + +def test_content_that_is_not_a_list_of_mappings_is_not_treated_as_cache_control(): + # Defensive: a plain string content, or a list of non-dict items, must not + # raise or be misread as carrying a breakpoint. + messages = [ + {"role": "assistant", "content": "plain string content"}, + {"role": "user", "content": ["not", "a", "dict", "list"]}, + {"role": "user", "content": "live instruction"}, + ] + + assert sorted(get_protected_indices(messages)) == [0, 2] + + +def test_compress_keeps_part_level_cache_control_row_verbatim(): + # compress() scores text-only copies of the rows, where a part-level marker + # is gone; protection has to read the original rows or the pinned row is stubbed. + stale_log = {"role": "user", "content": [{"type": "text", "text": "stale log line " * 2000}]} + pinned = { + "role": "user", + "content": [ + {"type": "text", "text": "cached tool result " * 2000, "cache_control": {"type": "ephemeral"}}, + ], + } + messages = [ + stale_log, + {"role": "assistant", "content": "old answer"}, + pinned, + {"role": "assistant", "content": "ack"}, + {"role": "user", "content": "live instruction"}, + ] + + result = compress( + messages, + model="gpt-4o", + call_type=CallTypes.anthropic_messages, + compression_trigger=1000, + compression_target=500, + ) + + assert len(result["messages"]) == len(messages) + assert result["messages"][2] == pinned + assert result["messages"][0] != stale_log + assert len(result["cache"]) >= 1 diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 2fd5fa76d8e..bb29bfed283 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -2625,7 +2625,7 @@ class TestLoggingOnlyApplyGuardrail: assert [e["guardrail_status"] for e in entries] == ["success"] @pytest.mark.asyncio - async def test_native_lifecycle_hook_guardrail_is_left_alone(self): + async def test_native_lifecycle_hook_guardrail_scans_in_logging_only(self): class _NativeHooks(_ApplyOnlyObserver): use_native_lifecycle_hooks = True @@ -2634,9 +2634,9 @@ class TestLoggingOnlyApplyGuardrail: out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) - assert guardrail.calls == [] - assert out_kwargs is kwargs + assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])] assert out_response is response + assert out_kwargs["standard_logging_object"]["guardrail_information"] @pytest.mark.asyncio async def test_aresponses_scans_logged_messages_when_input_is_cleared(self): @@ -2890,6 +2890,61 @@ class TestCustomGuardrailPostCallSuccessDeploymentHook: assert len(_guardrail_entries(request_data)) == 1 +class _NativeLifecycleLoggingGuardrail(CustomGuardrail): + """Native lifecycle guardrail that also implements apply_guardrail, like the azure guards.""" + + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__(self): + from litellm.types.guardrails import GuardrailEventHooks + + super().__init__( + guardrail_name="native-logging-guardrail", + event_hook=GuardrailEventHooks.logging_only, + ) + self.calls: list[tuple[Literal["request", "response"], list[str]]] = [] + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: "LiteLLMLoggingObj | None" = None, + ) -> GenericGuardrailAPIInputs: + self.calls.append((input_type, list(inputs.get("texts") or []))) + return inputs + + +@pytest.mark.asyncio +async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(): + """A use_native_lifecycle_hooks guardrail accepts mode logging_only and its + async_logging_hook scans kwargs["async_complete_streaming_response"], not the raw result.""" + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _NativeLifecycleLoggingGuardrail() + assembled = ModelResponse( + choices=[Choices(message=Message(role="assistant", content="assembled stream text"))] + ) + sentinel_result = object() + kwargs = { + "model": "gpt-5.4-mini", + "messages": [{"role": "user", "content": "hi"}], + "litellm_call_id": "call-1", + "litellm_params": {"metadata": {}}, + "optional_params": {}, + "standard_logging_object": {"guardrail_information": None}, + "async_complete_streaming_response": assembled, + } + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=sentinel_result, call_type=CallTypes.acompletion.value + ) + + assert out_result is sentinel_result + assert ("response", ["assembled stream text"]) in guardrail.calls + assert out_kwargs["standard_logging_object"]["guardrail_information"] + + class TestPreCallHookResponseIsNotLoggedVerbatim: """Regression for LIT-6935: a pre_call hook returning the request payload leaked the prompt into ``guardrail_response`` and from there onto OTEL guardrail spans.""" diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 2fadb1fd959..30b158e3b2c 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2,6 +2,7 @@ import json from datetime import datetime, timezone import pytest +from collections.abc import Mapping from fastapi.testclient import TestClient import litellm @@ -74,6 +75,108 @@ def test_missing_cache_read_policy_preserves_billing(prompt_tokens, read_rate, s assert prompt_cost == pytest.approx((prompt_tokens - 100) * billed[0] + 100 * billed[4]) +def test_generic_cost_per_token_prefers_audio_per_second_rate() -> None: + model_info: ModelInfo = { + "key": "gemini-embedding-2", + "max_tokens": None, + "max_input_tokens": None, + "max_output_tokens": None, + "input_cost_per_token": 2e-7, + "input_cost_per_audio_token": 6.5e-6, + "input_cost_per_audio_per_second": 0.00016, + "output_cost_per_token": 0.0, + "litellm_provider": "vertex_ai", + "mode": "embedding", + "supported_openai_params": None, + } + usage = Usage( + prompt_tokens=64, + completion_tokens=0, + total_tokens=64, + prompt_tokens_details=PromptTokensDetailsWrapper( + audio_tokens=64, + audio_length_seconds=2, + ), + ) + + prompt_cost, _ = generic_cost_per_token( + model="gemini-embedding-2", + usage=usage, + custom_llm_provider="vertex_ai", + model_info=model_info, + ) + + assert prompt_cost == pytest.approx(2 * 0.00016) + + +def test_generic_cost_per_token_prefers_image_per_image_rate() -> None: + model_info: ModelInfo = { + "key": "gemini-embedding-2", + "max_tokens": None, + "max_input_tokens": None, + "max_output_tokens": None, + "input_cost_per_token": 2e-7, + "input_cost_per_image_token": 4.5e-7, + "input_cost_per_image": 0.00012, + "output_cost_per_token": 0.0, + "litellm_provider": "vertex_ai", + "mode": "embedding", + "supported_openai_params": None, + } + usage = Usage( + prompt_tokens=258, + completion_tokens=0, + total_tokens=258, + prompt_tokens_details=PromptTokensDetailsWrapper( + image_tokens=258, + image_count=1, + ), + ) + + prompt_cost, _ = generic_cost_per_token( + model="gemini-embedding-2", + usage=usage, + custom_llm_provider="vertex_ai", + model_info=model_info, + ) + + assert prompt_cost == pytest.approx(0.00012) + + +def test_generic_cost_per_token_prefers_video_per_second_rate() -> None: + model_info: ModelInfo = { + "key": "gemini-embedding-2", + "max_tokens": None, + "max_input_tokens": None, + "max_output_tokens": None, + "input_cost_per_token": 2e-7, + "input_cost_per_video_token": 1.2e-5, + "input_cost_per_video_per_second": 0.00079, + "output_cost_per_token": 0.0, + "litellm_provider": "vertex_ai", + "mode": "embedding", + "supported_openai_params": None, + } + usage = Usage( + prompt_tokens=516, + completion_tokens=0, + total_tokens=516, + prompt_tokens_details=PromptTokensDetailsWrapper( + video_tokens=516, + video_length_seconds=2, + ), + ) + + prompt_cost, _ = generic_cost_per_token( + model="gemini-embedding-2", + usage=usage, + custom_llm_provider="vertex_ai", + model_info=model_info, + ) + + assert prompt_cost == pytest.approx(2 * 0.00079) + + def test_missing_cache_read_uses_off_peak_input_rate(): from datetime import datetime, timezone @@ -4039,7 +4142,7 @@ def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeyp cache_read_input_token_cost=6e-7, cache_read_input_audio_token_cost=6e-7, cache_creation_input_token_cost=7.5e-6, - cache_creation_input_token_cost_above_1hr=0.0, + cache_creation_input_token_cost_above_1hr=7.5e-6, output_cost_per_reasoning_token=3e-5, ) assert breakdown.cache_read_cost == pytest.approx(200_000 * rates.cache_read_input_token_cost) @@ -4560,7 +4663,7 @@ GEMINI_35_FLASH_LITE_TIER_RATES_BY_SURFACE = [ ("gemini", "priority", 5.4e-07, 4.5e-06, 5e-08), ("vertex_ai", None, 3e-07, 2.5e-06, 3e-08), ("vertex_ai", "flex", 1.5e-07, 1.25e-06, 1.5e-08), - ("vertex_ai", "priority", 5.4e-07, 4.5e-06, 5e-08), + ("vertex_ai", "priority", 5.4e-07, 4.5e-06, 5.4e-08), ] @@ -5334,3 +5437,72 @@ def test_realtime_models_bill_cached_text_and_audio_at_their_cache_read_rates( prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider=custom_llm_provider) assert prompt_cost == pytest.approx(expected_prompt_cost) + + +def test_generic_cost_per_token_bills_cache_creation_at_the_input_rate_without_a_write_price(): + """Azure and OpenAI publish no cache-write price and bill cache writes as ordinary input. + A deployment priced with only input, output, and cache-read rates must bill the creation + tokens the provider reports at the input rate, never at 0. The numbers are a cold 7,336-token + prompt on a deployment that reports all but 3 of them as cache creation.""" + model_info = { + "input_cost_per_token": 2e-7, + "output_cost_per_token": 1.25e-6, + "cache_read_input_token_cost": 2e-8, + } + usage = Usage( + prompt_tokens=7336, + completion_tokens=23, + total_tokens=7359, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, cache_creation_tokens=7333), + ) + + prompt_cost, completion_cost = generic_cost_per_token( + model="custom-priced-deployment", usage=usage, custom_llm_provider="azure", model_info=model_info + ) + + assert prompt_cost == pytest.approx(7336 * 2e-7) + assert completion_cost == pytest.approx(23 * 1.25e-6) + + +@pytest.mark.parametrize( + ("cache_rates", "current_time", "expected_creation", "expected_creation_1h"), + ( + pytest.param({}, None, 2e-7, 2e-7, id="no-write-price-uses-the-input-rate"), + pytest.param({"cache_creation_input_token_cost": 2.5e-7}, None, 2.5e-7, 2.5e-7, id="no-1h-price-uses-the-write-price"), + pytest.param({"cache_creation_input_token_cost": 0.0}, None, 0.0, 0.0, id="explicit-zero-stays-zero"), + pytest.param( + {"off_peak_pricing": {"hours_utc": "00:00-23:59", "input_cost_per_token": 1e-7}}, + datetime(2026, 9, 14, 12, tzinfo=timezone.utc), + 1e-7, + 1e-7, + id="no-write-price-uses-the-off-peak-input-rate", + ), + pytest.param( + { + "off_peak_pricing": { + "hours_utc": "00:00-23:59", + "input_cost_per_token": 1e-7, + "cache_creation_input_token_cost": 3e-7, + } + }, + datetime(2026, 9, 14, 12, tzinfo=timezone.utc), + 3e-7, + 3e-7, + id="no-1h-price-uses-the-off-peak-write-price", + ), + ), +) +def test_get_token_base_cost_resolves_missing_cache_write_rates_like_the_tiered_path( + cache_rates: Mapping[str, float | Mapping[str, float | str]], + current_time: datetime | None, + expected_creation: float, + expected_creation_1h: float, +) -> None: + model_info = {"input_cost_per_token": 2e-7, "output_cost_per_token": 1.25e-6, **cache_rates} + usage = Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11) + + _, _, creation, creation_1h, _ = _get_token_base_cost(model_info, usage, current_time=current_time) + + assert creation == pytest.approx(expected_creation) + assert creation_1h == pytest.approx(expected_creation_1h) + diff --git a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py index 42d3df76902..acc6248bf3e 100644 --- a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py @@ -1,4 +1,3 @@ - import httpx import openai import pytest @@ -178,9 +177,7 @@ class TestExceptionCheckers: ] for error_str in error_strings: - result = ExceptionCheckers.is_azure_content_policy_violation_error( - error_str - ) + result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str) assert result is True, f"Should detect policy violation in: {error_str}" def test_is_azure_content_policy_violation_error_case_insensitive(self): @@ -194,12 +191,8 @@ class TestExceptionCheckers: ] for error_str in error_strings: - result = ExceptionCheckers.is_azure_content_policy_violation_error( - error_str - ) - assert ( - result is True - ), f"Should detect policy violation in uppercase: {error_str}" + result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str) + assert result is True, f"Should detect policy violation in uppercase: {error_str}" def test_is_azure_content_policy_violation_error_with_non_policy_errors(self): """Test that non-policy violation errors are not detected as policy violations""" @@ -216,12 +209,8 @@ class TestExceptionCheckers: ] for error_str in error_strings: - result = ExceptionCheckers.is_azure_content_policy_violation_error( - error_str - ) - assert ( - result is False - ), f"Should NOT detect policy violation in: {error_str}" + result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str) + assert result is False, f"Should NOT detect policy violation in: {error_str}" def test_is_azure_content_policy_violation_error_with_partial_matches(self): """Test that partial keyword matches work correctly""" @@ -234,9 +223,7 @@ class TestExceptionCheckers: ] for error_str in positive_cases: - result = ExceptionCheckers.is_azure_content_policy_violation_error( - error_str - ) + result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str) assert result is True, f"Should detect policy violation in: {error_str}" # These should not match even though they contain similar words @@ -248,12 +235,8 @@ class TestExceptionCheckers: ] for error_str in negative_cases: - result = ExceptionCheckers.is_azure_content_policy_violation_error( - error_str - ) - assert ( - result is False - ), f"Should NOT detect policy violation in: {error_str}" + result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str) + assert result is False, f"Should NOT detect policy violation in: {error_str}" gemini_context_window_test_cases = [ @@ -271,12 +254,8 @@ gemini_context_window_test_cases = [ ] -@pytest.mark.parametrize( - "error_message, should_raise_context_window", gemini_context_window_test_cases -) -def test_gemini_context_window_error_mapping( - error_message, should_raise_context_window -): +@pytest.mark.parametrize("error_message, should_raise_context_window", gemini_context_window_test_cases) +def test_gemini_context_window_error_mapping(error_message, should_raise_context_window): """ Tests that the exception_type function correctly maps Gemini's context window exceeded errors to litellm.ContextWindowExceededError. @@ -421,9 +400,7 @@ vertex_rate_limit_test_cases = [ ] -@pytest.mark.parametrize( - "error_message, should_raise_rate_limit", vertex_rate_limit_test_cases -) +@pytest.mark.parametrize("error_message, should_raise_rate_limit", vertex_rate_limit_test_cases) def test_vertex_ai_rate_limit_error_mapping(error_message, should_raise_rate_limit): """ Tests that the exception_type function correctly maps Vertex AI's @@ -458,10 +435,7 @@ class TestGetBodyErrorCode: """Unit tests for _get_body_error_code helper.""" def test_parses_int_code(self): - body = ( - '{"error":{"message":"high demand","type":"upstream_error",' - '"param":"","code":429}}' - ) + body = '{"error":{"message":"high demand","type":"upstream_error","param":"","code":429}}' assert _get_body_error_code(body) == 429 def test_parses_string_code(self): @@ -498,8 +472,7 @@ gemini_body_code_429_test_cases = [ ), ( 503, - '{"error":{"message":"upstream unavailable","type":"upstream_error",' - '"param":"","code":429}}', + '{"error":{"message":"upstream unavailable","type":"upstream_error","param":"","code":429}}', litellm.RateLimitError, "HTTP 503 envelope with body code:429 -> RateLimitError", ), @@ -769,9 +742,7 @@ class _UpstreamHTTPError(Exception): self.message = "upstream failure" self.status_code = status_code self.request = httpx.Request("POST", "https://api.example.com/v1/chat/completions") - self.response = httpx.Response( - status_code=status_code, request=self.request, text="upstream failure" - ) + self.response = httpx.Response(status_code=status_code, request=self.request, text="upstream failure") UPSTREAM_STATUS_CODES = (400, 401, 403, 404, 408, 422, 429, 500, 503) @@ -892,15 +863,13 @@ PROVIDERS_WITHOUT_A_HANDLER = tuple( MINIMAX_401_BODY = ( '{"type":"error","error":{"type":"authorized_error","message":"login fail: Please carry the API secret key ' - "in the 'Authorization' field of the request header (1004)\",\"http_code\":\"401\"}," + 'in the \'Authorization\' field of the request header (1004)","http_code":"401"},' '"request_id":"06ddc9ba97ee6340e38f10e09787f547"}' ) def _expected_for(provider: str, status_code: int) -> tuple[type[Exception], int]: - return DEVIATIONS_FROM_THE_OPENAI_SHAPE.get(provider, {}).get( - status_code, OPENAI_SHAPED[status_code] - ) + return DEVIATIONS_FROM_THE_OPENAI_SHAPE.get(provider, {}).get(status_code, OPENAI_SHAPED[status_code]) @pytest.fixture @@ -910,9 +879,7 @@ def quiet_exception_mapping(monkeypatch): @pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES) @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_an_upstream_status_maps_to_one_exception_per_provider( - provider, status_code, quiet_exception_mapping -): +def test_an_upstream_status_maps_to_one_exception_per_provider(provider, status_code, quiet_exception_mapping): expected_class, expected_status = _expected_for(provider, status_code) with pytest.raises(openai.APIError) as raised: @@ -928,9 +895,7 @@ def test_an_upstream_status_maps_to_one_exception_per_provider( @pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES) @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from( - provider, status_code, quiet_exception_mapping -): +def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from(provider, status_code, quiet_exception_mapping): with pytest.raises(openai.APIError) as raised: exception_type( model="test-model", @@ -943,12 +908,8 @@ def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from( @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_an_already_mapped_litellm_exception_passes_through_untouched( - provider, quiet_exception_mapping -): - already_mapped = litellm.RateLimitError( - message="already mapped", llm_provider=provider, model="test-model" - ) +def test_an_already_mapped_litellm_exception_passes_through_untouched(provider, quiet_exception_mapping): + already_mapped = litellm.RateLimitError(message="already mapped", llm_provider=provider, model="test-model") returned = exception_type( model="test-model", @@ -961,9 +922,7 @@ def test_an_already_mapped_litellm_exception_passes_through_untouched( @pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES) @pytest.mark.parametrize("provider", PROVIDERS_WITHOUT_A_HANDLER) -def test_a_provider_without_a_handler_maps_by_the_upstream_status( - provider, status_code, quiet_exception_mapping -): +def test_a_provider_without_a_handler_maps_by_the_upstream_status(provider, status_code, quiet_exception_mapping): expected_class, expected_status = STATUS_KEYED[status_code] with pytest.raises(openai.APIError) as raised: @@ -1015,9 +974,7 @@ def test_an_unmapped_exception_with_no_model_or_provider_is_a_connection_error(q assert "boom" in raised.value.message -def _raise_and_map( - model: str | None, original_exception: Exception, custom_llm_provider: str | None -) -> None: +def _raise_and_map(model: str | None, original_exception: Exception, custom_llm_provider: str | None) -> None: """Calls exception_type() from inside the except block, as litellm/main.py does, so traceback.format_exc() has a real stack.""" try: @@ -1058,9 +1015,7 @@ def test_an_unmapped_exception_with_no_model_or_provider_message_keeps_traceback CONTEXT_WINDOW_MESSAGE = "This model's maximum context length is 4096 tokens." -CONTENT_POLICY_MESSAGE = ( - '{"error": {"type": "invalid_request_error", "code": "content_policy_violation"}}' -) +CONTENT_POLICY_MESSAGE = '{"error": {"type": "invalid_request_error", "code": "content_policy_violation"}}' TIMEOUT_MESSAGE = "Request timed out." PROVIDERS_THAT_RECOGNISE_A_FULL_CONTEXT_WINDOW = ( @@ -1103,15 +1058,11 @@ class _UpstreamErrorWithMessage(_UpstreamHTTPError): super().__init__(status_code=status_code) self.args = (message,) self.message = message - self.response = httpx.Response( - status_code=status_code, request=self.request, text=message - ) + self.response = httpx.Response(status_code=status_code, request=self.request, text=message) @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it( - provider, quiet_exception_mapping -): +def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it(provider, quiet_exception_mapping): if provider in PROVIDERS_THAT_RECOGNISE_A_FULL_CONTEXT_WINDOW: expected_class, expected_status = litellm.ContextWindowExceededError, 400 else: @@ -1129,9 +1080,7 @@ def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it( @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it( - provider, quiet_exception_mapping -): +def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it(provider, quiet_exception_mapping): if provider in PROVIDERS_THAT_RECOGNISE_A_CONTENT_POLICY_BLOCK: expected_class, expected_status = litellm.ContentPolicyViolationError, 400 else: @@ -1149,9 +1098,7 @@ def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it( @pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER) -def test_a_timed_out_request_is_a_timeout_for_every_provider( - provider, quiet_exception_mapping -): +def test_a_timed_out_request_is_a_timeout_for_every_provider(provider, quiet_exception_mapping): with pytest.raises(litellm.Timeout) as raised: exception_type( model="test-model", @@ -1409,3 +1356,97 @@ def test_bedrock_timeout_mapping_keeps_retry_after_readable(status_code): exception_headers = _get_response_headers(original_exception=exc_info.value) assert exception_headers is not None assert litellm.utils._get_retry_after_from_exception_header(response_headers=exception_headers) == 7 + + +_GUARDRAIL_BLOCK_ERROR = { + "message": "Content blocked: secret_project_codename pattern detected", + "param": "None", + "code": "400", + "provider_specific_fields": { + "error": "Content blocked: secret_project_codename pattern detected", + "pattern": "secret_project_codename", + "guardrail_name": "block-secret-project", + "guardrail_mode": "pre_call", + }, +} + + +def _openai_handler_error( + error_type: str, + headers: dict[str, str] | list[tuple[str, str]], + status_code: int = 400, + message: str = _GUARDRAIL_BLOCK_ERROR["message"], +) -> OpenAIError: + wire_error = {**_GUARDRAIL_BLOCK_ERROR, "type": error_type, "code": str(status_code), "message": message} + return OpenAIError( + status_code=status_code, + message=f"Error code: {status_code} - {{'error': {wire_error}}}", + headers=httpx.Headers(headers), + body=wire_error, + ) + + +_PROXY_HEADERS = {"x-litellm-call-id": "call-guardrail", "x-litellm-applied-guardrails": "block-secret-project"} + + +@pytest.mark.parametrize(("error_type", "status_code"), [("None", 400), ("invalid_request_error", 400), ("None", 422)]) +def test_litellm_proxy_guardrail_block_keeps_body_and_headers(error_type: str, status_code: int): + with pytest.raises(litellm.BadRequestError) as exc_info: + exception_type( + model="claude-haiku-4-5", + original_exception=_openai_handler_error(error_type, _PROXY_HEADERS, status_code=status_code), + custom_llm_provider="litellm_proxy", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert exc_info.value.body["provider_specific_fields"]["guardrail_name"] == "block-secret-project" + assert exc_info.value.body["type"] == error_type + assert dict(exc_info.value.response.headers) == _PROXY_HEADERS + + +@pytest.mark.parametrize("relayed_class", [litellm.BadRequestError, litellm.ContentPolicyViolationError]) +def test_litellm_proxy_relayed_litellm_error_keeps_body_and_headers(relayed_class: type[litellm.BadRequestError]): + message = f"litellm.{relayed_class.__name__}: {_GUARDRAIL_BLOCK_ERROR['message']}" + + with pytest.raises(relayed_class) as exc_info: + exception_type( + model="claude-haiku-4-5", + original_exception=_openai_handler_error("None", _PROXY_HEADERS, message=message), + custom_llm_provider="litellm_proxy", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert type(exc_info.value) is relayed_class + assert exc_info.value.body["provider_specific_fields"]["guardrail_name"] == "block-secret-project" + assert dict(exc_info.value.response.headers) == _PROXY_HEADERS + + +def test_openai_compatible_vendor_400_keeps_body_but_not_headers(): + with pytest.raises(litellm.BadRequestError) as exc_info: + exception_type( + model="gpt-5.4-mini", + original_exception=_openai_handler_error("vendor_specific_error", {"openai-organization": "org-1"}), + custom_llm_provider="openai", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert exc_info.value.body["type"] == "vendor_specific_error" + assert not exc_info.value.response.headers + + +def test_litellm_proxy_repeated_response_header_keeps_each_value(): + repeated = [("x-litellm-call-id", "call-guardrail"), ("set-cookie", "a=1"), ("set-cookie", "b=2")] + + with pytest.raises(litellm.BadRequestError) as exc_info: + exception_type( + model="claude-haiku-4-5", + original_exception=_openai_handler_error("None", repeated), + custom_llm_provider="litellm_proxy", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert exc_info.value.response.headers.multi_items() == repeated diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 9fe56f4dc65..7522e9a62e5 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -13,6 +13,7 @@ import pytest from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey from litellm.llms.anthropic.chat.guardrail_translation.handler import ( AnthropicMessagesHandler, @@ -635,14 +636,19 @@ class TestAnthropicMessagesHandlerInputProcessing: await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) assert guardrail.inputs is not None - assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"] + assert guardrail.inputs["texts"] == [ + "trusted top-level system prompt", + "safe text", + "prohibited correction", + ] structured = guardrail.inputs["structured_messages"] assert [m["role"] for m in structured] == ["system", "user", "system"] assert structured[0]["content"] == "trusted top-level system prompt" + assert data["system"] == "trusted top-level system prompt" assert data["messages"][1]["content"] == "[MASKED]" @pytest.mark.asyncio - async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included( + async def test_bedrock_masking_slice_lines_up_when_top_level_system_is_included( self, ): from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( @@ -668,25 +674,25 @@ class TestAnthropicMessagesHandlerInputProcessing: structured = guardrail.inputs["structured_messages"] bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1") - assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1 + assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) latest_user_index = bedrock._find_latest_message_index(structured, target_role="user") - assert ( - bedrock._locate_message_texts_slice( - structured_messages=structured, - target_index=latest_user_index, - texts=texts, - ) - is None - ) - assert ( - bedrock._merge_masked_texts( - masked_texts=["{MASKED}"], - texts=texts, - scanned_slice=None, - scanned_role_subset=True, - ) - == texts + scanned_slice = bedrock._locate_message_texts_slice( + structured_messages=structured, + target_index=latest_user_index, + texts=texts, ) + assert scanned_slice == (3, 1) + assert bedrock._merge_masked_texts( + masked_texts=["{MASKED}"], + texts=texts, + scanned_slice=scanned_slice, + scanned_role_subset=True, + ) == [ + "trusted top-level system prompt", + "safe text", + "prohibited correction", + "{MASKED}", + ] @pytest.mark.asyncio @pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None]) @@ -1611,7 +1617,8 @@ class TestAnthropicMessagesIncrementalScan: ) assert mock_api.call_count == 1 assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ - "What is the capital of France?" + "You are a helpful geography assistant.", + "What is the capital of France?", ] mock_api.reset_mock() await handler.process_input_messages( @@ -2150,6 +2157,213 @@ class TestAnthropicMessagesScanOnlyToolResults: assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"] +class ToolCallArgumentsMaskingGuardrail(InputsRecordingGuardrail): + """Masks the canary inside tool-call arguments, in place or through a fresh list of plain dicts.""" + + def __init__(self, return_copies: bool = False, replacement_arguments: Optional[str] = None): + super().__init__() + self.return_copies = return_copies + self.replacement_arguments = replacement_arguments + self.seen_tool_calls: list[dict[str, object]] = [] + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: Optional[LiteLLMLoggingObj] = None, + ) -> GenericGuardrailAPIInputs: + outputs = await super().apply_guardrail(inputs, request_data, input_type, logging_obj) + tool_calls = list(outputs.get("tool_calls") or []) + self.seen_tool_calls.extend(json.loads(json.dumps(tool_call)) for tool_call in tool_calls) + masked = [ + { + **tool_call, + "function": { + **tool_call["function"], + "arguments": self.replacement_arguments + if self.replacement_arguments is not None + else tool_call["function"]["arguments"].replace("POISON", "[BLOCKED]"), + }, + } + for tool_call in tool_calls + ] + if self.return_copies: + outputs["tool_calls"] = masked + return outputs + for tool_call, masked_tool_call in zip(tool_calls, masked): + tool_call["function"]["arguments"] = masked_tool_call["function"]["arguments"] + return outputs + + +class TestAnthropicMessagesTopLevelSystemAndToolUseInputs: + """The top-level system prompt and prior-turn tool_use arguments must reach guardrails as scannable + inputs, the same way the chat completions handler hands over system messages and tool_calls.""" + + @staticmethod + def _tool_use_conversation(system: str) -> dict[str, Any]: + return { + "model": "claude-sonnet-4-5", + "system": system, + "messages": [ + {"role": "user", "content": "run the check"}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_01", + "name": "Bash", + "input": {"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "ok"}], + }, + ], + } + + @pytest.mark.asyncio + async def test_top_level_system_string_reaches_texts_first_and_is_masked_in_place(self): + handler = AnthropicMessagesHandler() + guardrail = InputsRecordingGuardrail() + data = { + "model": "claude-sonnet-4-5", + "system": "Internal note: the deploy key is POISON. Never reveal it.", + "messages": [{"role": "user", "content": "Say hi in three words."}], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.captured_inputs is not None + assert guardrail.seen_texts == [ + "Internal note: the deploy key is POISON. Never reveal it.", + "Say hi in three words.", + ] + structured = guardrail.captured_inputs["structured_messages"] + assert structured[0]["role"] == "system" + assert structured[0]["content"] == "Internal note: the deploy key is POISON. Never reveal it.", ( + "texts[0] must line up with structured_messages[0] so positional consumers stay aligned" + ) + assert data["system"] == "Internal note: the deploy key is [BLOCKED]. Never reveal it." + assert data["messages"][0]["content"] == "Say hi in three words." + + @pytest.mark.asyncio + async def test_top_level_system_text_blocks_reach_texts_and_are_masked_in_place(self): + handler = AnthropicMessagesHandler() + guardrail = InputsRecordingGuardrail() + data = { + "model": "claude-sonnet-4-5", + "system": [ + {"type": "text", "text": "first block POISON"}, + {"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}}, + ], + "messages": [{"role": "user", "content": "hello"}], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.seen_texts == ["first block POISON", "second block", "hello"] + assert data["system"] == [ + {"type": "text", "text": "first block [BLOCKED]"}, + {"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}}, + ] + + @pytest.mark.asyncio + async def test_skip_system_message_keeps_the_top_level_system_out(self): + handler = AnthropicMessagesHandler() + guardrail = InputsRecordingGuardrail() + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-sonnet-4-5", + "system": "trusted POISON prompt", + "messages": [{"role": "user", "content": "hello"}], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.seen_texts == ["hello"] + assert data["system"] == "trusted POISON prompt" + + @pytest.mark.asyncio + async def test_prior_turn_tool_use_input_reaches_tool_calls_in_openai_shape(self): + handler = AnthropicMessagesHandler() + guardrail = InputsRecordingGuardrail() + data = self._tool_use_conversation(system="You are a careful agent harness.") + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.captured_inputs is not None + tool_calls = guardrail.captured_inputs.get("tool_calls") + assert tool_calls is not None and len(tool_calls) == 1 + assert tool_calls[0]["id"] == "toolu_01" + assert tool_calls[0]["type"] == "function" + assert tool_calls[0]["function"]["name"] == "Bash" + assert json.loads(tool_calls[0]["function"]["arguments"]) == { + "cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity" + } + assert data["messages"][1]["content"][0]["input"] == { + "cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity" + }, "a guardrail that leaves tool_calls alone must leave the tool_use input alone" + + @pytest.mark.asyncio + @pytest.mark.parametrize("return_copies", [False, True]) + async def test_masked_tool_call_arguments_write_back_into_the_tool_use_input(self, return_copies: bool): + handler = AnthropicMessagesHandler() + guardrail = ToolCallArgumentsMaskingGuardrail(return_copies=return_copies) + data = self._tool_use_conversation(system="You are a careful agent harness.") + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [tool_call["function"]["name"] for tool_call in guardrail.seen_tool_calls] == ["Bash"] + tool_use = data["messages"][1]["content"][0] + assert tool_use == { + "type": "tool_use", + "id": "toolu_01", + "name": "Bash", + "input": {"cmd": "AWS_ACCESS_KEY_ID=[BLOCKED] aws sts get-caller-identity"}, + } + assert data["messages"][2]["content"][0]["tool_use_id"] == "toolu_01" + + @pytest.mark.asyncio + async def test_non_json_rewritten_arguments_are_rejected_by_name(self): + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite + + handler = AnthropicMessagesHandler() + guardrail = ToolCallArgumentsMaskingGuardrail(replacement_arguments="[REDACTED]") + data = self._tool_use_conversation(system="Internal note: the deploy key is POISON. Never reveal it.") + data["messages"][2]["content"][0]["content"] = "fetched POISON page" + original = json.loads(json.dumps(data)) + + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert excinfo.value.guardrail_name == "scan-only-capture" + assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched" + assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched" + + @pytest.mark.asyncio + async def test_scan_only_tool_results_keeps_system_and_tool_use_out(self): + handler = AnthropicMessagesHandler() + guardrail = InputsRecordingGuardrail() + guardrail.scan_only_tool_results = True + data = self._tool_use_conversation(system="trusted POISON prompt") + data["messages"][2]["content"][0]["content"] = "fetched POISON page" + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.seen_texts == ["fetched POISON page"] + assert guardrail.captured_inputs is not None + assert guardrail.captured_inputs.get("tool_calls") is None + assert data["system"] == "trusted POISON prompt" + assert data["messages"][1]["content"][0]["input"] == { + "cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity" + } + assert data["messages"][2]["content"][0]["content"] == "fetched [BLOCKED] page" + + class TestStructuredWriteBackKeepsToolResults: """A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103).""" @@ -2272,6 +2486,116 @@ class TestAnthropicMessagesHandlerStreamingScanKey: assert ended_key != open_key +class PerRowTextGuardrail(CustomGuardrail): + """Answers one redacted text per chat row it was shown, the way a guardrail + that scans per message does, and hands back only texts.""" + + def __init__(self): + super().__init__(guardrail_name="per-row-redactor") + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + rows = inputs.get("structured_messages") or [] + return {**inputs, "texts": [str(row.get("content")).replace("123-45-6789", "") for row in rows]} + + +class PerSlotTextGuardrail(CustomGuardrail): + """Answers one redacted text per text slot of every chat row it was shown, the + way a guardrail that counts slots per message does, and hands back only texts.""" + + def __init__(self): + super().__init__(guardrail_name="per-slot-redactor") + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts + + rows = inputs.get("structured_messages") or [] + return { + **inputs, + "texts": [text.replace("123-45-6789", "") for row in rows for text in message_slot_texts(row)], + } + + +class TestPerMessageTextWriteBack: + """Texts that no longer pair one-to-one with what the handler extracted must be + rejected by name instead of sliding onto the wrong messages.""" + + @pytest.mark.asyncio + async def test_one_text_per_row_over_a_system_prompt_is_applied(self): + data = { + "model": "claude-sonnet-4-5", + "system": "Reply with exactly the SSN you were given.", + "messages": [{"role": "user", "content": "My SSN is 123-45-6789."}], + } + + await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail()) + + assert data["system"] == "Reply with exactly the SSN you were given." + assert data["messages"] == [{"role": "user", "content": "My SSN is ."}] + + @pytest.mark.asyncio + async def test_one_text_per_row_over_a_multi_block_system_prompt_is_rejected_by_name(self): + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite + + data = { + "model": "claude-sonnet-4-5", + "system": [ + {"type": "text", "text": "Reply with exactly the SSN you were given."}, + {"type": "text", "text": "Never apologize."}, + ], + "messages": [{"role": "user", "content": "My SSN is 123-45-6789."}], + } + original = json.loads(json.dumps(data)) + + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail()) + + assert excinfo.value.guardrail_name == "per-row-redactor" + assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched" + assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched" + + @pytest.mark.asyncio + async def test_one_text_per_slot_over_a_system_prompt_with_an_empty_block_is_applied(self): + data = { + "model": "claude-sonnet-4-5", + "system": [ + {"type": "text", "text": ""}, + {"type": "text", "text": "Reply with exactly the SSN you were given."}, + ], + "messages": [{"role": "user", "content": "My SSN is 123-45-6789."}], + } + + await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerSlotTextGuardrail()) + + assert data["system"] == [ + {"type": "text", "text": ""}, + {"type": "text", "text": "Reply with exactly the SSN you were given."}, + ] + assert data["messages"] == [{"role": "user", "content": "My SSN is ."}] + + @pytest.mark.asyncio + async def test_one_text_per_row_without_a_system_prompt_is_applied(self): + data = { + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "My SSN is 123-45-6789."}], + } + + await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail()) + + assert data["messages"] == [{"role": "user", "content": "My SSN is ."}] + + class TestAnthropicMessagesHandlerPostCallHookResponse: def test_openai_shaped_stream_assembly_reaches_the_hook_as_a_messages_response(self): from litellm.types.utils import Choices, Message, ModelResponse, Usage diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py index 29e9279731d..a20aaf2e324 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py @@ -8,6 +8,8 @@ Without the fix, the AnthropicStreamWrapper silently dropped these arguments, causing tool_use blocks to arrive with empty input {}. """ +import json + from typing import List from unittest.mock import MagicMock @@ -139,9 +141,7 @@ async def test_async_stream_emits_input_json_delta_for_bundled_tool_args(): # Verify the delta carries the tool arguments delta_event = events[input_json_delta_idx] - assert delta_event["delta"][ - "partial_json" - ], "input_json_delta should have non-empty partial_json" + assert json.loads(delta_event["delta"]["partial_json"]) == {"location": "Boston"} @pytest.mark.asyncio @@ -300,7 +300,7 @@ def test_sync_stream_emits_input_json_delta_for_bundled_tool_args(): assert ( input_json_delta_idx == tool_start_idx + 1 ), "input_json_delta should immediately follow the tool_use content_block_start" - assert events[input_json_delta_idx]["delta"]["partial_json"] + assert json.loads(events[input_json_delta_idx]["delta"]["partial_json"]) == {"location": "Boston"} def test_sync_stream_no_extra_delta_when_tool_args_empty(): diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_prompt_cache_prediction.py b/tests/test_litellm/llms/anthropic/test_anthropic_prompt_cache_prediction.py new file mode 100644 index 00000000000..2b36866a1a0 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/test_anthropic_prompt_cache_prediction.py @@ -0,0 +1,209 @@ +import json +from collections.abc import Mapping +from datetime import datetime +from types import SimpleNamespace +from typing import Final + +import httpx +import pytest +import respx +from pydantic import JsonValue + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.caching.llm_caching_handler import LLMClientCache +from litellm.llms.anthropic.count_tokens import handler as count_handler +from litellm.llms.anthropic.experimental_pass_through.messages.transformation import DEFAULT_ANTHROPIC_API_VERSION +from litellm.llms.anthropic.prompt_cache_prediction import ( + NativePredictionTarget, + cache_scope, + count_prompt_tokens, + parse_observed_cache, + parse_prompt, + resolve_prediction_target, + supported_prediction_headers, +) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.models.credentials import CredentialItem +from litellm.proxy import proxy_server +from litellm.proxy.hooks.prompt_cache_prediction import PromptCacheObserver, lookup +from litellm.proxy.management_endpoints.prompt_cache_prediction import predict_arm +from litellm.proxy.utils import InternalUsageCache +from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo +from litellm.types.utils import CacheCreationTokenDetails, ModelResponse, PromptTokensDetailsWrapper, Usage + +_MODEL: Final = "claude-sonnet-5" +_KEY: Final = "test-provider-key" +_CALLER: Final = "test-caller-hash" +_DEPLOYMENT: Final = "test-native-deployment" + + +def _body() -> dict[str, JsonValue]: + return { + "model": _MODEL, + "system": "Keep this context", + "tools": [{"name": "lookup", "input_schema": {"type": "object"}}], + "messages": [{"role": "user", "content": [ + {"type": "text", "text": "A cacheable prefix", "cache_control": {"type": "ephemeral"}} + ]}], + } + + +@pytest.mark.parametrize("version", [None, "2099-01-01", DEFAULT_ANTHROPIC_API_VERSION]) +@pytest.mark.asyncio +async def test_observer_records_only_version_supported_by_token_counter(version: str | None) -> None: + cache: Final = DualCache() + observer: Final = PromptCacheObserver(InternalUsageCache(dual_cache=cache), clock=lambda: 1010.0) + body: Final = _body() + prefix: Final = parse_prompt(body) + assert prefix is not None + headers: Final = {"x-api-key": _KEY, **({"anthropic-version": version} if version is not None else {})} + wire: Final = httpx.Request("POST", "https://api.anthropic.com/v1/messages", headers=headers, json=body) + response: Final = ModelResponse( + model=_MODEL, + usage=Usage( + prompt_tokens=311, + completion_tokens=2, + total_tokens=313, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=100, + cache_creation_tokens=200, + cache_creation_token_details=CacheCreationTokenDetails( + ephemeral_5m_input_tokens=200, ephemeral_1h_input_tokens=0 + ), + ), + ), + ) + await observer.async_log_success_event( + { + "call_type": "anthropic_messages", + "custom_llm_provider": "anthropic", + "httpx_response": httpx.Response(200, request=wire), + "first_api_call_start_time": datetime.fromtimestamp(1000.0), + "standard_logging_object": { + "status": "success", "model_id": _DEPLOYMENT, + "metadata": {"user_api_key_hash": _CALLER}, + }, + }, + response, + datetime.fromtimestamp(1010.0), + datetime.fromtimestamp(1010.0), + ) + default_scope: Final = cache_scope(_CALLER, _DEPLOYMENT, _KEY, _MODEL) + found: Final = await lookup(cache, default_scope, prefix, now=1010.0) + assert (found is not None) == (version == DEFAULT_ANTHROPIC_API_VERSION) + if version != DEFAULT_ANTHROPIC_API_VERSION: + other_scope: Final = cache_scope(_CALLER, _DEPLOYMENT, _KEY, _MODEL, version or "") + assert await lookup(cache, other_scope, prefix, now=1010.0) is None + + +@pytest.mark.parametrize("headers, supported", [ + ({}, True), + ({"Anthropic-Version": DEFAULT_ANTHROPIC_API_VERSION}, True), + ({"anthropic-version": "2099-01-01"}, False), + ({"Anthropic-Beta": ""}, False), + ({"anthropic-beta": "future-feature"}, False), +]) +def test_prediction_header_eligibility(headers: Mapping[str, str], supported: bool) -> None: + assert supported_prediction_headers(headers) is supported + + +@pytest.mark.asyncio +async def test_provider_count_uses_same_version_and_preserves_native_input(monkeypatch: pytest.MonkeyPatch) -> None: + body: Final = _body() + requests: Final[list[httpx.Request]] = [] + + def provider(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json={"input_tokens": 311}) + + client: Final = AsyncHTTPHandler() + await client.client.aclose() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(provider)) + monkeypatch.setattr(count_handler, "get_async_httpx_client", lambda **kwargs: client) + try: + assert await count_prompt_tokens(_MODEL, _KEY, body) == 311 + finally: + await client.client.aclose() + assert len(requests) == 1 + assert requests[0].headers["anthropic-version"] == DEFAULT_ANTHROPIC_API_VERSION + assert requests[0].url == "https://api.anthropic.com/v1/messages/count_tokens" + assert json.loads(requests[0].content) == body + + +@pytest.mark.parametrize("source", ["static", "database"]) +@pytest.mark.asyncio +async def test_environment_credential_matches_native_count_and_observed_scope( + source: str, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("LIT7658_PROVIDER_KEY", _KEY) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + params: Final = { + "model": f"anthropic/{_MODEL}", "api_key": "os.environ/LIT7658_PROVIDER_KEY", + "api_base": "https://api.anthropic.com", + } + router: Final = litellm.Router(model_list=[{ + "model_name": "test-native", "litellm_params": dict(params), "model_info": {"id": _DEPLOYMENT}, + }] if source == "static" else [], num_retries=0) + if source == "database": + monkeypatch.setattr(proxy_server, "llm_router", router) + assert proxy_server.ProxyConfig()._add_deployment([SimpleNamespace( + model_id=_DEPLOYMENT, model_name="test-native", model_info={}, litellm_params=dict(params), + )]) == 1 + deployment: Final = router.get_deployment(_DEPLOYMENT) + assert deployment is not None + target: Final = resolve_prediction_target(deployment.litellm_params) + assert isinstance(target, NativePredictionTarget) + body: Final = _body() + with respx.mock() as upstream: + native: Final = upstream.post("https://api.anthropic.com/v1/messages").respond(200, json={ + "id": "msg_test", "type": "message", "role": "assistant", "model": _MODEL, + "content": [{"type": "text", "text": "Hello"}], "stop_reason": "end_turn", "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 1, "cache_read_input_tokens": 300}, + }) + counter: Final = upstream.post("https://api.anthropic.com/v1/messages/count_tokens").respond( + 200, json={"input_tokens": 311}, + ) + await router.aanthropic_messages( + model="test-native", max_tokens=1, **{key: value for key, value in body.items() if key != "model"}, + ) + assert await count_prompt_tokens(target.model, target.api_key, body) == 311 + assert native.call_count == counter.call_count == 1 + assert native.calls.last.request.headers["x-api-key"] == counter.calls.last.request.headers["x-api-key"] == _KEY + observed: Final = parse_observed_cache(native.calls.last.request, ModelResponse( + model=_MODEL, usage=Usage( + prompt_tokens=311, completion_tokens=1, total_tokens=312, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=300), + ), + ), _CALLER, _DEPLOYMENT) + assert observed is not None + assert observed.scope == cache_scope(_CALLER, _DEPLOYMENT, target.api_key, target.model) + + +@pytest.mark.parametrize("inline_key", [None, _KEY]) +@pytest.mark.asyncio +async def test_named_credential_is_explicitly_unsupported_before_count( + inline_key: str | None, monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "credential_list", [CredentialItem( + credential_name="test-named", credential_info={}, credential_values={"api_key": "test-named-provider-key"}, + )]) + deployment: Final = Deployment( + model_name="test-native", + litellm_params=LiteLLM_Params( + model=f"anthropic/{_MODEL}", api_key=inline_key, litellm_credential_name="test-named", + ), + model_info=ModelInfo(id=_DEPLOYMENT), + ) + body: Final = _body() + prefix: Final = parse_prompt(body) + assert prefix is not None + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + pytest.fail("Unsupported named credentials must not reach provider counting") + + arm: Final = await predict_arm(deployment, body, prefix, _CALLER, DualCache(), count) + assert arm.cache_state == "unknown" + assert arm.reason == "unsupported_deployment_configuration" + assert arm.estimate is None and arm.cold is None and arm.warm is None diff --git a/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py b/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py index 3ea840519f9..cbb69b4ceed 100644 --- a/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py +++ b/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py @@ -32,9 +32,14 @@ action. import base64 import json from datetime import datetime, timedelta, timezone +from types import MappingProxyType +from typing import Final from unittest.mock import MagicMock, patch import pytest +from pydantic import TypeAdapter + +from litellm.llms.bedrock.base_aws_llm import WebIdentitySessionPolicy, _SessionPolicyStatement # Actions the Claude Platform on AWS service is documented to call. # Source: AWS IAM action reference + the #27678 surface area. @@ -49,9 +54,9 @@ _CLAUDE_PLATFORM_ACTIONS = { } -def _captured_policy() -> dict: - """Run _auth_with_web_identity_token under mocks + return the parsed - Policy dict that was actually sent to STS.""" +def _captured_policy_document() -> str: + """Run _auth_with_web_identity_token under mocks + return the Policy + JSON document that was actually sent to STS.""" from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM base = BaseAWSLLM() @@ -84,11 +89,21 @@ def _captured_policy() -> dict: mock_sts.assume_role_with_web_identity.assert_called_once() kwargs = mock_sts.assume_role_with_web_identity.call_args.kwargs - policy_str = kwargs["Policy"] - return json.loads(policy_str) + return kwargs["Policy"] -def _statement_by_sid(policy: dict, sid: str) -> dict: +_SESSION_POLICY_ADAPTER: Final = TypeAdapter(WebIdentitySessionPolicy) + + +def _captured_policy() -> WebIdentitySessionPolicy: + return _SESSION_POLICY_ADAPTER.validate_python(json.loads(_captured_policy_document())) + + +def _granted_actions(policy: WebIdentitySessionPolicy) -> frozenset[str]: + return frozenset(action for stmt in policy["Statement"] for action in stmt["Action"]) + + +def _statement_by_sid(policy: WebIdentitySessionPolicy, sid: str) -> _SessionPolicyStatement: for stmt in policy["Statement"]: if stmt.get("Sid") == sid: return stmt @@ -102,7 +117,6 @@ class TestWebIdentitySessionPolicyShape: def test_policy_parses_as_valid_iam_document(self): policy = _captured_policy() assert policy["Version"] == "2012-10-17" - assert isinstance(policy["Statement"], list) assert len(policy["Statement"]) >= 2 def test_bedrock_statement_actions_preserved(self): @@ -137,16 +151,7 @@ class TestClaudePlatformActionsCovered: @pytest.mark.parametrize("action", sorted(_CLAUDE_PLATFORM_ACTIONS)) def test_claude_platform_action_present(self, action: str): - policy = _captured_policy() - # Action may live in any Statement — search across all. - all_actions: set = set() - for stmt in policy["Statement"]: - stmt_actions = stmt.get("Action") - if isinstance(stmt_actions, str): - all_actions.add(stmt_actions) - elif isinstance(stmt_actions, list): - all_actions.update(stmt_actions) - assert action in all_actions, ( + assert action in _granted_actions(_captured_policy()), ( f"{action} missing from session policy — " f"bedrock/claude_platform/* requests will 403 on OIDC auth" ) @@ -179,15 +184,7 @@ class TestBedrockMantleActionsCovered: action" even when the role's identity policy grants it.""" def test_bedrock_mantle_create_inference_present(self): - policy = _captured_policy() - all_actions: set = set() - for stmt in policy["Statement"]: - stmt_actions = stmt.get("Action") - if isinstance(stmt_actions, str): - all_actions.add(stmt_actions) - elif isinstance(stmt_actions, list): - all_actions.update(stmt_actions) - assert "bedrock-mantle:CreateInference" in all_actions, ( + assert "bedrock-mantle:CreateInference" in _granted_actions(_captured_policy()), ( "bedrock-mantle:CreateInference missing from session policy — " "bedrock_mantle/* requests will 403 on OIDC/WIF auth" ) @@ -233,7 +230,7 @@ class TestInvalidIdentityTokenSurfacesAudience: operator can diagnose the mismatch without enabling LITELLM_LOG=DEBUG on a prod instance.""" - _AUD = "https://guidepoint.litellm-prod.ai" + _AUD = "https://gateway.example.com" _ISS = "https://accounts.google.com" _STS_MESSAGE = ( "An error occurred (InvalidIdentityToken) when calling the " @@ -308,3 +305,44 @@ class TestPolicyTransportConditions: "ClaudePlatformLiteLLM must require aws:SecureTransport=true " "to keep parity with the bedrock statement" ) + + +_STS_SESSION_POLICY_PLAINTEXT_LIMIT: Final = 2048 + +_BEDROCK_ROUTE_ACTIONS: Final = MappingProxyType( + { + "model/{model_id}/invoke": "bedrock:InvokeModel", + "model/{model_id}/invoke-with-response-stream": "bedrock:InvokeModelWithResponseStream", + "model/{model_id}/converse": "bedrock:InvokeModel", + "model/{model_id}/converse-stream": "bedrock:InvokeModelWithResponseStream", + "model/{model_id}/count-tokens": "bedrock:CountTokens", + "guardrail/{guardrail_id}/version/{version}/apply": "bedrock:ApplyGuardrail", + "rerank": "bedrock:Rerank", + "knowledgebases/{knowledge_base_id}/retrieve": "bedrock:Retrieve", + "knowledgebases": "bedrock:ListKnowledgeBases", + "agents/{agent_id}/agentAliases/{alias_id}/sessions/{session_id}/text": "bedrock:InvokeAgent", + "runtimes/{agent_runtime_arn}/invocations": "bedrock-agentcore:InvokeAgentRuntime", + "runtimes/{agent_runtime_arn}/invocations with X-Amzn-Bedrock-AgentCore-Runtime-User-Id": ( + "bedrock-agentcore:InvokeAgentRuntimeForUser" + ), + "mcp": "bedrock-agentcore:InvokeGateway", + } +) + + +class TestSessionPolicyGrantsEveryBedrockRoute: + """LIT-7348: ``/rerank`` authorizes against ``bedrock:Rerank``, which the + ceiling never granted, so rerank 403d on web identity auth while static + credentials and IRSA worked. Each route the bedrock package signs with the + web identity session maps to the IAM action it authorizes against, and the + ceiling must grant every one of them.""" + + @pytest.mark.parametrize(("route", "action"), sorted(_BEDROCK_ROUTE_ACTIONS.items())) + def test_route_action_is_granted_by_the_ceiling(self, route: str, action: str): + assert action in _granted_actions(_captured_policy()), ( + f"/{route} authorizes against {action}, which the session policy does not grant, " + "so it 403s on web identity auth" + ) + + def test_policy_document_fits_the_sts_plaintext_limit(self): + assert len(_captured_policy_document()) <= _STS_SESSION_POLICY_PLAINTEXT_LIMIT diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 40566261c84..a7aefa714aa 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -773,6 +773,12 @@ class TestBedrockMantleCodexAdditionalTools: assert body["input"] == codex_agentic_items assert "tools" not in body + def test_input_without_additional_tools_sanitizes_tools_on_the_caller_params_object(self): + params = {"tools": [{"type": "function", "name": "wait", "parameters": '{"type": "object"}'}]} + body = self._transform(input=[self._USER_MESSAGE], params=params) + assert body["tools"][0]["parameters"] == {"type": "object"} + assert params["tools"][0]["parameters"] == {"type": "object"} + def test_malformed_additional_tools_item_without_tools_list_is_stripped(self): body = self._transform( input=[ diff --git a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py index 2b3b6343fad..3eb4a70ee15 100644 --- a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py +++ b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py @@ -1,4 +1,6 @@ import json +from collections.abc import Mapping +from typing import cast from unittest.mock import MagicMock import pytest @@ -6,6 +8,7 @@ import pytest import litellm from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig +from litellm.types.llms.gemini import BidiGenerateContentServerMessage def test_gemini_realtime_transformation_session_created(): @@ -2178,3 +2181,71 @@ def test_unbilled_usage_on_session_close_flushes_trailing_audio(patch_gemini_tra } assert usage == expected assert config.unbilled_usage_on_session_close("gemini-3.5-transcribe-live") is None + + +def _grounded_live_frame(grounding_metadata: Mapping[str, object] | None) -> Mapping[str, object]: + """One Live server frame. Grounding metadata and usageMetadata arrive together, as Vertex sends them.""" + from typing import Final + + server_content: Final = { + "turnComplete": True, + **({} if grounding_metadata is None else {"groundingMetadata": grounding_metadata}), + } + return { + "serverContent": server_content, + "usageMetadata": { + "promptTokenCount": 19, + "candidatesTokenCount": 157, + "totalTokenCount": 176, + "promptTokensDetails": ({"modality": "TEXT", "tokenCount": 19},), + "candidatesTokensDetails": ({"modality": "AUDIO", "tokenCount": 157},), + }, + } + + +def _response_done_input_details(message: Mapping[str, object]) -> Mapping[str, object]: + """The ``input_tokens_details`` a ``response.done`` event carries, read off the emitted event.""" + from typing import Final + + config: Final = GeminiRealtimeConfig() + event: Final = config.transform_response_done_event( + message=cast( # cast-ok: a test fixture stands in for the server frame TypedDict + BidiGenerateContentServerMessage, message + ), + current_response_id="resp_grounding", + current_conversation_id="conv_grounding", + output_items=None, + ) + usage: Final = event["response"]["usage"] + assert usage, "response.done must carry a usage object" + return usage.get("input_tokens_details") or {} + + +def test_gemini_realtime_response_done_counts_web_grounding(): + """Regression: Live reports grounding in the server frames and never in usageMetadata. + + Nothing read those frames on the realtime path, so web_search_requests stayed unset and the + cost path's only trigger for Google's per-query grounding charge never fired. + + The counter is read off the emitted event, which is what the cost path is handed, so this covers + the grounding read and the usage bridge that carries it together + """ + input_details = _response_done_input_details( + _grounded_live_frame( + { + "webSearchQueries": ["who won the 2026 world cup final"], + "groundingChunks": [{"web": {"uri": "https://example.com"}}], + } + ) + ) + + assert input_details.get("web_search_requests") == 1, "a grounded turn must report its query" + assert input_details.get("text_tokens") == 19, "the modality breakdown must survive alongside it" + + +def test_gemini_realtime_response_done_reports_no_grounding_when_none_ran(): + """The counter must stay unset on an ordinary turn, or every session pays a grounding fee.""" + input_details = _response_done_input_details(_grounded_live_frame(None)) + + assert input_details.get("web_search_requests") is None + assert input_details.get("google_maps_grounding_requests") is None diff --git a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5a29a96829f..cb884fb7cc1 100644 --- a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -1893,6 +1893,183 @@ class TestScanOnlyToolResults: assert data["messages"][4]["content"] == "and then?" +class TestNoScannableContentRecordsNotRun: + """LIT-6314: a guardrail whose scoping leaves nothing to scan must still persist an evaluation record""" + + def _system_only_data(self) -> dict: + return {"messages": [{"role": "system", "content": "SYSTEM-PROMPT"}]} + + def _recorded_entries(self, data: dict) -> list: + metadata = data.get("metadata") or data.get("litellm_metadata") or {} + return metadata.get("standard_logging_guardrail_information") or [] + + @pytest.mark.asyncio + async def test_skipped_scan_records_not_run_entry(self): + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="skip-system-guardrail") + guardrail.skip_system_message_in_guardrail = True + data = self._system_only_data() + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is None, "nothing survived scoping, apply_guardrail must not run" + entries = self._recorded_entries(data) + assert len(entries) == 1 + assert entries[0]["guardrail_name"] == "skip-system-guardrail" + assert entries[0]["guardrail_status"] == "not_run" + assert entries[0]["guardrail_response"] == "no scannable content after message scoping" + + @pytest.mark.asyncio + @pytest.mark.parametrize("skip_system", [False, True]) + async def test_empty_content_does_not_blame_scoping(self, skip_system: bool): + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="unscoped-guardrail") + guardrail.skip_system_message_in_guardrail = skip_system + data = {"messages": [{"role": "user", "content": None}]} + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is None + entries = self._recorded_entries(data) + assert len(entries) == 1 + assert entries[0]["guardrail_status"] == "not_run" + assert entries[0]["guardrail_response"] == "no scannable content" + + @pytest.mark.asyncio + async def test_self_recording_guardrail_is_left_alone(self): + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="self-recording-guardrail") + guardrail.skip_system_message_in_guardrail = True + guardrail.records_own_guardrail_information = True + data = self._system_only_data() + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is None + assert self._recorded_entries(data) == [] + + @pytest.mark.asyncio + async def test_scannable_content_records_no_extra_entry(self): + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="normal-guardrail") + data = {"messages": [{"role": "user", "content": "hello"}]} + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is not None + assert all(e.get("guardrail_status") != "not_run" for e in self._recorded_entries(data)) + + @pytest.mark.asyncio + async def test_image_only_content_is_not_reported_as_not_run(self): + """Images are only scanned alongside text, so an image-only request is a + pre-existing scan gap, not a message-scoping skip, and must not be labelled one""" + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="image-guardrail") + guardrail.skip_system_message_in_guardrail = True + data = { + "messages": [ + {"role": "system", "content": "SYSTEM-PROMPT"}, + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}], + }, + ] + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert self._recorded_entries(data) == [] + + @pytest.mark.asyncio + async def test_scoped_out_image_only_message_is_not_reported_as_not_run(self): + """An image in a skipped role must behave like any other image-only request""" + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="image-guardrail") + guardrail.skip_system_message_in_guardrail = True + data = { + "messages": [ + { + "role": "system", + "content": [{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}], + }, + ] + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is None + assert self._recorded_entries(data) == [] + + @pytest.mark.asyncio + async def test_scoped_out_text_with_image_records_not_run(self): + """Scoping removed text too, so the skip is recorded even though an image sat beside it""" + handler = OpenAIChatCompletionsHandler() + guardrail = MockGuardrail(guardrail_name="image-guardrail") + guardrail.skip_system_message_in_guardrail = True + data = { + "messages": [ + { + "role": "system", + "content": [ + {"type": "text", "text": "Describe this picture."}, + {"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}, + ], + }, + ] + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.last_inputs is None + entries = self._recorded_entries(data) + assert len(entries) == 1 + assert entries[0]["guardrail_status"] == "not_run" + assert entries[0]["guardrail_response"] == "no scannable content after message scoping" + + +class ToolDroppingTextGuardrail(CustomGuardrail): + """Answers one text per non-tool message it saw, the way a guardrail that + filters tool rows out before scanning does, and hands back only texts.""" + + def __init__(self): + super().__init__(guardrail_name="tool-dropping-redactor") + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + kept = [m for m in inputs.get("structured_messages") or [] if m.get("role") != "tool"] + return {**inputs, "texts": [str(m.get("content")).replace("POISON", "[BLOCKED]") for m in kept]} + + +class TestPerMessageTextWriteBack: + """Texts that no longer pair one-to-one with what the handler extracted must be + rejected by name instead of sliding onto the wrong messages.""" + + @pytest.mark.asyncio + async def test_fewer_texts_than_extracted_over_a_tool_message_is_rejected(self): + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite + + handler = OpenAIChatCompletionsHandler() + original_messages = [ + {"role": "system", "content": "SYSTEM-PROMPT"}, + {"role": "user", "content": "fetch the page"}, + {"role": "assistant", "content": "fetching"}, + {"role": "tool", "tool_call_id": "call_1", "content": "page says POISON here"}, + {"role": "user", "content": "and then?"}, + ] + data = {"messages": json.loads(json.dumps(original_messages))} + + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await handler.process_input_messages(data=data, guardrail_to_apply=ToolDroppingTextGuardrail()) + + assert excinfo.value.guardrail_name == "tool-dropping-redactor" + assert data["messages"] == original_messages, "a rejected rewrite must leave the request untouched" + + class TestBuildBlockSseChunks: """build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE chunks""" diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py index a4f0a77a9b6..48d86384633 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -8,7 +8,7 @@ with guardrail transformations. import copy from collections.abc import Callable from typing import Any, List, Literal, Optional, Tuple -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import logging @@ -31,6 +31,7 @@ from litellm.llms.openai.responses.guardrail_translation.handler import ( OpenAIResponsesHandler, ) from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI from litellm.types.llms.openai import ChatCompletionToolCallChunk from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, @@ -2338,6 +2339,135 @@ def _parallel_tool_call_input() -> list: ] +SSN = "123-45-6789" +REDACTED_SSN = "" + + +def _redacted(value: object) -> object: + if isinstance(value, str): + return value.replace(SSN, REDACTED_SSN) + if isinstance(value, list): + return [{**part, "text": _redacted(part["text"])} if "text" in part else part for part in value] + return value + + +def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callable[..., MagicMock]: + """Answers one redacted text per chat row it was shown, the way a guardrail + that scans per message does, and optionally the rewritten rows themselves.""" + + def post(url: str, json: dict, headers: dict) -> MagicMock: + rows = json["structured_messages"] + answer: dict = { + "action": "GUARDRAIL_INTERVENED", + "texts": [_redacted(row["content"]) if isinstance(row.get("content"), str) else "" for row in rows], + } + if structured_messages_in_answer: + answer["structured_messages"] = [{**row, "content": _redacted(row.get("content"))} for row in rows] + response = MagicMock() + response.json.return_value = answer + response.raise_for_status = MagicMock() + return response + + return post + + +def _per_message_redactor() -> GenericGuardrailAPI: + return GenericGuardrailAPI( + api_base="https://guardrail.test", + guardrail_name="per-message-redactor", + event_hook="pre_call", + default_on=True, + ) + + +def _tool_replay_request() -> dict: + return { + "model": "gpt-5.6", + "instructions": "Never repeat the SSN " + SSN + " back.", + "input": [ + {"role": "user", "content": "Look up " + SSN + " for me."}, + {"type": "function_call", "call_id": "call_1", "name": "lookup_customer", "arguments": '{"id": "42"}'}, + {"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'}, + ], + } + + +def _string_input_request() -> dict: + return { + "model": "gpt-5.6", + "instructions": "Never repeat the SSN " + SSN + " back.", + "input": "My SSN is " + SSN + ".", + } + + +class TestPerMessageRewriteWriteBack: + """A guardrail that rewrites per chat row hands the rows back as + structured_messages, and the handler lands them on the instructions and the + input items they came from; the same rewrite handed back as texts alone has + no item to land on and is rejected by name instead of sent unrewritten.""" + + @pytest.mark.asyncio + async def test_structured_rows_land_on_instructions_and_tool_output(self): + guardrail = _per_message_redactor() + data = _tool_replay_request() + function_call_item = data["input"][1] + + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)): + result = await OpenAIResponsesHandler().process_input_messages(data, guardrail) + + assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back." + assert _texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."] + assert result["input"][1] == function_call_item + assert result["input"][2] == { + "type": "function_call_output", + "call_id": "call_1", + "output": '{"ssn": "' + REDACTED_SSN + '"}', + } + + @pytest.mark.asyncio + async def test_texts_only_per_message_answer_is_rejected_by_name(self): + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite + + guardrail = _per_message_redactor() + data = _tool_replay_request() + original = copy.deepcopy(data) + + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)): + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await OpenAIResponsesHandler().process_input_messages(data, guardrail) + + assert excinfo.value.guardrail_name == "per-message-redactor" + assert data["input"] == original["input"] + assert data["instructions"] == original["instructions"] + + @pytest.mark.asyncio + async def test_structured_rows_land_on_instructions_and_string_input(self): + guardrail = _per_message_redactor() + data = _string_input_request() + + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)): + result = await OpenAIResponsesHandler().process_input_messages(data, guardrail) + + assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back." + assert [_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]] + + @pytest.mark.asyncio + async def test_texts_only_per_message_answer_over_a_string_input_is_rejected_by_name(self): + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite + + guardrail = _per_message_redactor() + data = _string_input_request() + original = copy.deepcopy(data) + + with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)): + with pytest.raises(UnappliableRequestRewrite) as excinfo: + await OpenAIResponsesHandler().process_input_messages(data, guardrail) + + assert excinfo.value.guardrail_name == "per-message-redactor" + assert data["input"] == original["input"] + assert data["instructions"] == original["instructions"] + + class TestProvenancePatching: """The O(n) provenance pass must keep patching rewritten rows in place for the shapes real agent loops produce, and fall back safely everywhere else.""" diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py index 9c236d81f51..bbd0cdf97e3 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py @@ -133,7 +133,20 @@ def test_namespace_keeps_a_non_function_member_when_a_function_member_is_edited( assert merged[0]["tools"][1] == custom_member -def test_namespace_keeps_its_non_function_members_when_every_function_member_is_dropped(): +def test_namespace_keeps_its_custom_member_when_every_function_member_is_dropped(): + custom_member = {"type": "custom", "name": "grep", "description": "Grep", "format": {"type": "text"}} + original = [ + {"type": "namespace", "name": "ns", "description": "NS", "tools": [_function("read"), custom_member]}, + _function("a"), + ] + groups = _groups(original) + + merged = merge_guardrailed_tools(original, groups, [groups[0][1], groups[1][0]]) + + assert list(merged) == [{"type": "namespace", "name": "ns", "description": "NS", "tools": [custom_member]}, _function("a")] + + +def test_namespace_custom_member_is_dropped_when_the_guardrail_drops_its_chat_form(): custom_member = {"type": "custom", "name": "grep", "description": "Grep", "format": {"type": "text"}} original = [ {"type": "namespace", "name": "ns", "description": "NS", "tools": [_function("read"), custom_member]}, @@ -143,7 +156,37 @@ def test_namespace_keeps_its_non_function_members_when_every_function_member_is_ merged = merge_guardrailed_tools(original, groups, [groups[1][0]]) - assert list(merged) == [{"type": "namespace", "name": "ns", "description": "NS", "tools": [custom_member]}, _function("a")] + assert list(merged) == [_function("a")] + + +def test_custom_member_description_edit_lands_without_the_namespace_prefix_or_grammar_block(): + grammar = {"type": "grammar", "syntax": "lark", "definition": "start: X"} + custom_member = {"type": "custom", "name": "exec", "description": "Run a command", "format": grammar} + original = [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [custom_member]}] + groups = _groups(original) + assert groups[0][0]["function"]["description"] == "Shell\n\nRun a command\n\nFormat:\n```lark\nstart: X\n```" + edited = copy.deepcopy(_flat(groups)) + edited[0]["function"]["description"] = "Shell\n\nRun a command (guarded)\n\nFormat:\n```lark\nstart: X\n```" + + merged = merge_guardrailed_tools(original, groups, edited) + + guarded_member = {**custom_member, "description": "Run a command (guarded)"} + assert list(merged) == [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [guarded_member]}] + + +def test_text_appended_after_the_grammar_block_lands_on_the_member_without_the_block(): + grammar = {"type": "grammar", "syntax": "lark", "definition": "start: X"} + custom_member = {"type": "custom", "name": "exec", "description": "Run a command", "format": grammar} + original = [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [custom_member]}] + groups = _groups(original) + edited = copy.deepcopy(_flat(groups)) + edited[0]["function"]["description"] = edited[0]["function"]["description"] + " [checked]" + + merged = merge_guardrailed_tools(original, groups, edited) + + assert merged[0]["tools"][0]["description"] == "Run a command [checked]" + reflattened = _flat(_groups(merged)) + assert reflattened[0]["function"]["description"] == "Shell\n\nRun a command [checked]\n\nFormat:\n```lark\nstart: X\n```" def test_member_extras_edited_by_the_guardrail_land_on_that_member(): diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py index 232c6413e78..e6126b02790 100644 --- a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py @@ -17,9 +17,9 @@ from unittest.mock import patch import pytest - from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402 VertexAIBatchTransformation, + vertex_prompt_tokens_details, ) from litellm.llms.vertex_ai.common_utils import ( # noqa: E402 VertexAIError, @@ -41,6 +41,22 @@ ENDPOINT_INPUT_FILE = ( ) +def test_vertex_prompt_tokens_details_rejects_malformed_details(): + assert vertex_prompt_tokens_details({"promptTokensDetails": [1]}) is None + assert vertex_prompt_tokens_details({"promptTokensDetails": [{"modality": "AUDIO"}]}) is None + assert ( + vertex_prompt_tokens_details( + { + "promptTokensDetails": [ + {"modality": "AUDIO", "tokenCount": 1}, + "malformed", + ] + } + ) + is None + ) + + # =========================================================================== # # transform_openai_batch_request_to_vertex_ai_batch_request # =========================================================================== # diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py index 86b3f0976ab..fd8c2a9cf6a 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -10,6 +10,7 @@ Covers: import pytest +import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( _build_part_for_input, @@ -22,11 +23,19 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject from litellm.types.utils import EmbeddingResponse - IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII" GCS_URL = "gs://my-bucket/image.png" +@pytest.fixture(autouse=True) +def _local_model_cost_map(monkeypatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + litellm.get_model_info.cache_clear() + yield + litellm.get_model_info.cache_clear() + + class TestIsMultimodalInput: def test_text_only_string(self): assert _is_multimodal_input("hello world") is False @@ -324,7 +333,7 @@ class TestProcessEmbedContentResponseUsage: ) assert result.usage.prompt_tokens == 258 assert result.usage.total_tokens == 258 - assert result.usage.prompt_tokens_details.image_count == 1 + assert result.usage.prompt_tokens_details.image_tokens == 258 prompt_cost, _ = generic_cost_per_token( model=self.MODEL, @@ -358,7 +367,7 @@ class TestProcessEmbedContentResponseUsage: ) assert prompt_cost > 0 - def test_video_modality_derives_seconds_and_text_floor(self): + def test_video_modality_preserves_token_count(self): response_json = { "embedding": {"values": [0.1]}, "usageMetadata": { @@ -374,10 +383,8 @@ class TestProcessEmbedContentResponseUsage: response_json=response_json, ) assert result.usage.prompt_tokens == 516 - assert result.usage.prompt_tokens_details.video_length_seconds == pytest.approx( - 2.0 - ) - assert result.usage.prompt_tokens_details.text_tokens == 1 + assert result.usage.prompt_tokens_details.video_tokens == 516 + assert result.usage.prompt_tokens_details.text_tokens == 0 def test_missing_usage_metadata_does_not_estimate_from_base64(self): response_json = {"embedding": {"values": [0.1, 0.2]}} @@ -400,8 +407,7 @@ class TestProcessEmbedContentResponseUsage: ) assert result.usage.prompt_tokens > 0 - def test_file_reference_image_billed_per_image_not_text(self): - """files/... image refs must bill per-image, not at the text token rate.""" + def test_file_reference_image_billed_per_image_token_rate(self): response_json = { "embedding": {"values": [0.1, 0.2, 0.3]}, "usageMetadata": { @@ -422,7 +428,7 @@ class TestProcessEmbedContentResponseUsage: } }, ) - assert result.usage.prompt_tokens_details.image_count == 1 + assert result.usage.prompt_tokens_details.image_tokens == 258 assert result.usage.prompt_tokens_details.text_tokens == 0 prompt_cost, _ = generic_cost_per_token( @@ -430,10 +436,10 @@ class TestProcessEmbedContentResponseUsage: usage=result.usage, custom_llm_provider="vertex_ai", ) - assert prompt_cost == pytest.approx(0.00012) + assert prompt_cost == pytest.approx(258 * 4.5e-7) def test_file_reference_non_image_not_counted_as_image(self): - """A files/... ref resolving to a non-image mime must not be image-counted.""" + """A files/... ref resolving to a non-image mime keeps audio token billing.""" response_json = { "embedding": {"values": [0.1, 0.2]}, "usageMetadata": { @@ -454,21 +460,18 @@ class TestProcessEmbedContentResponseUsage: } }, ) - assert result.usage.prompt_tokens_details.image_count == 0 assert result.usage.prompt_tokens_details.audio_tokens == 64 - assert result.usage.prompt_tokens_details.audio_length_seconds == pytest.approx( - 2.0 - ) + assert result.usage.prompt_tokens_details.image_tokens == 0 prompt_cost, _ = generic_cost_per_token( model=self.MODEL, usage=result.usage, custom_llm_provider="vertex_ai", ) - assert prompt_cost == pytest.approx(2.0 * 0.00016) + assert prompt_cost == pytest.approx(64 * 6.5e-6) def test_video_plus_audio_does_not_double_bill_text(self): - """Video+audio responses must not get video tokens reassigned to text.""" + """Video and audio responses are billed from their respective token counts.""" response_json = { "embedding": {"values": [0.1]}, "usageMetadata": { @@ -486,18 +489,145 @@ class TestProcessEmbedContentResponseUsage: model=self.MODEL, response_json=response_json, ) - assert result.usage.prompt_tokens_details.text_tokens == 1 - assert result.usage.prompt_tokens_details.video_length_seconds == pytest.approx( - 2.0 - ) - assert result.usage.prompt_tokens_details.audio_length_seconds == pytest.approx( - 2.0 - ) + assert result.usage.prompt_tokens_details.text_tokens == 0 + assert result.usage.prompt_tokens_details.video_tokens == 516 + assert result.usage.prompt_tokens_details.audio_tokens == 64 prompt_cost, _ = generic_cost_per_token( model=self.MODEL, usage=result.usage, custom_llm_provider="vertex_ai", ) - # 1 floor text token at 2e-7 + 2s of video at 7.9e-4 + 2s of audio at 1.6e-4 - assert prompt_cost == pytest.approx(1 * 2e-7 + 2 * 0.00079 + 2 * 0.00016) + assert prompt_cost == pytest.approx(516 * 1.2e-5 + 64 * 6.5e-6) + + def test_preview_alias_bills_audio_per_token(self): + response_json = { + "embedding": {"values": [0.1]}, + "usageMetadata": { + "promptTokenCount": 64, + "totalTokenCount": 64, + "promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 64}], + }, + } + result = process_embed_content_response( + input="audio", + model_response=EmbeddingResponse(), + model="gemini-embedding-2-preview", + response_json=response_json, + ) + prompt_cost, _ = generic_cost_per_token( + model="gemini-embedding-2-preview", + usage=result.usage, + custom_llm_provider="vertex_ai", + ) + assert prompt_cost == pytest.approx(64 * 6.5e-6) + + def test_image_without_modality_details_uses_image_rate(self): + response_json = { + "embedding": {"values": [0.1]}, + "usageMetadata": { + "promptTokenCount": 258, + "totalTokenCount": 258, + }, + } + result = process_embed_content_response( + input=IMAGE_DATA_URI, + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json=response_json, + ) + assert result.usage.prompt_tokens_details.image_tokens == 258 + assert result.usage.prompt_tokens_details.text_tokens == 0 + + prompt_cost, _ = generic_cost_per_token( + model=self.MODEL, + usage=result.usage, + custom_llm_provider="vertex_ai", + ) + assert prompt_cost == pytest.approx(258 * 4.5e-7) + + @pytest.mark.parametrize( + "input_value,resolved_files,expected_image_tokens", + [ + (GCS_URL, {}, 258), + ("gs://my-bucket/clip.mp4", {}, 0), + ("gs://my-bucket/unknown.bin", {}, 0), + ("files/image-123", {"files/image-123": {"mime_type": "image/jpeg"}}, 258), + ("files/missing", {}, 0), + ("data:application/octet-stream;base64,abc", {}, 0), + ([[IMAGE_DATA_URI]], {}, 258), + ([], {}, 0), + ], + ) + def test_missing_modality_details_classifies_image_inputs(self, input_value, resolved_files, expected_image_tokens): + response_json = { + "embedding": {"values": [0.1]}, + "usageMetadata": { + "promptTokenCount": 258, + "totalTokenCount": 258, + }, + } + result = process_embed_content_response( + input=input_value, + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json=response_json, + resolved_files=resolved_files, + ) + assert result.usage.prompt_tokens_details.image_tokens == expected_image_tokens + assert result.usage.prompt_tokens_details.text_tokens == 0 + + prompt_cost, _ = generic_cost_per_token( + model=self.MODEL, + usage=result.usage, + custom_llm_provider="vertex_ai", + ) + expected_rate = 4.5e-7 if expected_image_tokens else 2e-7 + assert prompt_cost == pytest.approx(258 * expected_rate) + + def test_mixed_text_and_image_without_modality_details_not_billed_as_image(self): + response_json = { + "embedding": {"values": [0.1]}, + "usageMetadata": { + "promptTokenCount": 270, + "totalTokenCount": 270, + }, + } + result = process_embed_content_response( + input=["a short caption", IMAGE_DATA_URI], + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json=response_json, + ) + assert result.usage.prompt_tokens_details.image_tokens == 0 + + prompt_cost, _ = generic_cost_per_token( + model=self.MODEL, + usage=result.usage, + custom_llm_provider="vertex_ai", + ) + assert prompt_cost == pytest.approx(270 * 2e-7) + + def test_text_without_modality_details_uses_text_rate(self): + response_json = { + "embedding": {"values": [0.1]}, + "usageMetadata": { + "promptTokenCount": 12, + "totalTokenCount": 12, + }, + } + result = process_embed_content_response( + input="a short caption", + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json=response_json, + ) + assert result.usage.prompt_tokens_details.text_tokens == 0 + assert result.usage.prompt_tokens_details.image_tokens == 0 + + prompt_cost, _ = generic_cost_per_token( + model=self.MODEL, + usage=result.usage, + custom_llm_provider="vertex_ai", + ) + assert prompt_cost == pytest.approx(12 * 2e-7) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index cdf1f897707..bd6a14cad21 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -433,6 +433,232 @@ def test_get_model_from_request_no_request_extracts_model(): ) +def _cache_prediction_router(): + from litellm.router import Router + + return Router(model_list=[ + { + "model_name": group, + "litellm_params": {"model": "anthropic/claude-sonnet-5", "api_key": "test-provider-key"}, + "model_info": {"id": deployment_id, "team_id": team_id}, + } + for group, deployment_id, team_id in ( + ("current-group", "current-id", None), ("candidate-group", "candidate-id", None), + ("own-group", "own-id", "prediction-team"), ("foreign-group", "foreign-id", "foreign-team"), + ) + ]) + + +@pytest.mark.parametrize("candidate,team_id,expected", [ + ("candidate-id", None, ["current-group", "candidate-group"]), + ("current-id", None, "current-group"), + ("missing-id", None, None), + ("candidate-group", None, None), + ("own-id", None, None), + ("own-id", "prediction-team", ["current-group", "own-group"]), + ("foreign-id", "prediction-team", None), +]) +def test_cache_prediction_auth_resolves_only_exact_deployment_ids(candidate, team_id, expected): + assert get_model_from_request( + request_data={ + "current_deployment_id": "current-id", "candidate_deployment_id": candidate, + "request": {"model": "caller-controlled-provider-model"}, + }, + route="/cost/predict-cache", + llm_router=_cache_prediction_router(), + team_id=team_id, + ) == expected + + +def _cache_prediction_auth_app( + monkeypatch, allowed_routes, user_models, metadata=None, *, team_id=None, key_models=None, team_models=None +): + import importlib + from unittest.mock import AsyncMock + + from fastapi import FastAPI + + import litellm.proxy.proxy_server as proxy_server + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, ProxyException + from litellm.proxy.auth import auth_checks + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.management_endpoints import prompt_cache_prediction as endpoint + from litellm.proxy.utils import InternalUsageCache, ProxyLogging + + auth = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + router = _cache_prediction_router() + allowed_models = ["current-group", "candidate-group", "own-group"] + token = UserAPIKeyAuth( + api_key="test-proxy-key-hash", user_id="prediction-user", user_role=LitellmUserRoles.INTERNAL_USER, + models=allowed_models if key_models is None else key_models, team_id=team_id, + team_models=allowed_models if team_models is None else team_models, + allowed_routes=allowed_routes, metadata=metadata or {}, + ) + user = LiteLLM_UserTable( + user_id=token.user_id, user_role=LitellmUserRoles.INTERNAL_USER.value, models=user_models, + ) + async def authenticate(request, request_data, **_headers): + await auth._enforce_key_and_fallback_model_access( + valid_token=token, request_data=request_data, route=request.url.path, request=request, + llm_model_list=router.get_model_list(), llm_router=router, + ) + return token + + monkeypatch.setattr(auth, "_user_api_key_auth_builder", authenticate) + monkeypatch.setattr(auth, "get_user_object", AsyncMock(return_value=user)) + team = LiteLLM_TeamTableCachedObj(team_id=team_id, models=token.team_models) if team_id else None + monkeypatch.setattr(auth, "get_team_object", AsyncMock(return_value=team)) + monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=team)) + monkeypatch.setattr(auth_checks, "get_team_membership", AsyncMock(return_value=None)) + monkeypatch.setattr(auth, "get_global_proxy_spend", AsyncMock(return_value=0)) + monkeypatch.setattr(proxy_server, "master_key", "test-master-key") + monkeypatch.setattr(proxy_server, "user_custom_auth", None) + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", router.get_model_list()) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + logging = ProxyLogging(user_api_key_cache=DualCache()) + logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + InternalUsageCache(dual_cache=DualCache()) + ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) + counts = AsyncMock(return_value=6_000) + monkeypatch.setattr(endpoint, "count_prompt_tokens", counts) + app = FastAPI() + app.include_router(endpoint.router) + app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler) + return app, counts + + +def _cache_prediction_payload(candidate="candidate-id", current="current-id"): + return { + "current_deployment_id": current, "candidate_deployment_id": candidate, + "request": {"messages": [{"role": "user", "content": [{ + "type": "text", "text": "Stable cached context", + "cache_control": {"type": "ephemeral"}, + }]}]}, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("allowed_routes,user_models,candidate,status_code", [ + (["/chat/completions"], ["current-group", "candidate-group"], "candidate-id", 403), + (["/cost/predict-cache"], ["current-group"], "candidate-id", 403), + (["/cost/*"], ["current-group", "candidate-group"], "candidate-id", 200), + (["/cost/predict-cache"], ["current-group"], "current-id", 200), + (["/cost/predict-cache"], ["current-group"], "missing-id", 404), +]) +async def test_cache_prediction_authorizes_route_and_personal_models_before_provider_counts( + monkeypatch, allowed_routes, user_models, candidate, status_code +): + import httpx + + app, counts = _cache_prediction_auth_app(monkeypatch, allowed_routes, user_models) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/cost/predict-cache", json=_cache_prediction_payload(candidate)) + + assert response.status_code == status_code, response.text + if status_code == 200: + assert counts.await_count == (2 if candidate == "current-id" else 4) + else: + assert counts.await_count == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"]) +@pytest.mark.parametrize("team_id,key_models,user_models,team_models", [ + (None, ["*"], ["*"], None), + (None, ["current-group", "candidate-group"], ["*"], None), + (None, ["*"], ["current-group", "candidate-group"], None), + ("prediction-team", ["*"], ["*"], ["current-group", "candidate-group"]), +]) +async def test_cache_prediction_hides_foreign_and_missing_ids_before_model_authorization( + monkeypatch, arm, team_id, key_models, user_models, team_models +): + import httpx + + app, counts = _cache_prediction_auth_app( + monkeypatch, ["/cost/predict-cache"], user_models, + team_id=team_id, key_models=key_models, team_models=team_models, + ) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + missing = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "missing-id"}) + foreign = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "foreign-id"}) + + assert missing.status_code == foreign.status_code == 404, foreign.text + assert missing.json() == foreign.json() == {"detail": "Deployment not found"} + assert counts.await_count == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"]) +@pytest.mark.parametrize("key_models,team_models,status_code", [ + (["*"], ["*"], 200), + (["current-group", "candidate-group"], ["*"], 403), + (["*"], ["current-group", "candidate-group"], 403), +]) +async def test_cache_prediction_checks_each_visible_team_deployment_model( + monkeypatch, arm, key_models, team_models, status_code +): + import httpx + + app, counts = _cache_prediction_auth_app( + monkeypatch, ["/cost/predict-cache"], ["*"], + team_id="prediction-team", key_models=key_models, team_models=team_models, + ) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/cost/predict-cache", json={**_cache_prediction_payload(), arm: "own-id"}) + + assert response.status_code == status_code, response.text + assert counts.await_count == (4 if status_code == 200 else 0) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("arm", ["current_deployment_id", "candidate_deployment_id"]) +async def test_cache_prediction_checks_each_visible_personal_deployment_model(monkeypatch, arm): + import httpx + + app, counts = _cache_prediction_auth_app(monkeypatch, ["/cost/predict-cache"], ["current-group"]) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/cost/predict-cache", json={**_cache_prediction_payload(candidate="current-id"), arm: "candidate-id"} + ) + + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "user_model_access_denied" + assert counts.await_count == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("header_tag,key_tags,limit,status_code,provider_calls", [ + ("limited", [], 1, 429, 1), + (None, ["limited"], 1, 429, 1), + ("limited", ["limited"], 4, 200, 4), + ("unlimited", [], 1, 200, 4), +]) +async def test_cache_prediction_preserves_authenticated_header_and_key_tag_rpm( + monkeypatch, header_tag, key_tags, limit, status_code, provider_calls +): + import httpx + + app, counts = _cache_prediction_auth_app( + monkeypatch, ["/cost/predict-cache"], ["current-group", "candidate-group"], + metadata={"tag_rpm_limit": {"limited": limit}, "tags": key_tags}, + ) + headers = {"x-litellm-tags": header_tag} if header_tag else {} + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/cost/predict-cache", json=_cache_prediction_payload(), headers=headers) + assert response.status_code == status_code, response.text + assert counts.await_count == provider_calls + if limit == 4: + exhausted = await client.post("/cost/predict-cache", json=_cache_prediction_payload(), headers=headers) + assert exhausted.status_code == 429, exhausted.text + assert counts.await_count == 4 + assert all("metadata" not in call.args[2] for call in counts.await_args_list) + + def test_get_model_from_request_supports_google_model_names_with_slashes(): assert ( get_model_from_request( diff --git a/tests/test_litellm/proxy/client/cli/test_pi.py b/tests/test_litellm/proxy/client/cli/test_pi.py index 03e5d7dd197..6bb49566580 100644 --- a/tests/test_litellm/proxy/client/cli/test_pi.py +++ b/tests/test_litellm/proxy/client/cli/test_pi.py @@ -20,6 +20,13 @@ from litellm.proxy.client.cli.commands.pi import ( ) +def test_listing_failure_is_str_enum(): + assert issubclass(ListingFailure, str) + assert ListingFailure.REJECTED.value == "rejected" + assert ListingFailure("rejected") is ListingFailure.REJECTED + assert str(ListingFailure.REJECTED.value) == "rejected" + + class _FakeResponse: def __init__(self, status_code, payload=None): self.status_code = status_code diff --git a/tests/test_litellm/proxy/client/cli/test_statusline_script.py b/tests/test_litellm/proxy/client/cli/test_statusline_script.py index 5c0cf6b5703..39d0e24d7b0 100644 --- a/tests/test_litellm/proxy/client/cli/test_statusline_script.py +++ b/tests/test_litellm/proxy/client/cli/test_statusline_script.py @@ -252,11 +252,52 @@ class TestRender: def test_savings_header_and_bars_against_the_routers_baseline(self, config_dir): text = render("claude-sonnet-5", RECORDED, config_dir, use_color=False, bar_width=10) assert text.splitlines() == [ - "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5", - "LiteLLM ████░░░░░░ $0.14", + "Routed to: claude-sonnet-5 -63% vs Claude Opus 5", + "claude-auto ████░░░░░░ $0.14", "Claude Opus 5 ██████████ $0.38", ] + def test_a_long_router_name_keeps_both_cost_bars_aligned(self, config_dir: Path) -> None: + session: Final = RECORDED._replace(router_name="engineering-smart-router") + text: Final = render("claude-sonnet-5", session, config_dir, use_color=False, bar_width=10) + assert text.splitlines()[1:] == [ + "engineering-smart-router ████░░░░░░ $0.14", + "Claude Opus 5 ██████████ $0.38", + ] + + @pytest.mark.parametrize( + ("router_name", "baseline_name", "router_padding", "baseline_padding"), + ( + ("路由-router", "Claude Opus 5", 3, 1), + ("智能模型路由器", "Claude Opus 5", 1, 2), + ("ABC-router", "Claude Opus 5", 1, 1), + ("cafe\u0301-router", "Claude Opus 5", 3, 1), + ("a\u20dd-router", "Claude Opus 5", 6, 1), + ("カ\u3099-router", "Claude Opus 5", 5, 1), + ("auto", "基準モデル", 7, 1), + ("auto", "cafe\u0301", 1, 1), + ), + ) + @pytest.mark.parametrize("use_color", (False, True)) + def test_unicode_labels_align_cost_bars_by_terminal_columns( + self, + config_dir: Path, + router_name: str, + baseline_name: str, + router_padding: int, + baseline_padding: int, + use_color: bool, + ) -> None: + (config_dir / "cache" / "gateway-models.json").write_text( + json.dumps({"models": [{"id": "claude-opus-5", "display_name": baseline_name}]}) + ) + session: Final = RECORDED._replace(router_name=router_name) + text: Final = ANSI.sub("", render("claude-sonnet-5", session, config_dir, use_color, bar_width=10)) + assert text.splitlines()[1:] == [ + f"{router_name}{' ' * router_padding}████░░░░░░ $0.14", + f"{baseline_name}{' ' * baseline_padding}██████████ $0.38", + ] + def test_control_characters_in_any_externally_sourced_label_never_reach_the_terminal(self, tmp_path, config_dir): # The transcript, the proxy payload and Claude Code's model cache all feed labels straight into a # terminal, and none is under this script's control. Only the control bytes are dropped (ESC, BEL, @@ -289,7 +330,7 @@ class TestRender: assert "+25% vs Claude Opus 5" in render("m", dearer, config_dir, use_color=False) def test_without_a_baseline_only_the_routed_line_shows(self, config_dir): - assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "claude-auto · Routed to: m" + assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "Routed to: m" assert render("m", None, config_dir, False) == "Routed to: m" def test_color_wraps_the_same_text(self, config_dir): @@ -311,7 +352,8 @@ class TestClaudeCodeMode: return Fetched(RECORDED, definitive=True) text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch) - assert text.startswith("claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n") + assert text.startswith("Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n") + assert text.splitlines()[1].startswith("claude-auto ") def test_a_discovered_display_name_labels_the_sessions_model( self, tmp_path: Path, transcript: Path, config_dir: Path @@ -322,7 +364,7 @@ class TestClaudeCodeMode: return Fetched(session, definitive=True) text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch) - assert text.startswith("claude-auto · Routed to: Claude Opus 5 -63% vs Claude Opus 5\n") + assert text.startswith("Routed to: Claude Opus 5 -63% vs Claude Opus 5\n") def test_an_unrecorded_session_degrades_to_the_routed_line(self, tmp_path, transcript, config_dir): assert _run(_payload(transcript), _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == ( @@ -378,7 +420,8 @@ class TestCodexMode: out = _run({"hook_event_name": "Stop", "session_id": SESSION_ID, "transcript_path": "/nope"}, env, fetch) message = json.loads(out)["systemMessage"] - assert message.splitlines()[1] == "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5" + assert message.splitlines()[1] == "Routed to: claude-sonnet-5 -63% vs Claude Opus 5" + assert message.splitlines()[2].startswith("claude-auto ") assert message.startswith("\n") assert seen == [Credentials("http://127.0.0.1:4000", "sk-codex")] diff --git a/tests/test_litellm/proxy/common_utils/test_openai_error_payload.py b/tests/test_litellm/proxy/common_utils/test_openai_error_payload.py index 8b653ddfb71..90850840ab4 100644 --- a/tests/test_litellm/proxy/common_utils/test_openai_error_payload.py +++ b/tests/test_litellm/proxy/common_utils/test_openai_error_payload.py @@ -143,3 +143,18 @@ def test_a_status_carried_by_an_exception_drives_the_type_it_reports(): exc = HTTPException(status_code=403, detail="blocked by policy") assert openai_error_type(exc, error_status_code(exc, 400)) == "permission_error" + + +def test_a_stringified_none_type_or_param_is_treated_as_absent(): + from litellm.exceptions import BadRequestError + + carried = BadRequestError( + message="Content blocked", + model="claude-haiku-4-5", + llm_provider="litellm_proxy", + body={"message": "Content blocked", "type": "None", "param": "None", "code": "400"}, + ) + + assert carried.type == "None" + assert openai_error_type(carried, 400) == "invalid_request_error" + assert openai_error_param(carried) is None diff --git a/tests/test_litellm/proxy/common_utils/test_prompt_cache_pricing.py b/tests/test_litellm/proxy/common_utils/test_prompt_cache_pricing.py new file mode 100644 index 00000000000..994684a6005 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_prompt_cache_pricing.py @@ -0,0 +1,105 @@ +from typing import Final + +import pytest + +import litellm +from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.types.management_endpoints.prompt_cache_prediction import CacheTokenBuckets + + +@pytest.mark.parametrize( + ("model", "expected"), + [("anthropic/claude-sonnet-4-5", 1.26), ("anthropic/claude-sonnet-4-6", 0.63)], +) +def test_prices_all_cache_buckets_at_total_context_tier(model: str, expected: float) -> None: + tokens: Final = CacheTokenBuckets( + uncached_input_tokens=100_000, + cache_read_input_tokens=50_000, + cache_creation_5m_input_tokens=20_000, + cache_creation_1h_input_tokens=40_000, + ) + assert price_cache_tokens(model, "unconfigured-deployment", tokens) == pytest.approx(expected) + + +@pytest.mark.parametrize(("total", "expected"), [(200_000, 0.387), (200_001, 0.774006)]) +def test_long_context_tier_starts_above_threshold(total: int, expected: float) -> None: + tokens: Final = CacheTokenBuckets( + uncached_input_tokens=total - 100_000, + cache_creation_1h_input_tokens=10_000, + cache_read_input_tokens=90_000, + ) + actual: Final = price_cache_tokens("anthropic/claude-sonnet-4-5", "unconfigured-deployment", tokens) + assert actual == pytest.approx(expected) + + +def test_deployment_tariff_wins_without_proxy_discounts_or_margins(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "model_cost", litellm.model_cost.copy()) + litellm.Router( + model_list=[ + { + "model_name": "cache-pricing-test", + "litellm_params": { + "model": "anthropic/claude-sonnet-4-6", + "api_key": "test-only", + "input_cost_per_token": 0.00001, + "output_cost_per_token": 0.00002, + "cache_read_input_token_cost": 0.000001, + "cache_creation_input_token_cost": 0.0000125, + "cache_creation_input_token_cost_above_1hr": 0.00002, + }, + "model_info": {"id": "cache-pricing-test-a"}, + } + ] + ) + monkeypatch.setattr(litellm, "cost_discount_config", {"anthropic": 0.5}) + monkeypatch.setattr(litellm, "cost_margin_config", {"global": {"percentage": 0.3, "fixed_amount": 1.0}}) + tokens: Final = CacheTokenBuckets( + uncached_input_tokens=3_000, + cache_read_input_tokens=4_000, + cache_creation_5m_input_tokens=1_000, + cache_creation_1h_input_tokens=2_000, + ) + assert price_cache_tokens("anthropic/claude-sonnet-4-6", "cache-pricing-test-a", tokens) == pytest.approx(0.0865) + + +@pytest.mark.parametrize("rate", [None, -1.0, float("nan"), float("inf"), "0.00001", True]) +def test_unknown_for_absent_or_invalid_active_cache_rate(monkeypatch: pytest.MonkeyPatch, rate: object) -> None: + monkeypatch.setitem( + litellm.model_cost, + "cache-pricing-invalid", + { + "litellm_provider": "anthropic", + "mode": "chat", + "input_cost_per_token": 0.00001, + "output_cost_per_token": 0.00002, + "cache_creation_input_token_cost_above_1hr": rate, + }, + ) + tokens: Final = CacheTokenBuckets(cache_creation_1h_input_tokens=4_000) + assert price_cache_tokens("anthropic/claude-sonnet-4-6", "cache-pricing-invalid", tokens) is None + + +def test_missing_input_price_is_unknown_even_when_get_model_info_defaults_to_zero( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setitem(litellm.model_cost, "cache-pricing-missing", {"litellm_provider": "anthropic", "mode": "chat"}) + tokens: Final = CacheTokenBuckets(uncached_input_tokens=4_000) + assert price_cache_tokens("cache-pricing-missing", "unconfigured-deployment", tokens) is None + + +def test_explicit_free_pricing_is_not_unknown(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setitem( + litellm.model_cost, + "cache-pricing-free", + { + "litellm_provider": "anthropic", + "mode": "chat", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "cache_read_input_token_cost": 0.0, + "cache_creation_input_token_cost": 0.0, + "cache_creation_input_token_cost_above_1hr": 0.0, + }, + ) + tokens: Final = CacheTokenBuckets(uncached_input_tokens=100, cache_read_input_tokens=5_000) + assert price_cache_tokens("anthropic/claude-sonnet-4-6", "cache-pricing-free", tokens) == 0.0 diff --git a/tests/test_litellm/proxy/db/test_health_check_latest.py b/tests/test_litellm/proxy/db/test_health_check_latest.py index 529e4d8f2e9..6322891ae9e 100644 --- a/tests/test_litellm/proxy/db/test_health_check_latest.py +++ b/tests/test_litellm/proxy/db/test_health_check_latest.py @@ -8,6 +8,7 @@ from litellm.proxy.db.health_check_latest import ( LATEST_HEALTH_CHECKS_SQL, fetch_latest_health_checks, fetch_latest_health_checks_for_models, + query_latest_health_checks, ) @@ -83,6 +84,15 @@ async def test_fetch_all_degrades_to_no_rows_when_the_query_fails(): assert await fetch_latest_health_checks(prisma) == () +@pytest.mark.asyncio +async def test_query_all_raises_when_the_query_fails_instead_of_reading_as_an_empty_table(): + """The background save decides what to write from this read; a failure has to be told apart from no rows.""" + prisma = _prisma([]) + prisma.db.query_raw.side_effect = RuntimeError("db down") + with pytest.raises(RuntimeError, match="db down"): + await query_latest_health_checks(prisma) + + @pytest.mark.asyncio async def test_fetch_all_degrades_to_no_rows_for_a_malformed_row(): assert await fetch_latest_health_checks(_prisma([{"unexpected": "shape"}])) == () diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index c0762edec92..d173c5f5c70 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -2375,8 +2375,12 @@ _GROUNDING_QUERY_TEXT = "What is the capital of Japan?" _GROUNDING_RESPONSE_TEXT = "The capital of Japan is Tokyo." -def _grounding_guardrail() -> BedrockGuardrail: - return BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT") +def _grounding_guardrail(from_messages: bool = False) -> BedrockGuardrail: + return BedrockGuardrail( + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + contextual_grounding_from_messages=from_messages, + ) def _grounding_messages() -> list: @@ -2418,9 +2422,11 @@ def _input_request(messages: list) -> dict: return _grounding_guardrail().convert_to_bedrock_format(source="INPUT", messages=messages) -def _output_request(messages: list, response=None) -> dict: +def _output_request(messages: list, response=None, from_messages: bool = False) -> dict: """Arrange a guardrail and act: build the Bedrock OUTPUT payload.""" - return _grounding_guardrail().convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages) + return _grounding_guardrail(from_messages).convert_to_bedrock_format( + source="OUTPUT", response=response, messages=messages + ) def test_grounding_input_strips_grounding_and_query_qualifiers(): @@ -2474,6 +2480,131 @@ def test_grounding_output_keeps_legacy_payload_without_tags(): assert actual_request == expected_request +def test_grounding_output_derives_source_and_query_from_plain_messages(): + """Flag on: untagged system + user text is sent as grounding_source + query.""" + messages = [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ] + expected_request = { + "source": "OUTPUT", + "content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK], + } + + actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True) + + assert actual_request == expected_request + + +def test_grounding_output_plain_messages_stay_legacy_when_flag_is_off(): + """Default config: plain system + user text is never sent as grounding context.""" + messages = [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ] + expected_request = {"source": "OUTPUT", "content": [{"text": {"text": _GROUNDING_RESPONSE_TEXT}}]} + + actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT)) + + assert actual_request == expected_request + + +def test_grounding_output_derived_query_is_latest_user_turn_only(): + """Only the latest user turn is the query; system and developer turns are the source.""" + developer_text = "Answer in one sentence." + messages = [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "developer", "content": [{"type": "text", "text": developer_text}]}, + {"role": "user", "content": "Hi"}, + {"role": "assistant", "content": "Hello, how can I help?"}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ] + expected_request = { + "source": "OUTPUT", + "content": [ + _GROUNDING_SOURCE_BLOCK, + {"text": {"text": developer_text, "qualifiers": ["grounding_source"]}}, + _QUERY_BLOCK, + _GUARD_BLOCK, + ], + } + + actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True) + + assert actual_request == expected_request + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param([{"role": "system", "content": _GROUNDING_SOURCE_TEXT}], id="system-without-user"), + pytest.param( + [ + {"role": "tool", "content": _GROUNDING_SOURCE_TEXT, "tool_call_id": "c1"}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ], + id="tool-result-is-not-a-source", + ), + pytest.param( + [ + {"role": "system", "content": ""}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ], + id="empty-system-prompt", + ), + pytest.param( + [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://x.test/a.png"}}]}, + ], + id="image-only-user-turn", + ), + ], +) +def test_grounding_output_stays_legacy_when_plain_source_or_query_is_missing(messages): + """Bedrock rejects a source without a query and vice versa, so send neither.""" + expected_request = {"source": "OUTPUT", "content": [{"text": {"text": _GROUNDING_RESPONSE_TEXT}}]} + + actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True) + + assert actual_request == expected_request + + +def test_grounding_output_explicit_tags_take_precedence_over_plain_messages(): + """Tagged blocks win: untagged text around them is not added as source or query.""" + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + *_grounding_messages(), + {"role": "user", "content": "Please be brief."}, + ] + expected_request = { + "source": "OUTPUT", + "content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK], + } + + actual_request = _output_request(messages, _model_response(_GROUNDING_RESPONSE_TEXT), from_messages=True) + + assert actual_request == expected_request + + +def test_grounding_input_ignores_plain_message_derivation(): + """INPUT scans never derive grounding qualifiers from plain messages.""" + messages = [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ] + expected_request = { + "source": "INPUT", + "content": [{"text": {"text": _GROUNDING_SOURCE_TEXT}}, {"text": {"text": _GROUNDING_QUERY_TEXT}}], + } + + actual_request = _grounding_guardrail(from_messages=True).convert_to_bedrock_format( + source="INPUT", messages=messages + ) + + assert actual_request == expected_request + + def test_grounding_output_combines_multiple_sources(): """Every grounding_source block is emitted; Bedrock combines them into one corpus.""" uk_source_text = "London is the capital of UK." @@ -2597,6 +2728,56 @@ async def test_grounding_output_blocked_raises_400(): assert exc_info.value.status_code == 400 +@pytest.mark.asyncio +@pytest.mark.parametrize( + "from_messages, request_messages", + [ + ( + True, + [ + {"role": "system", "content": _GROUNDING_SOURCE_TEXT}, + {"role": "user", "content": _GROUNDING_QUERY_TEXT}, + ], + ), + ( + False, + [ + {"role": "system", "content": [{"type": "grounding_source", "text": _GROUNDING_SOURCE_TEXT}]}, + {"role": "user", "content": [{"type": "query", "text": _GROUNDING_QUERY_TEXT}]}, + ], + ), + ], + ids=["plain-messages-flag-on", "tagged-messages-flag-off"], +) +async def test_apply_guardrail_response_forwards_request_messages_for_grounding(from_messages, request_messages): + guardrail = _grounding_guardrail(from_messages=from_messages) + expected_request = { + "source": "OUTPUT", + "content": [_GROUNDING_SOURCE_BLOCK, _QUERY_BLOCK, _GUARD_BLOCK], + } + + mock_credentials = MagicMock() + mock_credentials.access_key = "test-access-key" + mock_credentials.secret_key = "test-secret-key" + mock_credentials.token = None + + with ( + patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post, + patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")), + patch.object(guardrail, "_prepare_request", return_value=MagicMock()) as mock_prepare, + ): + mock_post.return_value = _passing_bedrock_httpx_response(_GROUNDING_RESPONSE_TEXT) + + await guardrail.apply_guardrail( + inputs={"texts": [_GROUNDING_RESPONSE_TEXT]}, + request_data={"messages": request_messages}, + input_type="response", + ) + + assert mock_prepare.call_count == 1 + assert json.loads(json.dumps(mock_prepare.call_args.kwargs["data"])) == expected_request + + ############################################################################### # LIT-4186: disable_exception_on_block regression tests # diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index a9ca13a463d..9849ad7ec88 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -1820,7 +1820,7 @@ async def test_unalignable_rewrite_is_rejected_never_sent_unredacted( Skipping the write-back would hand the model the unredacted text, so a guardrail could be bypassed by adding ``instructions`` or a tool call. """ - from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite + from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite data: dict[str, object] = {"model": "gpt-4o", "input": responses_input} if instructions is not None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 83cc9ae8bb9..a5e79f84ef1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -582,6 +582,145 @@ class TestGuardrailActions: assert result_images is None +class TestStructuredMessagesInResponse: + """A guardrail server that rewrites per chat row answers with the rewritten + rows as structured_messages, which the endpoint handlers write back by row.""" + + @pytest.mark.asyncio + async def test_returned_rows_are_handed_back_as_structured_messages( + self, generic_guardrail, mock_request_data_input + ): + rewritten_rows = [ + {"role": "system", "content": "Never repeat an SSN."}, + {"role": "user", "content": "Look up [REDACTED] for me."}, + {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'}, + ] + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["Never repeat an SSN.", "Look up [REDACTED] for me.", '{"ssn": "[REDACTED]"}'], + "structured_messages": rewritten_rows, + } + mock_response.raise_for_status = MagicMock() + + with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert guardrailed_inputs["structured_messages"] == rewritten_rows + assert guardrailed_inputs["texts"] == mock_response.json.return_value["texts"] + + @pytest.mark.asyncio + async def test_rows_echoed_back_as_shown_keep_their_original_keys( + self, generic_guardrail, mock_request_data_input + ): + tool_call_row = { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}, "index": 0} + ], + } + original_rows = [ + {"role": "user", "content": "Look up 123-45-6789 for me.", "name": "pat"}, + tool_call_row, + {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'}, + ] + + def echo_with_tool_output_redacted(url, json, headers): + shown_rows = json["structured_messages"] + assert "index" not in shown_rows[1]["tool_calls"][0] + assert "name" not in shown_rows[0] + answer = MagicMock() + answer.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["Look up 123-45-6789 for me."], + "structured_messages": [ + shown_rows[0], + shown_rows[1], + {**shown_rows[2], "content": '{"ssn": "[REDACTED]"}'}, + ], + } + answer.raise_for_status = MagicMock() + return answer + + with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_with_tool_output_redacted): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."], "structured_messages": original_rows}, + request_data=mock_request_data_input, + input_type="request", + ) + + returned_rows = guardrailed_inputs["structured_messages"] + assert returned_rows[0] is original_rows[0] + assert returned_rows[1] is tool_call_row + assert returned_rows[2] == {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'} + + @pytest.mark.asyncio + async def test_rows_all_echoed_back_as_shown_leave_the_rewrite_to_texts( + self, generic_guardrail, mock_request_data_input + ): + """A server written against the texts contract that echoes the request rows back + untouched while rewriting texts still gets its texts rewrite applied.""" + original_rows = [ + {"role": "system", "content": "Never repeat an SSN."}, + {"role": "user", "content": "Look up 123-45-6789 for me."}, + ] + + def echo_rows_and_rewrite_texts(url, json, headers): + answer = MagicMock() + answer.json.return_value = { + "action": "NONE", + "texts": [text.replace("123-45-6789", "[REDACTED]") for text in json["texts"]], + "structured_messages": json["structured_messages"], + } + answer.raise_for_status = MagicMock() + return answer + + with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_rows_and_rewrite_texts): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={ + "texts": ["Never repeat an SSN.", "Look up 123-45-6789 for me."], + "structured_messages": original_rows, + }, + request_data=mock_request_data_input, + input_type="request", + ) + + assert "structured_messages" not in guardrailed_inputs + assert guardrailed_inputs["texts"] == ["Never repeat an SSN.", "Look up [REDACTED] for me."] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "structured_messages", + [[], [{"content": "a row with no role"}], "not a list"], + ids=["empty", "no_role", "not_a_list"], + ) + async def test_rows_that_are_not_chat_messages_are_ignored( + self, generic_guardrail, mock_request_data_input, structured_messages + ): + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "GUARDRAIL_INTERVENED", + "texts": ["[REDACTED]"], + "structured_messages": structured_messages, + } + mock_response.raise_for_status = MagicMock() + + with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response): + guardrailed_inputs = await generic_guardrail.apply_guardrail( + inputs={"texts": ["Look up 123-45-6789 for me."]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert "structured_messages" not in guardrailed_inputs + assert guardrailed_inputs["texts"] == ["[REDACTED]"] + + class TestImageSupport: """Test image handling in guardrail requests""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index d8eeb8d2b8a..d4531398ba1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -1797,12 +1797,8 @@ PARTS_MESSAGES = [ { "role": "user", "content": [ - {"type": "text", "text": "Earlier turn.", "cache_control": {"type": "ephemeral"}}, - { - "type": "text", - "text": "Second block. " + "B" * 5000, - "cache_control": {"type": "ephemeral", "ttl": "1h"}, - }, + {"type": "text", "text": "Earlier turn."}, + {"type": "text", "text": "Second block. " + "B" * 5000}, ], }, { @@ -1891,14 +1887,9 @@ async def test_apply_guardrail_restores_rewritten_all_text_row( messages = result["structured_messages"] history_content = messages[1]["content"] - # Rewritten all-text row collapses to one part carrying the LAST declared - # breakpoint: an Anthropic breakpoint caches the prefix ending at its - # part, so after the merge the last one (and its TTL) still describes the - # row. assert isinstance(history_content, list) assert len(history_content) == 1 assert history_content[0]["text"] == "compressed history. Retrieve more: hash=b573993006976af767214fac" - assert history_content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"} # Mixed row passes through byte-identical. assert messages[2]["content"] == PARTS_MESSAGES[2]["content"] # The service-declared hash still drives retrieve-tool injection on a restored row. @@ -2523,6 +2514,35 @@ async def test_history_is_still_compressed(guardrail: HeadroomGuardrail): assert messages[3] == compressed_history[1] +CACHE_MARKED_HISTORY_MESSAGES = [ + {"role": "system", "content": "You are Claude Code. " + "S" * 5000}, + {"role": "user", "content": "old question " + "Q" * 5000}, + { + "role": "assistant", + "content": "Reading the file now.", + "tool_calls": [{"id": "old_1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}], + }, + { + "role": "tool", + "tool_call_id": "old_1", + "content": [{"type": "text", "text": "large cached file body " + "F" * 5000}], + "cache_control": {"type": "ephemeral"}, + }, + {"role": "assistant", "content": "Summarized the file for you."}, + {"role": "user", "content": "live instruction"}, +] + + +@pytest.mark.asyncio +async def test_mid_history_cache_control_row_is_never_sent_for_compression(guardrail: HeadroomGuardrail): + wire, result = await _wire_and_result(guardrail, CACHE_MARKED_HISTORY_MESSAGES) + + cached_row = CACHE_MARKED_HISTORY_MESSAGES[3] + assert cached_row not in wire + assert not any(row.get("tool_call_id") == "old_1" for row in wire) + assert result["structured_messages"][3] == cached_row + + # --------------------------------------------------------------------------- # #38558: a client that runs its own tool loop (e.g. Claude Code via the MCP # gateway) executes headroom_retrieve and echoes the recovered original content diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 47089b7b1b1..5aa3c7dfc17 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -2,6 +2,8 @@ import asyncio import base64 import io import json +from collections.abc import Iterator, Sequence +from typing import cast from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -14,7 +16,7 @@ import litellm import litellm.types.utils from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache -from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, MaskedHTTPStatusError from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail @@ -4929,3 +4931,390 @@ def test_every_responses_delta_event_is_in_the_scanned_set(): } assert not missing assert "response.mcp_call_arguments.delta" in _RESPONSES_DELTA_EVENT_TYPES + + +def _clean_armor_response() -> dict[str, object]: + return { + "sanitizationResult": { + "filterMatchState": "NO_MATCH_FOUND", + "filterResults": {}, + } + } + + +def _flagged_armor_response() -> dict[str, object]: + return { + "sanitizationResult": { + "filterMatchState": "MATCH_FOUND", + "filterResults": {"rai": {"raiFilterResult": {"matchState": "MATCH_FOUND"}}}, + } + } + + +class _FakeArmorHandler(AsyncHTTPHandler): + def __init__(self, responses: Sequence[dict[str, object] | Exception]): + self.responses: Iterator[dict[str, object] | Exception] = iter(responses) + self.calls: list[dict[str, object]] = [] + self.raise_on_call: Exception | None = None + + async def post( + self, + url: str, + json: dict[str, object] | None = None, + headers: dict[str, str] | None = None, + **kwargs: object, + ) -> httpx.Response: + if self.raise_on_call is not None: + raise self.raise_on_call + if json is not None: + self.calls.append(json) + response: dict[str, object] | Exception = next(self.responses) + if isinstance(response, Exception): + raise response + return httpx.Response(200, json=response, request=httpx.Request("POST", url)) + + +async def _async_token_provider() -> tuple[str, str]: + return ("test-token", "test-project") + + +def _logging_only_guardrail( + responses: Sequence[dict[str, object] | Exception] = (_clean_armor_response(), _clean_armor_response()), +) -> ModelArmorGuardrail: + handler = _FakeArmorHandler(responses) + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-logging", + event_hook=GuardrailEventHooks.logging_only, + async_handler=handler, + access_token_provider=_async_token_provider, + ) + return guardrail + + +def _logged_kwargs() -> dict[str, object]: + return { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "litellm_call_id": "call-1", + "litellm_params": {"metadata": {}}, + "optional_params": {}, + "standard_logging_object": {"guardrail_information": None}, + } + + +def _chat_response(text: str) -> litellm.ModelResponse: + return litellm.ModelResponse( + choices=[ + litellm.types.utils.Choices( + message=litellm.types.utils.Message(role="assistant", content=text) + ) + ] + ) + + +def _stream_chunk(text: str) -> litellm.ModelResponseStream: + return litellm.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + delta=litellm.types.utils.Delta(content=text) + ) + ] + ) + + +def _metadata_entries(kwargs: dict[str, object]) -> list[dict[str, object]]: + standard_logging_object = cast(dict[str, object], kwargs["standard_logging_object"]) + entries = standard_logging_object.get("guardrail_information") or [] + return cast(list[dict[str, object]], entries) + + +def test_logging_only_mode_is_accepted_and_keeps_native_hooks(): + guardrail = _logging_only_guardrail() + assert guardrail.event_hook == GuardrailEventHooks.logging_only + assert guardrail.use_native_lifecycle_hooks is True + assert GuardrailEventHooks.logging_only in ModelArmorGuardrail.get_supported_event_hooks() + + post_call_guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-post", + event_hook=GuardrailEventHooks.post_call, + ) + assert post_call_guardrail._deployment_hook_target() is post_call_guardrail + + +@pytest.mark.asyncio +async def test_logging_only_stream_yields_chunks_without_waiting_for_scan(): + """A logging_only guardrail must pass stream chunks straight through; the scan happens + afterwards on the assembled response via async_logging_hook.""" + guardrail = _logging_only_guardrail( + [_clean_armor_response(), _clean_armor_response()] + ) + handler = cast(_FakeArmorHandler, guardrail.async_handler) + handler.raise_on_call = AssertionError("logging_only must not scan the stream") + + produced = 0 + + async def gen(): + nonlocal produced + for i in range(3): + produced += 1 + yield _stream_chunk(f"chunk-{i} ") + + hook_iter = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=gen(), + request_data={"metadata": {}, "guardrails": ["model-armor-logging"]}, + ) + first = await hook_iter.__anext__() + assert produced == 1 + chunks = [first] + async for chunk in hook_iter: + chunks.append(chunk) + assert len(chunks) == 3 + assert handler.calls == [] + handler.raise_on_call = None + + response = _chat_response("all clear") + kwargs = _logged_kwargs() + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + assert out_result is response + entries = _metadata_entries(out_kwargs) + assert len(entries) >= 1 + entry = entries[-1] + assert entry["guardrail_status"] == "success" + assert entry["guardrail_mode"] == "logging_only" + assert entry["guardrail_provider"] == "model_armor" + + +@pytest.mark.asyncio +async def test_logging_only_records_flagged_verdict_without_altering_response(): + guardrail = _logging_only_guardrail( + [_flagged_armor_response(), _flagged_armor_response()] + ) + response = _chat_response("flagged output") + kwargs = _logged_kwargs() + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + + assert out_result is response + entries = _metadata_entries(out_kwargs) + assert entries[-1]["guardrail_status"] == "guardrail_flagged" + assert entries[-1]["guardrail_mode"] == "logging_only" + + +@pytest.mark.asyncio +async def test_logging_only_records_model_armor_api_error(): + guardrail = _logging_only_guardrail( + [ + ModelArmorAPIError("Model Armor API error (upstream 500)"), + ModelArmorAPIError("Model Armor API error (upstream 500)"), + ] + ) + response = _chat_response("some output") + kwargs = _logged_kwargs() + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + + assert out_result is response + entries = _metadata_entries(out_kwargs) + assert entries[-1]["guardrail_status"] == "guardrail_failed_to_respond" + + +@pytest.mark.asyncio +async def test_logging_only_scans_assembled_responses_api_stream(): + """The terminal ResponseCompletedEvent is an envelope; the scan must run on the + assembled ResponsesAPIResponse kept in kwargs.""" + from openai.types.responses import ResponseOutputMessage, ResponseOutputText + + from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + assembled = ResponsesAPIResponse( + id="resp-1", + created_at=1700000000, + output=[ + ResponseOutputMessage( + id="msg-1", + type="message", + role="assistant", + status="completed", + content=[ + ResponseOutputText( + annotations=[], text="assembled output text", type="output_text" + ) + ], + ) + ], + ) + event = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, response=assembled + ) + + guardrail = _logging_only_guardrail() + kwargs = _logged_kwargs() + del kwargs["messages"] + kwargs["input"] = "hello" + kwargs["async_complete_streaming_response"] = assembled + + out_kwargs, _ = await guardrail.async_logging_hook( + kwargs=kwargs, result=event, call_type="aresponses" + ) + + handler = cast(_FakeArmorHandler, guardrail.async_handler) + response_scans = [call for call in handler.calls if "modelResponseData" in call] + assert response_scans, "expected a model_response scan of the assembled response" + assert "assembled output text" in response_scans[0]["modelResponseData"]["text"] + assert _metadata_entries(out_kwargs) + + +@pytest.mark.asyncio +async def test_logging_only_scans_anthropic_messages_model_response(): + """/v1/messages logs a ModelResponse; the output scan must extract the assistant text.""" + guardrail = _logging_only_guardrail() + kwargs = _logged_kwargs() + kwargs["messages"] = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}] + response = _chat_response("anthropic assembled text") + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="anthropic_messages" + ) + + assert out_result is response + handler = cast(_FakeArmorHandler, guardrail.async_handler) + response_scans = [call for call in handler.calls if "modelResponseData" in call] + assert response_scans + assert "anthropic assembled text" in response_scans[0]["modelResponseData"]["text"] + assert _metadata_entries(out_kwargs) + + +@pytest.mark.asyncio +async def test_logging_only_skips_output_scan_when_no_assembled_response(): + guardrail = _logging_only_guardrail() + kwargs = _logged_kwargs() + + await guardrail.async_logging_hook(kwargs=kwargs, result=None, call_type="acompletion") + + handler = cast(_FakeArmorHandler, guardrail.async_handler) + assert all("modelResponseData" not in call for call in handler.calls) + + +@pytest.mark.asyncio +async def test_native_post_call_mode_ignores_logging_hook(): + handler = _FakeArmorHandler([_clean_armor_response()]) + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-post", + event_hook=GuardrailEventHooks.post_call, + async_handler=handler, + access_token_provider=_async_token_provider, + ) + response = _chat_response("some output") + kwargs = _logged_kwargs() + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + + assert out_kwargs is kwargs + assert out_result is response + assert handler.calls == [] + + +@pytest.mark.asyncio +async def test_apply_guardrail_records_flagged_without_raising(): + guardrail = _logging_only_guardrail([_flagged_armor_response()]) + request_data = {"metadata": {}} + inputs = {"texts": ["forbidden output"]} + + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + ) + + assert result == inputs + entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert entries[-1]["guardrail_status"] == "guardrail_flagged" + + +@pytest.mark.asyncio +async def test_logging_only_records_transport_error(): + guardrail = _logging_only_guardrail([httpx.ConnectError("boom"), httpx.ConnectError("boom")]) + response = _chat_response("some output") + kwargs = _logged_kwargs() + + out_kwargs, out_result = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + + assert out_result is response + entries = _metadata_entries(out_kwargs) + failed = [e for e in entries if e["guardrail_status"] == "guardrail_failed_to_respond"] + assert failed + assert all(e["guardrail_provider"] == "model_armor" for e in failed) + + +@pytest.mark.asyncio +async def test_logging_only_flagged_prompt_still_scans_response(): + """A flagged input scan must not abort the output scan; both verdicts are recorded.""" + guardrail = _logging_only_guardrail( + [_flagged_armor_response(), _flagged_armor_response()] + ) + response = _chat_response("flagged output") + kwargs = _logged_kwargs() + + out_kwargs, _ = await guardrail.async_logging_hook( + kwargs=kwargs, result=response, call_type="acompletion" + ) + + handler = cast(_FakeArmorHandler, guardrail.async_handler) + sources = ["user_prompt" if "userPromptData" in call else "model_response" for call in handler.calls] + assert sources == ["user_prompt", "model_response"] + entries = _metadata_entries(out_kwargs) + flagged = [e for e in entries if e["guardrail_status"] == "guardrail_flagged"] + assert len(flagged) == 2 + + +@pytest.mark.asyncio +async def test_apply_guardrail_raises_on_flagged_when_not_logging_only(): + """The /guardrails/apply_guardrail endpoint calls apply_guardrail directly; a + non-logging_only instance must signal the block so flagged text is not returned as clean.""" + handler = _FakeArmorHandler([_flagged_armor_response()]) + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-pre", + event_hook=GuardrailEventHooks.pre_call, + async_handler=handler, + access_token_provider=_async_token_provider, + ) + request_data = {"metadata": {}} + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["forbidden prompt"]}, + request_data=request_data, + input_type="request", + ) + + assert exc_info.value.status_code == 400 + entries = request_data["metadata"]["standard_logging_guardrail_information"] + flagged = [e for e in entries if e["guardrail_status"] == "guardrail_flagged"] + assert len(flagged) == 1 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 3d7c6e06d94..f25727ebd9a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -4620,46 +4620,27 @@ class TestPanwAirsLatestRoleMessageOnly: @pytest.mark.asyncio async def test_anthropic_system_plus_multiturn_no_fallback(self): - """Anthropic with top-level system + multi-turn messages[] - — latest-user works, no scan-all fallback. + """Anthropic with a top-level system prompt and multi-turn messages[] + scans only the latest user turn, with no scan-all fallback. - Key scenario: Anthropic top-level `system` field causes - structured_messages to have an injected system entry, but - request_data["messages"] does NOT include it. + The Anthropic handler hoists the top-level `system` field into both + `texts` and `structured_messages`, so the latest-user walk has to + count the same entries the framework flattened. """ - handler = PanwPrismaAirsHandler( - guardrail_name="test_panw_airs", - api_key="test_api_key", - profile_name="test_profile", - default_on=True, + from litellm.llms.anthropic.chat.guardrail_translation.handler import ( + AnthropicMessagesHandler, ) - # Original Anthropic messages (no system in messages array) - original_messages = [ - {"role": "user", "content": "First user turn"}, - {"role": "assistant", "content": "First assistant turn"}, - {"role": "user", "content": "Latest user turn"}, - ] - - # texts extracted from original_messages (3 text entries) - texts = ["First user turn", "First assistant turn", "Latest user turn"] - - # structured_messages has an INJECTED system message from translation - structured_messages = [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "First user turn"}, - {"role": "assistant", "content": "First assistant turn"}, - {"role": "user", "content": "Latest user turn"}, - ] - - inputs: GenericGuardrailAPIInputs = { - "texts": texts, - "structured_messages": structured_messages, - } + handler = make_handler() request_data = { "litellm_call_id": "test-call-id", "model": "anthropic/claude-sonnet-4-20250514", - "messages": original_messages, + "system": "You are a helpful assistant.", + "messages": [ + {"role": "user", "content": "First user turn"}, + {"role": "assistant", "content": "First assistant turn"}, + {"role": "user", "content": "Latest user turn"}, + ], "proxy_server_request": { "url": "http://localhost:4000/v1/messages", }, @@ -4670,13 +4651,11 @@ class TestPanwAirsLatestRoleMessageOnly: ) as mock_api: mock_api.return_value = {"action": "allow", "category": "benign"} - await handler.apply_guardrail( - inputs=inputs, - request_data=request_data, - input_type="request", + await AnthropicMessagesHandler().process_input_messages( + data=request_data, + guardrail_to_apply=handler, ) - # Should scan ONLY the latest user message, not fall back to scan-all assert mock_api.call_count == 1 assert mock_api.call_args.kwargs["content"] == "Latest user turn" diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index c550a0a41d2..6295469c066 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -16,14 +16,21 @@ Streaming: CSW.__anext__ stores args on logging_obj at stream end. import asyncio import logging -from typing import Any, Final +from collections.abc import Callable, Mapping +from datetime import datetime +from typing import Any, Final, cast from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +import respx import litellm from litellm.caching.caching import DualCache +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.types.utils import StandardLoggingPayload +from litellm.utils import _dispatch_success_logging from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth @@ -54,6 +61,27 @@ def _attach_mock_success_dispatch(mock_logging_obj, async_success_fn): mock_logging_obj.async_success_handler = async_success_fn +async def _wait_until(condition: Callable[[], bool]) -> None: + """Give the logging worker a bounded window to run what the closure enqueued.""" + for _ in range(200): + if condition(): + return + await asyncio.sleep(0.01) + + +class _RecordingLogger(CustomLogger): + """Keeps what the async success callback was handed, the way a spend logger sees it.""" + + def __init__(self) -> None: + super().__init__() + self.standard_logging_object: StandardLoggingPayload | None = None + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self.standard_logging_object = cast(StandardLoggingPayload, kwargs["standard_logging_object"]) + + class PostCallGuardrail(CustomGuardrail): """A post-call guardrail.""" @@ -259,6 +287,120 @@ async def test_deferred_flag_stores_and_executes_closure(): pass +@pytest.mark.asyncio +async def test_deferred_slot_keeps_the_innermost_wrapper_result(): + """Nested @client wrappers exit through _dispatch_success_logging with one shared logging + object. The deferred slot must keep the first stored result, the way the immediate path's + has_logged dedupe keeps the first fired task, so the spend log reads usage from the + innermost provider-shaped response and never from an outer wrapper's translation of it.""" + logging_obj: Final = MagicMock() + logging_obj._defer_async_logging = True + logging_obj._enqueue_deferred_logging = None + logging_obj.async_success_handler = AsyncMock() + inner_result: Final = object() + outer_result: Final = object() + + for result in (inner_result, outer_result): + _dispatch_success_logging( + logging_obj=logging_obj, + result=result, + start_time=datetime.now(), + end_time=datetime.now(), + is_completion_with_fallbacks=False, + is_litellm_internal_call=False, + ) + + logging_obj._enqueue_deferred_logging() + await _wait_until(lambda: logging_obj.async_success_handler.await_count > 0) + + logging_obj.async_success_handler.assert_awaited_once() + assert logging_obj.async_success_handler.await_args.kwargs["result"] is inner_result + assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_count == 2 + + +@pytest.mark.asyncio +async def test_deferred_anthropic_messages_bridged_to_the_responses_api_logs_the_provider_usage( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + """/v1/messages on an Azure gpt-5.4+ deployment with function tools runs three nested + wrappers: anthropic_messages, the chat adapter's acompletion, and the Responses bridge + acompletion hands the call to, which retags the call as ``responses``. With logging + deferred for a post-call guardrail the stored closure must carry the innermost provider + response: logging the Anthropic-shaped reply under Responses semantics books this + 7,336-token prompt as 3 tokens, since Anthropic's input_tokens excludes the cache hit.""" + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.post(url__regex=r"https://deferred-nested\.openai\.azure\.com/openai/.*responses.*").mock( + return_value=httpx.Response( + 200, + json={ + "id": "resp_deferred_nested", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.4-nano", + "output": [ + { + "type": "message", + "id": "msg_deferred_nested", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello!", "annotations": []}], + } + ], + "usage": { + "input_tokens": 7336, + "input_tokens_details": {"cached_tokens": 7333}, + "output_tokens": 23, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": 7359, + }, + }, + ) + ) + recorder: Final = _RecordingLogger() + logging_obj: Final = Logging( + model="azure/gpt-5.4-nano", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id="deferred-nested-anthropic-messages", + function_id="deferred-nested-anthropic-messages", + dynamic_async_success_callbacks=[recorder], + ) + logging_obj._defer_async_logging = True + + response: Final = await litellm.anthropic_messages( + model="azure/gpt-5.4-nano", + messages=[{"role": "user", "content": "hi"}], + max_tokens=16, + tools=[ + { + "name": "lookup_volume", + "description": "Look up a storage volume by name", + "input_schema": {"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}, + } + ], + api_key="sk-deferred-nested", + api_base="https://deferred-nested.openai.azure.com", + api_version="2025-04-01-preview", + litellm_logging_obj=logging_obj, + ) + assert response["content"] == [{"type": "text", "text": "Hello!"}] + assert response["usage"]["input_tokens"] == 3 + assert response["usage"]["cache_read_input_tokens"] == 7333 + + logging_obj._enqueue_deferred_logging() + await _wait_until(lambda: recorder.standard_logging_object is not None) + + assert recorder.standard_logging_object is not None + assert recorder.standard_logging_object["prompt_tokens"] == 7336 + assert recorder.standard_logging_object["metadata"]["usage_object"]["prompt_tokens_details"]["cached_tokens"] == 7333 + assert recorder.standard_logging_object["response_cost"] == pytest.approx(3 * 2e-7 + 7333 * 2e-8 + 23 * 1.25e-6) + + # --------------------------------------------------------------------------- # 3. Non-streaming regression: without flag, create_task fires normally # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index ceb084b4a4d..8377db57b6e 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -71,6 +71,52 @@ def test_initialize_bedrock_forwards_chunk_budget_chars(): assert initialized[-1].chunk_budget_chars == 60_000 +def test_initialize_bedrock_forwards_contextual_grounding_from_messages(): + """`contextual_grounding_from_messages: true` in config.yaml must make the post-call + payload carry the plain system prompt and user turn as grounding_source and query.""" + import litellm + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + from litellm.types.utils import Choices, Message, ModelResponse + + test_guardrail = { + "guardrail_name": "test_bedrock_grounding_from_messages", + "litellm_params": { + "guardrail": SupportedGuardrailIntegrations.BEDROCK.value, + "mode": "post_call", + "guardrailIdentifier": "test-guardrail", + "guardrailVersion": "DRAFT", + "contextual_grounding_from_messages": True, + }, + } + messages = [ + {"role": "system", "content": "Returns are accepted for 30 days."}, + {"role": "user", "content": "How long is the return window?"}, + ] + response = ModelResponse( + choices=[Choices(index=0, message=Message(role="assistant", content="30 days."), finish_reason="stop")] + ) + expected_request = { + "source": "OUTPUT", + "content": [ + {"text": {"text": "Returns are accepted for 30 days.", "qualifiers": ["grounding_source"]}}, + {"text": {"text": "How long is the return window?", "qualifiers": ["query"]}}, + {"text": {"text": "30 days.", "qualifiers": ["guard_content"]}}, + ], + } + + guardrail_handler = InMemoryGuardrailHandler() + guardrail_handler.initialize_guardrail(guardrail=test_guardrail) + + initialized = [ + callback + for callback in litellm.callbacks + if isinstance(callback, BedrockGuardrail) and callback.guardrail_name == "test_bedrock_grounding_from_messages" + ] + assert initialized, "bedrock guardrail was not registered as a callback" + actual_request = initialized[-1].convert_to_bedrock_format(source="OUTPUT", response=response, messages=messages) + assert json.loads(json.dumps(actual_request)) == expected_request + + def test_initialize_guardrail_preserves_guardrail_info(): """ Regression (LIT-2529): initialize_guardrail must carry guardrail_info into the diff --git a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py index e650f796f29..3218632a8d2 100644 --- a/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_prompt_security_guardrails.py @@ -1,5 +1,6 @@ import asyncio import base64 +from collections.abc import Mapping, Sequence from unittest.mock import AsyncMock, patch import pytest @@ -12,6 +13,7 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im PromptSecurityGuardrailMissingSecrets, ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.llms.openai import AllMessageValues def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch): @@ -174,6 +176,123 @@ async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch): assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"] +def _modify_response(modified_messages: Sequence[Mapping[str, object]]) -> Response: + mock_response = Response( + json={"result": {"prompt": {"action": "modify", "modified_messages": modified_messages}}}, + status_code=200, + request=Request(method="POST", url="https://test.prompt.security/api/protect"), + ) + mock_response.raise_for_status = lambda: None + return mock_response + + +def _tool_replay_messages() -> list[AllMessageValues]: + return [ + {"role": "system", "content": "Never echo an SSN like 123-45-6789."}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Look up 123-45-6789"}, + {"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}}, + ], + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'}, + {"role": "user", "content": "Summarize what you found."}, + ] + + +@pytest.mark.asyncio +async def test_modify_returns_structured_messages_with_tool_rows_kept(monkeypatch: pytest.MonkeyPatch): + """A per-message modify verdict comes back as structured_messages so the + endpoint handler can write it back by message, with the rows Prompt Security + never saw (tool results) and the non-text parts (images) left in place.""" + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True) + messages = _tool_replay_messages() + inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages} + modified_messages = [ + {"role": "system", "content": "Never echo an SSN like [REDACTED]."}, + {"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}]}, + {"role": "assistant", "content": None}, + {"role": "user", "content": "Summarize what you found."}, + ] + + with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)): + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={"messages": messages}, input_type="request" + ) + + assert result["structured_messages"] == [ + {"role": "system", "content": "Never echo an SSN like [REDACTED]."}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Look up [REDACTED]"}, + {"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}}, + ], + }, + messages[2], + messages[3], + {"role": "user", "content": "Summarize what you found."}, + ] + assert result["structured_messages"] is not messages + assert result["texts"] == [ + "Never echo an SSN like [REDACTED].", + "Look up [REDACTED]", + "Summarize what you found.", + ] + + +@pytest.mark.asyncio +async def test_modify_with_unexpected_message_count_keeps_texts_only(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True) + messages = _tool_replay_messages() + inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages} + modified_messages = [{"role": "user", "content": "Look up [REDACTED]"}] + + with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)): + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={"messages": messages}, input_type="request" + ) + + assert result["structured_messages"] is messages + assert result["texts"] == ["Look up [REDACTED]"] + + +@pytest.mark.asyncio +async def test_modify_keeps_empty_text_parts_as_slots(monkeypatch: pytest.MonkeyPatch): + """The chat handler counts an empty text part as a slot, so a modify verdict + that echoes the empty part still lines up with the row and its texts.""" + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True) + messages: list[AllMessageValues] = [ + {"role": "user", "content": [{"type": "text", "text": "Look up 123-45-6789"}, {"type": "text", "text": ""}]} + ] + inputs = {"texts": ["Look up 123-45-6789", ""], "structured_messages": messages} + modified_messages = [ + {"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}, {"type": "text", "text": ""}]} + ] + + with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)): + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={"messages": messages}, input_type="request" + ) + + assert result["structured_messages"] == modified_messages + assert result["texts"] == ["Look up [REDACTED]", ""] + + @pytest.mark.asyncio async def test_apply_guardrail_allow_request(monkeypatch: pytest.MonkeyPatch): """Test that apply_guardrail allows safe prompts""" @@ -497,6 +616,98 @@ async def test_file_sanitization_modify_can_rewrite_when_blocking_disabled(monke assert base64.b64decode(result["file"]["data"]) == b"name,email\nAlice,[REDACTED]\n" +@pytest.mark.asyncio +async def test_file_sanitization_keeps_polling_through_queued_statuses(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + + guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True) + guardrail.poll_interval = 0 + upload_response = Response( + json={"jobId": "queued-job"}, + status_code=200, + request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"), + ) + poll_request = Request(method="GET", url="https://test.prompt.security/api/sanitizeFile") + poll_responses = [ + Response(json={"status": "created"}, status_code=200, request=poll_request), + Response(json={"status": "in progress"}, status_code=200, request=poll_request), + Response( + json={"status": "done", "content": "clean", "metadata": {"action": "allow", "violations": []}}, + status_code=200, + request=poll_request, + ), + ] + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)): + with patch.object(guardrail.async_handler, "get", AsyncMock(side_effect=poll_responses)) as poll_mock: + result = await guardrail.sanitize_file_content(b"image-content", "image.png") + + assert poll_mock.await_count == 3 + assert result["action"] == "allow" + assert result["content"] == "clean" + + +@pytest.mark.asyncio +async def test_file_sanitization_never_finishing_job_times_out(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + + guardrail = PromptSecurityGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True, file_sanitization_fail_open=False + ) + guardrail.poll_interval = 0 + guardrail.max_poll_attempts = 3 + upload_response = Response( + json={"jobId": "stuck-job"}, + status_code=200, + request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"), + ) + poll_response = Response( + json={"status": "created"}, + status_code=200, + request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"), + ) + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)): + with patch.object(guardrail.async_handler, "get", AsyncMock(return_value=poll_response)) as poll_mock: + with pytest.raises(HTTPException) as exc_info: + await guardrail.sanitize_file_content(b"file-content", "document.pdf") + + assert poll_mock.await_count == 3 + assert exc_info.value.status_code == 408 + assert exc_info.value.detail == "File sanitization timeout" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("poll_body", [{"status": "failed"}, {}]) +async def test_file_sanitization_terminal_failure_does_not_fail_open(monkeypatch: pytest.MonkeyPatch, poll_body): + monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key") + monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security") + + guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True) + guardrail.poll_interval = 0 + upload_response = Response( + json={"jobId": "failed-job"}, + status_code=200, + request=Request(method="POST", url="https://test.prompt.security/api/sanitizeFile"), + ) + poll_response = Response( + json=poll_body, + status_code=200, + request=Request(method="GET", url="https://test.prompt.security/api/sanitizeFile"), + ) + + with patch.object(guardrail.async_handler, "post", AsyncMock(return_value=upload_response)): + with patch.object(guardrail.async_handler, "get", AsyncMock(return_value=poll_response)) as poll_mock: + with pytest.raises(HTTPException) as exc_info: + await guardrail.sanitize_file_content(b"file-content", "document.pdf") + + assert poll_mock.await_count == 1 + assert exc_info.value.status_code == 500 + assert exc_info.value.detail == f"Unexpected sanitization status: {poll_body.get('status')}" + + @pytest.mark.asyncio @pytest.mark.parametrize( "timeout", diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index 644b213d3a3..db87e12ac88 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -682,3 +682,67 @@ async def test_detail_prev_trend_query_is_bounded(): prev_wheres = [w for w in wheres if "lt" in w.get("date", {})] assert prev_wheres assert all("gte" in w["date"] for w in prev_wheres) + + +@pytest.mark.asyncio +async def test_logs_report_not_run_entries_as_not_run_not_passed(): + """LIT-6314: a guardrail that never scanned must not be reported as a pass in the drill-down.""" + index_row = MagicMock() + index_row.request_id = "req-nr" + index_row.guardrail_id = "db-1" + index_row.start_time = datetime(2026, 4, 22) + spend_log = MagicMock() + spend_log.request_id = "req-nr" + spend_log.model = "gpt-4o-mini" + spend_log.startTime = datetime(2026, 4, 22) + spend_log.metadata = { + "guardrail_information": [ + {"guardrail_name": "db-1", "guardrail_status": "not_run", "duration": 0.0}, + ] + } + prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row]) + prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log]) + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_logs( + guardrail_id="db-1", + policy_id=None, + page=1, + page_size=50, + action=None, + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + assert [log.action for log in resp.logs] == ["not_run"] + + +@pytest.mark.asyncio +async def test_logs_action_passed_filter_excludes_not_run_entries(): + """LIT-6314: filtering the drill-down for passes must not return unscanned requests.""" + index_row = MagicMock() + index_row.request_id = "req-nr" + index_row.guardrail_id = "db-1" + index_row.start_time = datetime(2026, 4, 22) + spend_log = MagicMock() + spend_log.request_id = "req-nr" + spend_log.model = "gpt-4o-mini" + spend_log.startTime = datetime(2026, 4, 22) + spend_log.metadata = {"guardrail_information": [{"guardrail_name": "db-1", "guardrail_status": "not_run"}]} + prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row]) + prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log]) + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_logs( + guardrail_id="db-1", + policy_id=None, + page=1, + page_size=50, + action="passed", + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + assert resp.logs == [] diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py index 85f22e1f307..69ec098b840 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -349,6 +349,95 @@ async def test_zero_and_non_int_usage_counters_are_skipped(): } +@pytest.mark.asyncio +async def test_not_run_entries_are_indexed_but_not_counted_as_evaluations(): + """ + LIT-6314 records a not_run entry when message scoping leaves a guardrail + nothing to scan. The guardrail never evaluated the request, so counting it + as a passed evaluation would inflate daily pass rates; it still gets an + index row so per-request drill-down finds the spend log. + """ + prisma = _prisma() + logs = [_payload("r1", guardrail_status="not_run"), _payload("r2")] + + await process_spend_logs_guardrail_usage(prisma, logs) + + metrics_create = prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"] + assert metrics_create["requests_evaluated"] == 1 + assert metrics_create["passed_count"] == 1 + index_rows = prisma.db.litellm_spendlogguardrailindex.create_many.call_args.kwargs["data"] + assert sorted(row["request_id"] for row in index_rows) == ["r1", "r2"] + + +@pytest.mark.asyncio +async def test_not_run_entry_shares_index_key_with_evaluated_sibling_of_same_name(): + """ + The not_run entry from the shared base guardrail carries only guardrail_name, + while the evaluated entry from the same guardrail (e.g. content filter on the + output of a logging_only run) carries its guardrail_id. Keying them differently + lists one request twice in the monitor, once as not_run and once as passed. + """ + prisma = _prisma() + payload = _payload("r1") + payload["metadata"] = json.dumps( + { + "guardrail_information": [ + {"guardrail_name": "cf", "guardrail_status": "not_run"}, + { + "guardrail_name": "cf", + "guardrail_id": "cf-uuid", + "policy_id": "pol-1", + "guardrail_status": "success", + }, + {"guardrail_name": "other", "guardrail_status": "not_run"}, + ] + } + ) + + await process_spend_logs_guardrail_usage(prisma, [payload]) + + index_rows = prisma.db.litellm_spendlogguardrailindex.create_many.call_args.kwargs["data"] + assert sorted((row["guardrail_id"], row["policy_id"]) for row in index_rows) == [ + ("cf-uuid", "pol-1"), + ("other", None), + ] + metrics_create = prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"] + assert (metrics_create["guardrail_id"], metrics_create["requests_evaluated"]) == ("cf-uuid", 1) + + +@pytest.mark.asyncio +async def test_malformed_not_run_entry_does_not_drop_the_batch(): + prisma = _prisma() + payload = _payload("r1") + payload["metadata"] = json.dumps( + { + "guardrail_information": [ + {"guardrail_name": ["not", "a", "string"], "guardrail_status": "success"}, + {"guardrail_name": "", "guardrail_id": "cf-uuid", "guardrail_status": "success"}, + {"guardrail_status": "success"}, + ] + } + ) + + await process_spend_logs_guardrail_usage(prisma, [payload]) + + index_rows = prisma.db.litellm_spendlogguardrailindex.create_many.call_args.kwargs["data"] + assert [row["guardrail_id"] for row in index_rows] == ["cf-uuid"] + metrics_create = prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"] + assert (metrics_create["guardrail_id"], metrics_create["requests_evaluated"]) == ("cf-uuid", 1) + + +@pytest.mark.asyncio +async def test_batch_of_only_not_run_entries_writes_no_metrics_row(): + prisma = _prisma() + + await process_spend_logs_guardrail_usage(prisma, [_payload("r1", guardrail_status="not_run")]) + + assert prisma.db.litellm_dailyguardrailmetrics.upsert.call_count == 0 + index_rows = prisma.db.litellm_spendlogguardrailindex.create_many.call_args.kwargs["data"] + assert [row["request_id"] for row in index_rows] == ["r1"] + + @pytest.mark.asyncio async def test_payload_without_request_id_is_skipped_like_the_metrics_path(): prisma = _prisma() diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 10c0bb88a82..48f980086fd 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -6284,3 +6284,248 @@ async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_ assert isinstance(values, list) assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == [] assert any("circuit breaker is open" in record.getMessage() for record in caplog.records) + + +@pytest.mark.parametrize( + "limits, request_data, counter_scope", + [ + ({"rpm_limit": 1}, {}, "api_key"), + ({"user_id": "u", "user_rpm_limit": 1}, {}, "user"), + ({"team_id": "t", "team_rpm_limit": 1}, {}, "team"), + ( + {"team_id": "t", "user_id": "u", "team_member_rpm_limit": 1}, + {}, + "team_member", + ), + ({"end_user_id": "e", "end_user_rpm_limit": 1}, {}, "end_user"), + ( + {"metadata": {"model_rpm_limit": {"test-model": 1}}}, + {}, + "model_per_key", + ), + ( + {"metadata": {"tag_rpm_limit": {"test-tag": 1}}}, + {"metadata": {"tags": ["test-tag"]}}, + "tag_per_key", + ), + ( + { + "team_id": "t", + "metadata": {"model_rpm_limit": {"test-model": 100}}, + "team_metadata": {"model_rpm_limit": {"test-model": 1}}, + }, + {}, + "model_per_team", + ), + ( + {"project_id": "p", "project_metadata": {"model_rpm_limit": {"test-model": 1}}}, + {}, + "model_per_project", + ), + ({"org_id": "o", "organization_rpm_limit": 1}, {}, "organization"), + ( + {"org_id": "o", "organization_metadata": {"model_rpm_limit": {"test-model": 1}}}, + {}, + "model_per_organization", + ), + ], +) +@pytest.mark.parametrize("request_kind", ["count", "generation"]) +@pytest.mark.asyncio +async def test_request_capacity_enforces_shared_rpm_scopes( + limits, request_data, counter_scope, request_kind +): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key=hash_token("sk-count-rpm"), **limits) + async def request(): + if request_kind == "generation": + await handler.async_pre_call_hook( + user_api_key_dict=auth, + cache=cache, + data={**request_data, "model": "test-model"}, + call_type="acompletion", + ) + return + async with handler.request_capacity(auth, "test-model", request_data=request_data): + pass + + await request() + with pytest.raises(HTTPException) as exc: + await request() + assert exc.value.status_code == 429 + assert counter_scope in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_request_capacity_keeps_dynamic_rpm_policy(monkeypatch): + import litellm.proxy.proxy_server as proxy_server + + router = Router(model_list=[{ + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-test", "api_key": "test-key"}, + "model_info": {"id": "test-deployment"}, + }]) + monkeypatch.setattr(proxy_server, "llm_router", router) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + auth = UserAPIKeyAuth( + api_key=hash_token("sk-count-dynamic"), + rpm_limit=1, + metadata={"rpm_limit_type": "dynamic"}, + ) + for _ in range(2): + async with handler.request_capacity(auth, "test-model"): + pass + router.cache.set_cache("test-deployment:fails", 100, ttl=60, local_only=True) + async with handler.request_capacity(auth, "test-model"): + pass + with pytest.raises(HTTPException) as exc: + async with handler.request_capacity(auth, "test-model"): + pytest.fail("dynamic RPM must enforce after deployment failures") + assert exc.value.status_code == 429 + + +@pytest.mark.asyncio +async def test_request_capacity_skips_tokens_and_preserves_parent_stash(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + auth = UserAPIKeyAuth( + api_key=hash_token("sk-count-tpm"), + rpm_limit=5, + tpm_limit=1, + max_parallel_requests=1, + project_id="p", + project_metadata={ + "model_tpm_limit": {"test-model": 1}, + "model_itpm_limit": {"test-model": 1}, + "model_otpm_limit": {"test-model": 1}, + }, + ) + token_scopes = ( + ("api_key", auth.api_key), + ("model_per_project", "p:test-model"), + ("model_per_project_itpm", "p:test-model"), + ("model_per_project_otpm", "p:test-model"), + ) + for scope, value in token_scopes: + token_key = handler.create_rate_limit_keys(scope, value, "tokens") + await cache.async_set_cache(token_key, 100, ttl=60) + await cache.async_set_cache(f"{{{scope}:{value}}}:window", int(time.time()), ttl=60) + parent = get_or_create_request_stash() + parent.reserved_tokens = 123 + parent.parallel_slot = ParallelSlotAcquisition(slot_id="parent", counter_keys=["parent-gauge"]) + for _ in range(2): + async with handler.request_capacity(auth, "test-model"): + assert get_request_stash() is parent + assert parent.parallel_slot["slot_id"] == "parent" + assert parent.reserved_tokens == 123 + for scope, value in token_scopes: + assert await cache.async_get_cache(handler.create_rate_limit_keys(scope, value, "tokens")) == 100 + + +@pytest.mark.parametrize("exit_mode", ["success", "failure", "cancel"]) +@pytest.mark.asyncio +async def test_request_capacity_releases_exact_parallel_slot(exit_mode): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key=hash_token("sk-count-parallel"), max_parallel_requests=1) + entered = asyncio.Event() + finish = asyncio.Event() + + async def provider(): + async with handler.request_capacity(auth, "test-model"): + entered.set() + await finish.wait() + if exit_mode == "failure": + raise RuntimeError("provider failed") + + task = asyncio.create_task(provider()) + await asyncio.wait_for(entered.wait(), timeout=2) + try: + for _ in range(2): + with pytest.raises(HTTPException) as exc: + async with handler.request_capacity(auth, "test-model"): + pytest.fail("rejected request freed the occupied slot") + assert exc.value.status_code == 429 + finally: + if exit_mode == "cancel": + task.cancel() + else: + finish.set() + if exit_mode == "success": + await task + else: + with pytest.raises(asyncio.CancelledError if exit_mode == "cancel" else RuntimeError): + await task + async with handler.request_capacity(auth, "test-model"): + pass + + +class _DelayedCapacityUsageCache: + def __init__(self): + self.delegate = InternalUsageCache(DualCache()) + self.dual_cache = self.delegate.dual_cache + self.acquired = asyncio.Event() + self.finish_admission = asyncio.Event() + self.releasing = asyncio.Event() + self.finish_release = asyncio.Event() + + async def async_get_cache(self, *args, **kwargs): + return await self.delegate.async_get_cache(*args, **kwargs) + + async def async_batch_get_cache(self, *args, **kwargs): + return await self.delegate.async_batch_get_cache(*args, **kwargs) + + async def async_set_cache(self, key, value, **kwargs): + await self.delegate.async_set_cache(key=key, value=value, **kwargs) + if not key.endswith(":max_parallel_requests"): + return + if value: + self.acquired.set() + await self.finish_admission.wait() + else: + self.releasing.set() + await self.finish_release.wait() + + +@pytest.mark.asyncio +async def test_request_capacity_finishes_admission_and_release_despite_repeated_cancel(): + cache = _DelayedCapacityUsageCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + auth = UserAPIKeyAuth(api_key=hash_token("sk-count-cancel-admission"), max_parallel_requests=1) + + async def provider(): + async with handler.request_capacity(auth, "test-model"): + pytest.fail("cancelled admission entered provider body") + + task = asyncio.create_task(provider()) + await asyncio.wait_for(cache.acquired.wait(), timeout=2) + task.cancel() + await asyncio.sleep(0) + cache.finish_admission.set() + await asyncio.wait_for(cache.releasing.wait(), timeout=2) + task.cancel() + await asyncio.sleep(0) + task.cancel() + await asyncio.sleep(0) + assert not task.done() + cache.finish_release.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=2) + async with handler.request_capacity(auth, "test-model"): + pass + + +@pytest.mark.asyncio +async def test_request_capacity_rejection_keeps_existing_redis_mirror(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key=hash_token("sk-count-mirror"), max_parallel_requests=1) + counter_key = handler.create_rate_limit_keys("api_key", auth.api_key, "max_parallel_requests") + await cache.async_set_cache(counter_key, 1, ttl=60, local_only=True) + for _ in range(2): + with pytest.raises(HTTPException) as exc: + async with handler.request_capacity(auth, "test-model"): + pytest.fail("rejection released another request's mirrored slot") + assert exc.value.status_code == 429 + assert await cache.async_get_cache(counter_key, local_only=True) == 1 diff --git a/tests/test_litellm/proxy/hooks/test_prompt_cache_observer.py b/tests/test_litellm/proxy/hooks/test_prompt_cache_observer.py new file mode 100644 index 00000000000..82af3e9a6ef --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_prompt_cache_observer.py @@ -0,0 +1,300 @@ +import asyncio +import json +import time +from datetime import datetime + +import httpx +import pytest + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.llms.anthropic.chat.transformation import AnthropicConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.anthropic.prompt_cache_prediction import cache_scope, parse_prompt +from litellm.proxy.hooks.prompt_cache_prediction import ( + PromptCacheObserver, + lookup, +) +from litellm.proxy.utils import InternalUsageCache +from litellm.types.utils import ModelResponse + +MODEL = "claude-sonnet-5" +CALLER = "a" * 64 +DEPLOYMENT = "native-deployment" +KEY = "test-provider-key" + + +def body(ttl="5m", texts=("private cache prefix",)): + return { + "model": MODEL, + "max_tokens": 2, + "system": "private system instructions", + "tools": [{"name": "lookup", "input_schema": {"type": "object"}}], + "messages": [{"role": "user", "content": [ + {"type": "text", "text": text, **( + {"cache_control": {"type": "ephemeral", "ttl": ttl}} + if index == len(texts) - 1 else {} + )} + for index, text in enumerate(texts) + ]}], + } + + +def usage(ttl="5m", read=100, write=200): + return { + "input_tokens": 11, + "output_tokens": 2, + "cache_read_input_tokens": read, + "cache_creation_input_tokens": write, + "cache_creation": { + "ephemeral_5m_input_tokens": write if ttl == "5m" else 0, + "ephemeral_1h_input_tokens": write if ttl == "1h" else 0, + }, + } + + +def event(request_body, started=1000.0, headers=None, **overrides): + request = httpx.Request( + "POST", "https://api.anthropic.com/v1/messages", json=request_body, + headers={"x-api-key": KEY, "anthropic-version": "2023-06-01", **(headers or {})}, + ) + return { + "call_type": "anthropic_messages", + "custom_llm_provider": "anthropic", + "cache_hit": False, + "httpx_response": httpx.Response(200, request=request), + "first_api_call_start_time": datetime.fromtimestamp(started), + "standard_logging_object": { + "status": "success", "model_id": DEPLOYMENT, + "metadata": {"user_api_key_hash": CALLER}, + }, + **overrides, + } + + +async def observe(cache, request_body=None, native_usage=None, now=1010.0, **overrides): + observer = PromptCacheObserver(InternalUsageCache(dual_cache=cache), clock=lambda: now) + response = ModelResponse( + model=MODEL, + usage=AnthropicConfig().calculate_usage(native_usage or usage(), reasoning_content=None), + ) + await observer.async_log_success_event( + event(request_body or body(), **overrides), response, + datetime.fromtimestamp(now), datetime.fromtimestamp(now), + ) + + +def scope(**overrides): + return cache_scope(**{ + "caller_key_hash": CALLER, "deployment_id": DEPLOYMENT, + "provider_key": KEY, "model": MODEL, **overrides, + }) + + +@pytest.mark.parametrize("ttl,expires", [("5m", 1300), ("1h", 4600)]) +@pytest.mark.asyncio +async def test_observed_cache_count_and_request_start_expiry_survive_as_stale(ttl, expires): + cache = DualCache() + request_body = body(ttl=ttl) + await observe(cache, request_body, usage(ttl=ttl)) + prefix = parse_prompt(request_body) + observed = await lookup(cache, scope(), prefix, now=1200) + assert observed.cached_tokens == 300 + assert observed.observed_at == 1010 + assert observed.expires_at == expires + assert await lookup(cache, scope(), prefix, now=expires) == observed + saved = json.dumps(cache.in_memory_cache.cache_dict) + assert "private cache prefix" not in saved + assert "private system instructions" not in saved + assert KEY not in saved + assert CALLER not in saved + + +@pytest.mark.parametrize("changed", [ + {"caller_key_hash": "b" * 64}, {"deployment_id": "other"}, + {"provider_key": "rotated"}, {"model": "claude-opus-5"}, + {"anthropic_version": "different"}, +]) +@pytest.mark.asyncio +async def test_cache_evidence_is_isolated_by_every_scope_dimension(changed): + cache = DualCache() + await observe(cache) + assert await lookup(cache, scope(**changed), parse_prompt(body()), now=1010) is None + + +@pytest.mark.asyncio +async def test_append_only_prefix_finds_prior_evidence_but_edit_or_context_change_does_not(): + cache = DualCache() + await observe(cache) + extended = parse_prompt(body(texts=("private cache prefix", "new turn"))) + prior = await lookup(cache, scope(), extended, now=1010) + assert prior.cached_tokens == 300 + assert prior.fingerprint != extended.fingerprint + for changed in ( + body(texts=("edited prefix", "new turn")), + {**body(), "system": "different system"}, + {**body(), "tools": [{"name": "other", "input_schema": {"type": "object"}}]}, + body(ttl="1h"), + ): + assert await lookup(cache, scope(), parse_prompt(changed), now=1010) is None + outside_lookback = parse_prompt(body(texts=("private cache prefix", *[str(i) for i in range(20)]))) + assert await lookup(cache, scope(), outside_lookback, now=1010) is None + + +@pytest.mark.parametrize("change", [ + {"thinking": {"type": "enabled", "budget_tokens": 1024}}, + {"tool_choice": {"type": "auto"}}, + {"cache_control": {"type": "ephemeral"}}, + {"tools": [{"type": "web_search_20250305", "name": "web_search"}]}, + {"system": [{"type": "text", "text": "system", "cache_control": {"type": "ephemeral"}}]}, + {"messages": [{"role": "user", "content": [{"type": "image", "source": {}}]}]}, + {"messages": [{"role": "user", "content": "no breakpoint"}]}, +]) +def test_unsupported_or_ambiguous_shapes_have_no_cache_identity(change): + assert parse_prompt({**body(), **change}) is None + duplicate = body() + duplicate["messages"][0]["content"].append(duplicate["messages"][0]["content"][0]) + assert parse_prompt(duplicate) is None + + +@pytest.mark.parametrize("overrides", [ + {"cache_hit": True}, {"call_type": "completion"}, + {"custom_llm_provider": "bedrock"}, {"stream": True}, + {"headers": {"anthropic-beta": "unverified-feature"}}, + {"headers": {"x-custom-header": "unverified"}}, + {"standard_logging_object": {"status": "success", "model_id": DEPLOYMENT, "metadata": {}}}, +]) +@pytest.mark.asyncio +async def test_unverified_source_never_creates_observations(overrides): + cache = DualCache() + await observe(cache, **overrides) + assert await lookup(cache, scope(), parse_prompt(body()), now=1010) is None + + +@pytest.mark.parametrize("native_usage", [ + usage(write=0), + {**usage(), "cache_creation": None}, + {**usage(), "cache_creation": {"ephemeral_5m_input_tokens": 199, "ephemeral_1h_input_tokens": 0}}, + usage(ttl="1h"), + {**usage(), "cache_creation_input_tokens": -200}, +]) +@pytest.mark.asyncio +async def test_missing_or_contradictory_telemetry_cannot_create_observations(native_usage): + cache = DualCache() + await observe(cache, native_usage=native_usage) + assert await lookup(cache, scope(), parse_prompt(body()), now=1010) is None + + +@pytest.mark.asyncio +async def test_pure_read_refresh_requires_prior_matching_evidence(): + cache = DualCache() + await observe(cache, native_usage=usage(read=300, write=0)) + assert await lookup(cache, scope(), parse_prompt(body()), now=1010) is None + await observe(cache) + await observe(cache, native_usage=usage(read=300, write=0), started=1100, now=1110) + assert (await lookup(cache, scope(), parse_prompt(body()), now=1110)).expires_at == 1400 + + +class RecordingObserver(PromptCacheObserver): + def __init__(self, cache): + super().__init__(InternalUsageCache(dual_cache=cache)) + self.finished = asyncio.Event() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await super().async_log_success_event(kwargs, response_obj, start_time, end_time) + self.finished.set() + + +def native_response(): + return { + "id": "msg_prediction", "type": "message", "role": "assistant", "model": MODEL, + "content": [{"type": "text", "text": "ok"}], "stop_reason": "end_turn", + "stop_sequence": None, "usage": usage(ttl="1h"), + } + + +def stream_response(completed, provider_error=False): + response = native_response() + events = [ + {"type": "message_start", "message": {**response, "content": [], "stop_reason": None}}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "ok"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}}, + ] + if completed: + events.append({"type": "message_stop"}) + if provider_error: + events.append({"type": "error", "error": {"type": "overloaded_error", "message": "temporary failure"}}) + return "".join(f"event: {item['type']}\ndata: {json.dumps(item)}\n\n" for item in events) + + +class TransportChunks(httpx.AsyncByteStream): + def __init__(self, payload, chunk_size, fragment_error_only=False): + self.payload = payload.encode() + self.chunk_size = chunk_size or len(self.payload) + self.prefix_length = self.payload.index(b"event: error") if fragment_error_only else 0 + + async def __aiter__(self): + if self.prefix_length: + yield self.payload[:self.prefix_length] + for offset in range(self.prefix_length, len(self.payload), self.chunk_size): + yield self.payload[offset:offset + self.chunk_size] + + +@pytest.mark.parametrize("stream,completed,provider_error,transport", [ + (False, True, False, "whole"), + (True, True, False, "whole"), + (True, False, False, "whole"), + (True, True, True, "whole"), + (True, True, False, "fragmented"), + (True, False, False, "fragmented"), + (True, True, True, "fragmented"), + (True, True, True, "fragmented_error"), + (True, True, False, "unterminated"), +]) +@pytest.mark.asyncio +async def test_native_production_callback_records_only_completed_wire_requests(stream, completed, provider_error, transport): + cache = DualCache() + observer = RecordingObserver(cache) + litellm.logging_callback_manager.add_litellm_callback(observer) + + def provider(request): + if stream: + payload = stream_response(completed, provider_error) + if transport == "unterminated": + payload = payload.removesuffix("\n\n") + return httpx.Response( + 200, request=request, headers={"content-type": "text/event-stream"}, + stream=TransportChunks( + payload, 1 if transport.startswith("fragmented") else None, + fragment_error_only=transport == "fragmented_error", + ), + ) + return httpx.Response(200, request=request, json=native_response()) + + client = AsyncHTTPHandler() + await client.client.aclose() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(provider)) + try: + request_body = body(ttl="1h") + before = time.time() + result = await litellm.anthropic_messages( + **{**request_body, "model": f"anthropic/{MODEL}"}, + api_key=KEY, client=client, stream=stream, model_info={"id": DEPLOYMENT}, + litellm_metadata={"user_api_key_hash": CALLER, "model_info": {"id": DEPLOYMENT}}, + ) + if stream: + async for _ in result: + pass + await asyncio.wait_for(observer.finished.wait(), timeout=5) + found = await lookup(cache, scope(), parse_prompt(request_body)) + if completed and not provider_error and transport != "unterminated": + assert found is not None + assert found.cached_tokens == 300 + assert before + 3600 <= found.expires_at <= time.time() + 3600 + else: + assert found is None + finally: + litellm.logging_callback_manager.remove_callback_from_all_lists(observer) + await client.client.aclose() diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_users.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_users.py new file mode 100644 index 00000000000..edd1d315093 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_users.py @@ -0,0 +1,122 @@ +"""The HTTP contract of `POST /management/v1/users/bulk`: envelope, problem documents and strict bodies. + +The batching behaviour itself is covered next to the helper, in +`tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py`, whose in-memory Prisma this reuses. +""" + +import pytest +from fastapi import FastAPI, Request +from fastapi.exceptions import RequestValidationError +from fastapi.testclient import TestClient + +from litellm.proxy._types import LitellmUserRoles, Member +from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth +from litellm.proxy.list_api.common import ManagementProblem, problem_response, request_validation_problem +from litellm.proxy.management_endpoints.management_v1 import router +from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX +from tests.test_litellm.proxy.management_helpers.test_bulk_user_creation import _FakePrisma, _License, _team + +app = FastAPI() + + +@app.exception_handler(ManagementProblem) +async def management_problem_exception_handler(request: Request, exc: ManagementProblem): + return problem_response(exc.problem) + + +@app.exception_handler(RequestValidationError) +async def validation_exception_handler(request: Request, exc: RequestValidationError): + return problem_response(request_validation_problem(exc.errors())) + + +app.include_router(router) +client = TestClient(app) + +USERS_BULK_PATH = f"{MANAGEMENT_V1_PREFIX}/users/bulk" + + +@pytest.fixture +def as_proxy_admin(): + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + yield + app.dependency_overrides.clear() + + +@pytest.fixture +def prisma(monkeypatch): + fake = _FakePrisma(teams=[_team("t1", [Member(user_id="existing", role="admin")])]) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake) + monkeypatch.setattr("litellm.proxy.proxy_server._license_check", _License()) + return fake + + +def _post(body: object): + return client.post(USERS_BULK_PATH, json=body, headers={"Authorization": "Bearer k"}) + + +def test_returns_one_result_per_row_in_order_inside_the_data_meta_envelope(prisma, as_proxy_admin): + response = _post( + { + "users": [ + {"user_id": "u1", "user_email": "a@example.com", "teams": ["t1"]}, + {"user_id": "u2", "teams": ["missing-team"]}, + {"user_id": "u3"}, + ] + } + ) + + assert response.status_code == 200 + body = response.json() + assert set(body) == {"data", "meta"} + assert body["meta"] == {"total_requested": 3, "created": 2, "failed": 1} + assert [row["user_id"] for row in body["data"]] == ["u1", "u2", "u3"] + assert [row["success"] for row in body["data"]] == [True, False, True] + assert body["data"][0]["teams"] == ["t1"] + assert "missing-team" in body["data"][1]["error"] + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["existing", "u1"] + + +def test_an_unknown_field_anywhere_in_the_body_is_a_422_problem(prisma, as_proxy_admin): + for body, field in ( + ({"users": [{"user_email": "a@example.com", "user_emial": "typo"}]}, "users.0.user_emial"), + ({"users": [{"user_email": "a@example.com"}], "dry_run": True}, "dry_run"), + ): + response = _post(body) + + assert response.status_code == 422, body + assert response.headers["content-type"] == "application/problem+json" + assert response.json()["type"] == "urn:litellm:error:invalid-request-body" + assert response.json()["detail"] == f"{field}: Extra inputs are not permitted" + assert prisma.db.litellm_usertable.rows == {} + + +def test_empty_and_oversized_batches_are_422_problems(prisma, as_proxy_admin): + for users in ([], [{"user_email": f"{i}@example.com"} for i in range(501)]): + response = _post({"users": users}) + + assert response.status_code == 422, len(users) + assert response.json()["type"] == "urn:litellm:error:invalid-request-body" + assert prisma.db.litellm_usertable.rows == {} + + +def test_license_limit_is_a_403_problem_and_creates_nothing(prisma, as_proxy_admin, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server._license_check", _License(max_users=1)) + + response = _post({"users": [{"user_id": "u1"}, {"user_id": "u2"}]}) + + assert response.status_code == 403 + assert response.headers["content-type"] == "application/problem+json" + assert response.json()["type"] == "urn:litellm:error:license-limit-exceeded" + assert prisma.db.litellm_usertable.rows == {} + + +def test_no_database_is_a_503_problem(as_proxy_admin, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + + response = _post({"users": [{"user_id": "u1"}]}) + + assert response.status_code == 503 + assert response.headers["content-type"] == "application/problem+json" + assert response.json()["type"] == "urn:litellm:error:database-not-connected" diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index da8fc760787..7352ca0e9ee 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -8,6 +8,10 @@ users can intentionally clear previously-set fields. """ from datetime import datetime, timezone +from types import SimpleNamespace + +from fastapi import HTTPException +from litellm import Router from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -1120,3 +1124,41 @@ class TestUpdateMetadataFieldsPremiumCheck: } _update_metadata_fields(updated_kv) mock_check.assert_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("db_model,stored_name,owner,public_name,error", [ + (False, None, None, None, None), + (True, "group", None, None, None), + (True, None, None, None, "Unknown deployment ID in router weights: id"), + (False, "renamed", None, None, "Deployment id does not belong to model group group"), + (False, None, "other-team", None, "Unknown deployment ID in router weights: id"), + (True, "internal", "team", "group", None), + (True, "group", "team", "public", "Deployment id does not belong to model group group"), + (True, "group", None, "unrelated-public-name", None), +]) +async def test_router_weights_validate_current_deployment_scope( + db_model: bool, stored_name: str | None, owner: str | None, + public_name: str | None, error: str | None, +) -> None: + from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights + + info = {"team_id": owner, "team_public_model_name": public_name} + router = Router(model_list=[{ + "model_name": "group", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "test"}, + "model_info": {"id": "id", "db_model": db_model, **info}, + }]) + rows = [SimpleNamespace(model_id="id", model_name=stored_name, model_info=info)] if stored_name else [] + table = SimpleNamespace(find_many=AsyncMock(return_value=rows)) + db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)) + validation = validate_router_settings_weights( + {"weights": {"group": {"id": 1}}}, team_id="team", prisma_client=db, llm_router=router, + ) + if error: + with pytest.raises(HTTPException, match=error) as exc: + await validation + assert exc.value.status_code == 400 + assert exc.value.detail == error + else: + await validation diff --git a/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py index dcbe515d5de..8382a5ada96 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py @@ -2,10 +2,8 @@ Unit tests for compliance check endpoints (EU AI Act and GDPR). """ - import pytest - from litellm.proxy.compliance_checks import ComplianceChecker from litellm.types.proxy.compliance_endpoints import ComplianceCheckRequest @@ -591,3 +589,37 @@ class TestModeMatching: continue if matched: assert mode in _guaranteed_modes(g_mode), (g_mode, mode) + + +class TestNotRunGuardrails: + """LIT-6314 logs a not_run entry for a guardrail that message scoping left nothing to scan.""" + + def test_not_run_alone_never_evidences_compliance(self): + data = ComplianceCheckRequest( + request_id="req-601", + user_id="user-1", + model="gpt-4", + timestamp="2026-02-17T00:00:00Z", + guardrail_information=[ + {"guardrail_name": "pii_detection", "guardrail_status": "not_run", "guardrail_mode": "pre_call"}, + ], + ) + results = {c.check_name: c.passed for c in ComplianceChecker(data).check_eu_ai_act()} + assert results["Guardrails applied"] is False + assert results["Content screened before LLM"] is False + assert results["Audit record complete"] is False + + def test_not_run_sibling_does_not_fail_a_passing_request(self): + data = ComplianceCheckRequest( + request_id="req-602", + user_id="user-1", + model="gpt-4", + timestamp="2026-02-17T00:00:00Z", + pii_detected=True, + guardrail_information=[ + {"guardrail_name": "pii_detection", "guardrail_status": "success", "guardrail_mode": "pre_call"}, + {"guardrail_name": "system_only", "guardrail_status": "not_run", "guardrail_mode": "pre_call"}, + ], + ) + results = {c.check_name: c.passed for c in ComplianceChecker(data).check_gdpr()} + assert results["Sensitive data protected"] is True diff --git a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py index 7ece35ceedf..c73d29e78b2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py +++ b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py @@ -975,8 +975,9 @@ class TestEstimateCostCacheAndReasoningTokens: @pytest.mark.asyncio async def test_a_model_without_cache_or_reasoning_prices_estimates_what_the_proxy_bills(self, monkeypatch): - """The cost calculator bills cache tokens of a cost-map model without cache prices at zero - and its reasoning tokens at the output rate. The estimate reports those effective rates.""" + """The cost calculator bills cache reads of a cost-map model without cache prices at zero, + its cache writes at the input rate, and its reasoning tokens at the output rate. The estimate + reports those effective rates.""" monkeypatch.setitem( litellm.model_cost, A_MAPPED_MODEL, @@ -986,12 +987,14 @@ class TestEstimateCostCacheAndReasoningTokens: response = await _estimate_with_cache_and_reasoning(None, model=A_MAPPED_MODEL) assert response.cache_read_cost_per_request == 0.0 - assert response.cache_creation_cost_per_request == 0.0 + assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 5e-6) assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 6e-6) - assert response.input_cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6) - assert response.cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6 + OUTPUT_TOKENS * 6e-6) + assert response.input_cost_per_request == pytest.approx((TEXT_INPUT_TOKENS + CACHE_CREATION_TOKENS) * 5e-6) + assert response.cost_per_request == pytest.approx( + (TEXT_INPUT_TOKENS + CACHE_CREATION_TOKENS) * 5e-6 + OUTPUT_TOKENS * 6e-6 + ) assert response.cache_read_input_token_cost == 0.0 - assert response.cache_creation_input_token_cost == 0.0 + assert response.cache_creation_input_token_cost == pytest.approx(5e-6) assert response.output_cost_per_reasoning_token == pytest.approx(6e-6) @pytest.mark.asyncio 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 2ac52da57df..c590a24203c 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 @@ -1,4 +1,7 @@ +from collections.abc import Mapping +from contextlib import ExitStack from typing import Final +from types import SimpleNamespace import json from datetime import datetime, timedelta, timezone @@ -17,27 +20,36 @@ from litellm.proxy._types import ( GenerateKeyRequest, NewUserRequest, LiteLLM_BudgetTable, + LiteLLM_ObjectPermissionBase, LiteLLM_OrganizationTable, LiteLLM_ProjectTableCachedObj, LiteLLM_TeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LiteLLM_VerificationToken, + LiteLLMKeyType, LitellmUserRoles, Member, ProxyException, + RegenerateKeyRequest, ResetSpendRequest, UpdateKeyRequest, ) +from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.proxy.auth.auth_checks import _delete_cache_key_object, _project_cache_key from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_org_key_limits, _check_project_key_limits, _check_team_key_limits, _common_key_generation_helper, + _effective_key_after_update, + _effective_key_for_generate, + _enforce_custom_key_policy, _enforce_upperbound_key_params, + _execute_virtual_key_regeneration, _get_and_validate_existing_key, _list_key_helper, _persist_deleted_verification_tokens, @@ -62,6 +74,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( validate_key_team_change, ) from litellm.proxy.proxy_server import app +from litellm.types.proxy.management_endpoints.key_management_endpoints import CustomKeyPolicyRequest client = TestClient(app) @@ -1026,7 +1039,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): @pytest.mark.asyncio -async def test_key_update_object_permissions_existing_permission(monkeypatch): +async def test_key_update_object_permissions_existing_permission(): """ Test updating object permissions when a key already has an existing object_permission_id. @@ -1046,9 +1059,7 @@ async def test_key_update_object_permissions_existing_permission(monkeypatch): _handle_update_object_permission, ) - # Mock prisma client mock_prisma_client = AsyncMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock existing key with object_permission_id existing_key_row = LiteLLM_VerificationToken( @@ -1088,6 +1099,7 @@ async def test_key_update_object_permissions_existing_permission(monkeypatch): result = await _handle_update_object_permission( data_json=data_json, existing_key_row=existing_key_row, + prisma_client=mock_prisma_client, ) # Verify the object_permission was removed from data_json and object_permission_id was set @@ -1102,7 +1114,7 @@ async def test_key_update_object_permissions_existing_permission(monkeypatch): @pytest.mark.asyncio -async def test_key_update_object_permissions_no_existing_permission(monkeypatch): +async def test_key_update_object_permissions_no_existing_permission(): """ Test creating object permissions when a key has no existing object_permission_id. @@ -1122,9 +1134,7 @@ async def test_key_update_object_permissions_no_existing_permission(monkeypatch) _handle_update_object_permission, ) - # Mock prisma client mock_prisma_client = AsyncMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) existing_key_row_no_perm = LiteLLM_VerificationToken( token="test_token_hash_2", @@ -1155,6 +1165,7 @@ async def test_key_update_object_permissions_no_existing_permission(monkeypatch) result = await _handle_update_object_permission( data_json=data_json, existing_key_row=existing_key_row_no_perm, + prisma_client=mock_prisma_client, ) # Verify new object_permission_id was set @@ -1165,7 +1176,7 @@ async def test_key_update_object_permissions_no_existing_permission(monkeypatch) @pytest.mark.asyncio -async def test_key_update_object_permissions_missing_permission_record(monkeypatch): +async def test_key_update_object_permissions_missing_permission_record(): """ Test creating object permissions when existing object_permission_id record is not found. @@ -1185,9 +1196,7 @@ async def test_key_update_object_permissions_missing_permission_record(monkeypat _handle_update_object_permission, ) - # Mock prisma client mock_prisma_client = AsyncMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) existing_key_row_missing_perm = LiteLLM_VerificationToken( token="test_token_hash_3", @@ -1218,6 +1227,7 @@ async def test_key_update_object_permissions_missing_permission_record(monkeypat result = await _handle_update_object_permission( data_json=data_json, existing_key_row=existing_key_row_missing_perm, + prisma_client=mock_prisma_client, ) # Verify new object_permission_id was set @@ -6615,6 +6625,9 @@ async def test_generate_key_with_router_settings(monkeypatch): return_value=[] ) mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[ + SimpleNamespace(model_id="weighted-id", model_name="gpt-4", model_info={}) + ]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -6630,6 +6643,7 @@ async def test_generate_key_with_router_settings(monkeypatch): "routing_strategy": "usage-based", "num_retries": 3, "model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}, + "weights": {"gpt-4": {"weighted-id": 1}}, } request_data = GenerateKeyRequest( @@ -6679,21 +6693,37 @@ async def test_generate_key_with_router_settings(monkeypatch): # Verify router_settings matches input (regardless of serialization state) assert actual_settings == router_settings_data + mock_prisma_client.insert_data.reset_mock() + with pytest.raises(ProxyException, match="Unknown deployment ID"): + await generate_key_fn( + data=GenerateKeyRequest(router_settings={"weights": {"gpt-4": {"unknown-id": 1}}}), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="user-router-1"), + ) + mock_prisma_client.insert_data.assert_not_awaited() @pytest.mark.asyncio -async def test_update_key_with_router_settings(monkeypatch): +@pytest.mark.parametrize("request_type", [UpdateKeyRequest, RegenerateKeyRequest]) +@pytest.mark.parametrize("target_team", ["new-team", None]) +async def test_update_key_with_router_settings( + monkeypatch: pytest.MonkeyPatch, + request_type: type[UpdateKeyRequest | RegenerateKeyRequest], target_team: str | None, +) -> None: """ Test that /key/update correctly handles router_settings by: 1. Accepting router_settings as a dict parameter 2. Serializing router_settings to JSON when updating database 3. Updating router_settings in the key record """ - from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest + from litellm.proxy._types import LiteLLM_VerificationToken from litellm.proxy.management_endpoints.key_management_endpoints import ( prepare_key_update_data, ) + model = SimpleNamespace(model_id="weighted-id", model_name="gpt-4", model_info={}) + table = SimpleNamespace(find_many=AsyncMock(return_value=[model])) + db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)) + # Mock existing key existing_key = LiteLLM_VerificationToken( token="test-token-router", @@ -6710,14 +6740,16 @@ async def test_update_key_with_router_settings(monkeypatch): router_settings_data = { "routing_strategy": "latency-based", "num_retries": 2, + "weights": {"gpt-4": {"weighted-id": 1}}, } - update_request = UpdateKeyRequest( + update_request = request_type( key="test-token-router", router_settings=router_settings_data ) result = await prepare_key_update_data( - data=update_request, existing_key_row=existing_key + data=update_request, existing_key_row=existing_key, + prisma_client=db, llm_router=None, ) # Verify router_settings is serialized to JSON string @@ -6728,6 +6760,28 @@ async def test_update_key_with_router_settings(monkeypatch): deserialized_settings = json.loads(result["router_settings"]) assert deserialized_settings == router_settings_data + with pytest.raises(HTTPException, match="Unknown deployment ID"): + await prepare_key_update_data( + request_type(key=existing_key.token, router_settings={"weights": {"gpt-4": {"unknown-id": 1}}}), + existing_key, + prisma_client=db, llm_router=None, + ) + existing_key.team_id = "old-team" + existing_key.router_settings = router_settings_data + move = request_type(key=existing_key.token, team_id=target_team) + retained = await prepare_key_update_data(move, existing_key, prisma_client=db, llm_router=None) + assert retained["team_id"] == target_team + assert "router_settings" not in retained + model.model_info = {"team_id": "old-team"} + with pytest.raises(HTTPException, match="Unknown deployment ID"): + await prepare_key_update_data(move, existing_key, prisma_client=db, llm_router=None) + cleared = await prepare_key_update_data( + request_type(key=existing_key.token, team_id=target_team, router_settings={}), existing_key, + prisma_client=db, llm_router=None, + ) + assert cleared["team_id"] == target_team + assert json.loads(cleared["router_settings"]) == {} + @pytest.mark.asyncio async def test_validate_max_budget(): @@ -11935,6 +11989,10 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monk "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", new_callable=AsyncMock, ), + patch( # test-quality-ok: archival path is outside upperbound rejection + "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + new_callable=AsyncMock, + ) as persist_deleted_verification_tokens, patch( "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", new_callable=AsyncMock, @@ -11955,6 +12013,7 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monk assert exc_info.value.status_code == 400 assert "duration" in str(exc_info.value.detail) # Rejected regenerate must not reach the DB update. + persist_deleted_verification_tokens.assert_not_awaited() assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 @@ -12014,6 +12073,1011 @@ async def test_execute_virtual_key_regeneration_allows_within_limit_duration(mon assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 +@pytest.mark.asyncio +async def test_execute_virtual_key_regeneration_rejects_when_custom_key_update_hook_denies(): + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(duration="3000d") + mock_prisma_client = _make_regenerate_mock_prisma() + received_data: list[UpdateKeyRequest] = [] + + async def hook(data: UpdateKeyRequest) -> dict[str, object]: + received_data.append(data) + if data.duration and duration_in_seconds(data.duration) > duration_in_seconds("7d"): + return {"decision": False, "message": "duration must be <= 7d"} + return {"decision": True} + + with ( + patch( # test-quality-ok: deterministic token setup for policy rejection + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( # test-quality-ok: grace-period path is outside policy rejection + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ) as insert_deprecated_key, + patch( # test-quality-ok: archival path is outside policy rejection + "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + new_callable=AsyncMock, + ) as persist_deleted_verification_tokens, + patch( # test-quality-ok: cache eviction is outside policy rejection + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: rotation callback is outside policy rejection + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.user_custom_key_update", hook), # test-quality-ok: inject policy hook + ): + with pytest.raises(HTTPException) as exc_info: + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == "duration must be <= 7d" + insert_deprecated_key.assert_not_awaited() + persist_deleted_verification_tokens.assert_not_awaited() + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + assert len(received_data) == 1 + assert received_data[0].key == "abc123" + assert received_data[0].duration == "3000d" + + +@pytest.mark.asyncio +async def test_execute_virtual_key_regeneration_allows_when_custom_key_update_hook_approves(): + existing_key = _make_regenerate_existing_key() + data = RegenerateKeyRequest(duration="5d") + mock_prisma_client = _make_regenerate_mock_prisma() + received_data: list[UpdateKeyRequest] = [] + + async def hook(data: UpdateKeyRequest) -> dict[str, object]: + received_data.append(data) + if data.duration and duration_in_seconds(data.duration) > duration_in_seconds("7d"): + return {"decision": False, "message": "duration must be <= 7d"} + return {"decision": True} + + with ( + patch( # test-quality-ok: deterministic token setup for policy approval + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( # test-quality-ok: grace-period path is outside policy approval + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: verify archival follows policy approval + "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + new_callable=AsyncMock, + ) as persist_deleted_verification_tokens, + patch( # test-quality-ok: cache eviction is outside policy approval + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: rotation callback is outside policy approval + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.user_custom_key_update", hook), # test-quality-ok: inject policy hook + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 + persist_deleted_verification_tokens.assert_awaited_once() + assert persist_deleted_verification_tokens.call_args.kwargs["keys"] == [existing_key] + assert len(received_data) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "data", + [None, RegenerateKeyRequest(), RegenerateKeyRequest(duration=""), RegenerateKeyRequest(budget_duration="")], +) +async def test_execute_virtual_key_regeneration_skips_custom_key_update_hook_without_changes(data): + mock_prisma_client = _make_regenerate_mock_prisma() + + async def hook(data: UpdateKeyRequest) -> dict[str, object]: + raise AssertionError(f"custom key update hook called with {data}") + + with ( + patch( # test-quality-ok: deterministic token setup for unchanged request + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( # test-quality-ok: grace-period path is outside unchanged request + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: cache eviction is outside unchanged request + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: rotation callback is outside unchanged request + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.user_custom_key_update", hook), # test-quality-ok: inject policy hook + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=_make_regenerate_existing_key(), + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 + + +@pytest.mark.asyncio +async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry_from_the_custom_key_update_hook(): + mock_prisma_client = _make_regenerate_mock_prisma() + untouched_modal_body = RegenerateKeyRequest( + key_alias=None, max_budget=None, tpm_limit=None, rpm_limit=None, duration="", grace_period="" + ) + received_data: list[UpdateKeyRequest] = [] + + async def hook(data: UpdateKeyRequest) -> dict[str, object]: + received_data.append(data) + if data.duration is not None and duration_in_seconds(data.duration) > duration_in_seconds("7d"): + return {"decision": False, "message": "duration must be <= 7d"} + return {"decision": True} + + with ( + patch( # test-quality-ok: deterministic token setup for the untouched modal body + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + patch( # test-quality-ok: grace-period path is outside the hook input + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: cache eviction is outside the hook input + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: rotation callback is outside the hook input + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.user_custom_key_update", hook), # test-quality-ok: inject policy hook + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=_make_regenerate_existing_key(), + hashed_api_key="abc123", + key="abc123", + data=untouched_modal_body, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 + assert len(received_data) == 1 + assert "duration" not in received_data[0].model_fields_set + assert received_data[0].model_fields_set >= {"key", "key_alias", "max_budget", "tpm_limit", "rpm_limit"} + + +_POLICY_DENIAL_MESSAGE = "key duration must be 7d or less" +_POLICY_HASHED_TOKEN = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" +_POLICY_GENERATED_KEY = {"key": "sk-test-key", "expires": None, "user_id": "test-user", "team_id": None} + + +def _seven_day_policy(received: list[CustomKeyPolicyRequest]): + async def policy(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + received.append(policy_request) + expires = policy_request.effective_key.expires + if isinstance(expires, datetime) and expires > datetime.now(timezone.utc) + timedelta(days=7): + return {"decision": False, "message": _POLICY_DENIAL_MESSAGE} + return {"decision": True} + + return policy + + +def _assert_expires_in(effective_key: LiteLLM_VerificationToken, duration: str) -> None: + expires = effective_key.expires + assert isinstance(expires, datetime) + assert expires.tzinfo is not None + expected = datetime.now(timezone.utc) + timedelta(seconds=duration_in_seconds(duration=duration)) + assert abs((expires - expected).total_seconds()) < 60 + + +def _regenerate_policy_mocks(policy, insert_deprecated_key: AsyncMock, persist: AsyncMock) -> ExitStack: + stack = ExitStack() + stack.enter_context( + patch( # test-quality-ok: deterministic token setup for the policy path + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ) + ) + stack.enter_context( + patch( # test-quality-ok: grace-period write must not run on a denied regenerate + "litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key", + insert_deprecated_key, + ) + ) + stack.enter_context( + patch( # test-quality-ok: archival write must not run on a denied regenerate + "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + persist, + ) + ) + stack.enter_context( + patch( # test-quality-ok: cache eviction is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ) + ) + stack.enter_context( + patch( # test-quality-ok: rotation callback is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook", + new_callable=AsyncMock, + ) + ) + stack.enter_context( + patch("litellm.proxy.proxy_server.user_custom_key_policy", policy) # test-quality-ok: inject policy hook + ) + return stack + + +async def _regenerate_under_policy(mock_prisma_client, existing_key, data): + return await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + +@pytest.mark.asyncio +async def test_regenerate_rejects_when_custom_key_policy_denies_the_effective_expiry(): + existing_key = _make_regenerate_existing_key() + mock_prisma_client = _make_regenerate_mock_prisma() + received: list[CustomKeyPolicyRequest] = [] + insert_deprecated_key = AsyncMock() + persist = AsyncMock() + + with _regenerate_policy_mocks(_seven_day_policy(received), insert_deprecated_key, persist): + with pytest.raises(HTTPException) as exc_info: + await _regenerate_under_policy(mock_prisma_client, existing_key, RegenerateKeyRequest(duration="3000d")) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == _POLICY_DENIAL_MESSAGE + insert_deprecated_key.assert_not_awaited() + persist.assert_not_awaited() + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + assert [policy_request.operation for policy_request in received] == ["regenerate"] + assert received[0].existing_key is not None + assert received[0].existing_key.token == "abc123" + assert isinstance(received[0].request, RegenerateKeyRequest) + assert received[0].request.duration == "3000d" + _assert_expires_in(received[0].effective_key, "3000d") + + +@pytest.mark.asyncio +async def test_regenerate_within_custom_key_policy_rotates_the_key(): + existing_key = _make_regenerate_existing_key() + mock_prisma_client = _make_regenerate_mock_prisma() + received: list[CustomKeyPolicyRequest] = [] + persist = AsyncMock() + + with _regenerate_policy_mocks(_seven_day_policy(received), AsyncMock(), persist): + await _regenerate_under_policy(mock_prisma_client, existing_key, RegenerateKeyRequest(duration="5d")) + + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 1 + persist.assert_awaited_once() + assert persist.call_args.kwargs["keys"] == [existing_key] + assert [policy_request.operation for policy_request in received] == ["regenerate"] + _assert_expires_in(received[0].effective_key, "5d") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("data", [None, RegenerateKeyRequest()]) +async def test_regenerate_without_changes_still_runs_custom_key_policy(data): + existing_key = _make_regenerate_existing_key() + mock_prisma_client = _make_regenerate_mock_prisma() + received: list[CustomKeyPolicyRequest] = [] + + async def freeze_rotation(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + received.append(policy_request) + return {"decision": False, "message": "key rotation is frozen"} + + with _regenerate_policy_mocks(freeze_rotation, AsyncMock(), AsyncMock()): + with pytest.raises(HTTPException) as exc_info: + await _regenerate_under_policy(mock_prisma_client, existing_key, data) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == "key rotation is frozen" + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + assert [policy_request.operation for policy_request in received] == ["regenerate"] + assert received[0].existing_key == existing_key + assert received[0].effective_key == existing_key + + +def _policy_existing_team_key() -> LiteLLM_VerificationToken: + return LiteLLM_VerificationToken( + token=_POLICY_HASHED_TOKEN, user_id="test-user", team_id="team-a", max_budget=200.0 + ) + + +def _setup_update_key_fn_policy_mocks(monkeypatch, existing_key: LiteLLM_VerificationToken) -> AsyncMock: + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(return_value=None) + mock_prisma_client.update_data = AsyncMock(return_value={"data": {"max_budget": 50.0, "team_id": "team-a"}}) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", AsyncMock(return_value=None) + ) + return mock_prisma_client + + +def _assert_update_policy_request(policy_request: CustomKeyPolicyRequest, request: UpdateKeyRequest) -> None: + assert policy_request.operation == "update" + assert policy_request.request is request + assert policy_request.existing_key is not None + assert policy_request.existing_key.max_budget == 200.0 + assert policy_request.effective_key.team_id == "team-a" + assert policy_request.effective_key.user_id == "test-user" + assert policy_request.effective_key.max_budget == 50.0 + _assert_expires_in(policy_request.effective_key, request.duration or "") + + +@pytest.mark.asyncio +async def test_update_key_fn_runs_custom_key_policy_on_the_effective_row(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn + + mock_prisma_client = _setup_update_key_fn_policy_mocks(monkeypatch, _policy_existing_team_key()) + received: list[CustomKeyPolicyRequest] = [] + policy = _seven_day_policy(received) + data = UpdateKeyRequest( + key=_POLICY_HASHED_TOKEN, duration="5d", max_budget=50.0, auto_rotate=True, rotation_interval="30d" + ) + + with ( + patch( # test-quality-ok: cache eviction is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch("litellm.proxy.proxy_server.user_custom_key_policy", policy), # test-quality-ok: inject policy hook + ): + await update_key_fn( + request=MagicMock(), + data=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + ) + + mock_prisma_client.update_data.assert_awaited_once() + assert len(received) == 1 + _assert_update_policy_request(received[0], data) + key_rotation_at = received[0].effective_key.key_rotation_at + assert key_rotation_at is not None + assert abs(key_rotation_at - (datetime.now(timezone.utc) + timedelta(days=30))) < timedelta(seconds=60) + + +@pytest.mark.asyncio +async def test_update_key_fn_rejects_when_custom_key_policy_denies(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn + + mock_prisma_client = _setup_update_key_fn_policy_mocks(monkeypatch, _policy_existing_team_key()) + received: list[CustomKeyPolicyRequest] = [] + policy = _seven_day_policy(received) + + with patch("litellm.proxy.proxy_server.user_custom_key_policy", policy): # test-quality-ok: inject policy hook + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key=_POLICY_HASHED_TOKEN, duration="3000d", max_budget=50.0), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "403" + assert exc_info.value.message == _POLICY_DENIAL_MESSAGE + mock_prisma_client.update_data.assert_not_awaited() + assert [policy_request.operation for policy_request in received] == ["update"] + _assert_expires_in(received[0].effective_key, "3000d") + + +async def _process_single_key_update_under_policy(prisma_client: AsyncMock, data: UpdateKeyRequest, policy): + with ( + patch( # test-quality-ok: cache eviction is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: update callback is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook", + new_callable=AsyncMock, + ), + ): + return await _process_single_key_update( + update_key_request=data, + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + prisma_client=prisma_client, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + llm_router=None, + existing_key_row=_policy_existing_team_key(), + user_custom_key_policy=policy, + ) + + +@pytest.mark.asyncio +async def test_process_single_key_update_runs_custom_key_policy_on_the_effective_row(): + mock_prisma_client = AsyncMock() + updated_row = MagicMock() + updated_row.model_dump.return_value = {"max_budget": 50.0, "team_id": "team-a"} + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) + received: list[CustomKeyPolicyRequest] = [] + data = UpdateKeyRequest(key=_POLICY_HASHED_TOKEN, duration="5d", max_budget=50.0) + + result = await _process_single_key_update_under_policy(mock_prisma_client, data, _seven_day_policy(received)) + + assert result["max_budget"] == 50.0 + mock_prisma_client.update_data.assert_awaited_once() + assert len(received) == 1 + _assert_update_policy_request(received[0], data) + + +@pytest.mark.asyncio +async def test_process_single_key_update_rejects_when_custom_key_policy_denies(): + mock_prisma_client = AsyncMock() + mock_prisma_client.update_data = AsyncMock() + received: list[CustomKeyPolicyRequest] = [] + data = UpdateKeyRequest(key=_POLICY_HASHED_TOKEN, duration="3000d", max_budget=50.0) + + with pytest.raises(HTTPException) as exc_info: + await _process_single_key_update_under_policy(mock_prisma_client, data, _seven_day_policy(received)) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == _POLICY_DENIAL_MESSAGE + mock_prisma_client.update_data.assert_not_awaited() + assert [policy_request.operation for policy_request in received] == ["update"] + + +_OBJECT_PERMISSION_ID_AFTER_POLICY = "perm-after-policy" + + +def _record_object_permission_writes(mock_prisma_client: AsyncMock, events: list[str]) -> None: + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + + async def upsert(**_kwargs: object) -> MagicMock: + events.append("permission row upsert") + return MagicMock(object_permission_id=_OBJECT_PERMISSION_ID_AFTER_POLICY) + + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(side_effect=upsert) + + +def _recording_policy(events: list[str], allowed: bool): + async def policy(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + events.append("policy") + return {"decision": allowed, "message": "key max_budget must be 1000 or less"} + + return policy + + +def _assert_permission_row_written_after_policy(events: list[str], written: Mapping[str, object]) -> None: + assert events == ["policy", "permission row upsert"] + assert written["object_permission_id"] == _OBJECT_PERMISSION_ID_AFTER_POLICY + assert "object_permission" not in written + + +def _assert_permission_row_untouched(mock_prisma_client: AsyncMock, events: list[str]) -> None: + assert events == ["policy"] + mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + + +def _update_with_object_permission(max_budget: float) -> UpdateKeyRequest: + return UpdateKeyRequest( + key=_POLICY_HASHED_TOKEN, + max_budget=max_budget, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["vs-1"]), + ) + + +def _setup_update_key_fn_object_permission_mocks(monkeypatch, allowed: bool) -> tuple[AsyncMock, list[str]]: + mock_prisma_client = _setup_update_key_fn_policy_mocks(monkeypatch, _policy_existing_team_key()) + events: list[str] = [] + _record_object_permission_writes(mock_prisma_client, events) + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_policy", _recording_policy(events, allowed)) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", AsyncMock() + ) + return mock_prisma_client, events + + +async def _update_key_fn_with_object_permission(max_budget: float): + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn + + return await update_key_fn( + request=MagicMock(), + data=_update_with_object_permission(max_budget=max_budget), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + ) + + +@pytest.mark.asyncio +async def test_update_key_fn_writes_the_object_permission_row_only_after_the_policy_allows(monkeypatch): + mock_prisma_client, events = _setup_update_key_fn_object_permission_mocks(monkeypatch, allowed=True) + + await _update_key_fn_with_object_permission(max_budget=50.0) + + _assert_permission_row_written_after_policy(events, mock_prisma_client.update_data.await_args.kwargs["data"]) + + +@pytest.mark.asyncio +async def test_update_key_fn_denied_by_the_policy_leaves_the_object_permission_row_untouched(monkeypatch): + mock_prisma_client, events = _setup_update_key_fn_object_permission_mocks(monkeypatch, allowed=False) + + with pytest.raises(ProxyException) as exc_info: + await _update_key_fn_with_object_permission(max_budget=5000.0) + + assert str(exc_info.value.code) == "403" + _assert_permission_row_untouched(mock_prisma_client, events) + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_process_single_key_update_writes_the_object_permission_row_only_after_the_policy_allows(): + mock_prisma_client = AsyncMock() + updated_row = MagicMock() + updated_row.model_dump.return_value = {"max_budget": 50.0, "team_id": "team-a"} + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) + events: list[str] = [] + _record_object_permission_writes(mock_prisma_client, events) + + await _process_single_key_update_under_policy( + mock_prisma_client, _update_with_object_permission(max_budget=50.0), _recording_policy(events, allowed=True) + ) + + _assert_permission_row_written_after_policy(events, mock_prisma_client.update_data.await_args.kwargs["data"]) + + +@pytest.mark.asyncio +async def test_process_single_key_update_denied_by_the_policy_leaves_the_object_permission_row_untouched(): + mock_prisma_client = AsyncMock() + mock_prisma_client.update_data = AsyncMock() + events: list[str] = [] + _record_object_permission_writes(mock_prisma_client, events) + + with pytest.raises(HTTPException) as exc_info: + await _process_single_key_update_under_policy( + mock_prisma_client, _update_with_object_permission(max_budget=5000.0), _recording_policy(events, allowed=False) + ) + + assert exc_info.value.status_code == 403 + _assert_permission_row_untouched(mock_prisma_client, events) + mock_prisma_client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_regenerate_writes_the_object_permission_row_only_after_the_policy_allows(): + mock_prisma_client = _make_regenerate_mock_prisma() + events: list[str] = [] + _record_object_permission_writes(mock_prisma_client, events) + data = RegenerateKeyRequest(max_budget=50.0, object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["vs-1"])) + + with _regenerate_policy_mocks(_recording_policy(events, allowed=True), AsyncMock(), AsyncMock()): + await _regenerate_under_policy(mock_prisma_client, _make_regenerate_existing_key(), data) + + _assert_permission_row_written_after_policy( + events, mock_prisma_client.db.litellm_verificationtoken.update.await_args.kwargs["data"] + ) + + +@pytest.mark.asyncio +async def test_regenerate_denied_by_the_policy_leaves_the_object_permission_row_untouched(): + mock_prisma_client = _make_regenerate_mock_prisma() + events: list[str] = [] + _record_object_permission_writes(mock_prisma_client, events) + data = RegenerateKeyRequest(max_budget=5000.0, object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["vs-1"])) + + with _regenerate_policy_mocks(_recording_policy(events, allowed=False), AsyncMock(), AsyncMock()): + with pytest.raises(HTTPException) as exc_info: + await _regenerate_under_policy(mock_prisma_client, _make_regenerate_existing_key(), data) + + assert exc_info.value.status_code == 403 + _assert_permission_row_untouched(mock_prisma_client, events) + assert mock_prisma_client.db.litellm_verificationtoken.update.await_count == 0 + + +@pytest.mark.asyncio +async def test_bulk_update_keys_runs_custom_key_policy_per_key(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateKeyRequest, + BulkUpdateKeyRequestItem, + ) + + existing_keys = [ + LiteLLM_VerificationToken(token="test-key-1", user_id="user-123", max_budget=None), + LiteLLM_VerificationToken(token="test-key-2", user_id="user-123", max_budget=50.0), + ] + updated_row = MagicMock() + updated_row.model_dump.return_value = {"user_id": "user-123", "max_budget": 100.0} + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(side_effect=existing_keys) + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) + mock_prisma_client.get_data = AsyncMock(return_value=None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + received: list[CustomKeyPolicyRequest] = [] + + async def cap_max_budget(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + received.append(policy_request) + max_budget = policy_request.effective_key.max_budget + if max_budget is not None and max_budget > 100: + return {"decision": False, "message": "max_budget must be 100 or less"} + return {"decision": True} + + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_policy", cap_max_budget) + + with ( + patch( # test-quality-ok: cache eviction is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: update callback is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook", + new_callable=AsyncMock, + ), + ): + response = await bulk_update_keys( + data=BulkUpdateKeyRequest( + keys=[ + BulkUpdateKeyRequestItem(key="test-key-1", max_budget=100.0), + BulkUpdateKeyRequestItem(key="test-key-2", max_budget=500.0), + ] + ), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + ) + + assert [update.key for update in response.successful_updates] == ["test-key-1"] + assert [(failed.key, failed.failed_reason) for failed in response.failed_updates] == [ + ("test-key-2", "max_budget must be 100 or less") + ] + assert mock_prisma_client.update_data.await_count == 1 + assert [policy_request.operation for policy_request in received] == ["update", "update"] + assert [policy_request.effective_key.max_budget for policy_request in received] == [100.0, 500.0] + assert [ + policy_request.existing_key.max_budget if policy_request.existing_key is not None else "missing" + for policy_request in received + ] == [None, 50.0] + + +def _policy_generate_prisma() -> MagicMock: + mock_prisma = MagicMock() + mock_prisma.db.litellm_budgettable.create = AsyncMock(return_value=MagicMock(budget_id="budget-1")) + mock_prisma.jsonify_object = MagicMock(side_effect=lambda data: json.loads(data) if isinstance(data, str) else data) + return mock_prisma + + +def _generate_policy_mocks(mock_prisma: MagicMock, generate_key_helper: AsyncMock, policy) -> ExitStack: + stack = ExitStack() + stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)) # test-quality-ok: fake DB + stack.enter_context(patch("litellm.proxy.proxy_server.llm_router", None)) # test-quality-ok: no router in test + stack.enter_context(patch("litellm.proxy.proxy_server.premium_user", True)) # test-quality-ok: premium fields + stack.enter_context(patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")) # test-quality-ok: admin + stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())) # test-quality-ok: cache + stack.enter_context( + patch( # test-quality-ok: the key write must not run on a denied generate + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + generate_key_helper, + ) + ) + stack.enter_context( + patch("litellm.proxy.proxy_server.user_custom_key_policy", policy) # test-quality-ok: inject policy hook + ) + return stack + + +def _generate_request(duration: str, organization_id: str | None) -> GenerateKeyRequest: + return GenerateKeyRequest( + duration=duration, + organization_id=organization_id, + guardrails=["g1"], + tags=["t1"], + soft_budget=10.0, + max_budget=20.0, + ) + + +def _assert_generate_policy_request( + policy_request: CustomKeyPolicyRequest, duration: str, organization_id: str | None +) -> None: + assert policy_request.operation == "generate" + assert policy_request.existing_key is None + assert policy_request.effective_key.org_id == organization_id + assert policy_request.effective_key.max_budget == 20.0 + assert policy_request.effective_key.metadata["guardrails"] == ["g1"] + assert policy_request.effective_key.metadata["tags"] == ["t1"] + _assert_expires_in(policy_request.effective_key, duration) + + +@pytest.mark.asyncio +async def test_generate_key_rejects_when_custom_key_policy_denies_before_any_write(): + mock_prisma = _policy_generate_prisma() + generate_key_helper = AsyncMock(return_value=_POLICY_GENERATED_KEY) + received: list[CustomKeyPolicyRequest] = [] + data = _generate_request("3000d", organization_id="org-1") + + with _generate_policy_mocks(mock_prisma, generate_key_helper, _seven_day_policy(received)): + with pytest.raises(ProxyException) as exc_info: + await generate_key_fn( + data=data, user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None + ) + + assert str(exc_info.value.code) == "403" + assert exc_info.value.message == _POLICY_DENIAL_MESSAGE + mock_prisma.db.litellm_budgettable.create.assert_not_awaited() + generate_key_helper.assert_not_awaited() + assert len(received) == 1 + _assert_generate_policy_request(received[0], "3000d", organization_id="org-1") + assert received[0].request is data + assert data.duration == "3000d" + assert data.guardrails == ["g1"] + assert data.tags == ["t1"] + assert data.organization_id == "org-1" + + +@pytest.mark.asyncio +async def test_generate_key_within_custom_key_policy_creates_the_key(): + mock_prisma = _policy_generate_prisma() + generate_key_helper = AsyncMock(return_value=_POLICY_GENERATED_KEY) + received: list[CustomKeyPolicyRequest] = [] + data = _generate_request("5d", organization_id=None) + + with _generate_policy_mocks(mock_prisma, generate_key_helper, _seven_day_policy(received)): + await generate_key_fn( + data=data, user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None + ) + + mock_prisma.db.litellm_budgettable.create.assert_awaited_once() + generate_key_helper.assert_awaited_once() + assert len(received) == 1 + _assert_generate_policy_request(received[0], "5d", organization_id=None) + assert received[0].request is data + + +@pytest.mark.asyncio +async def test_service_account_generate_rejects_when_custom_key_policy_denies(): + from litellm.proxy.management_endpoints.key_management_endpoints import generate_service_account_key_fn + + mock_prisma = _policy_generate_prisma() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=MagicMock()) + generate_key_helper = AsyncMock(return_value=_POLICY_GENERATED_KEY) + received: list[CustomKeyPolicyRequest] = [] + + with ( + _generate_policy_mocks(mock_prisma, generate_key_helper, _seven_day_policy(received)), + patch( # test-quality-ok: team lookup is outside the policy path + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=None, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await generate_service_account_key_fn( + data=GenerateKeyRequest(team_id="team-1", duration="3000d"), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == _POLICY_DENIAL_MESSAGE + generate_key_helper.assert_not_awaited() + mock_prisma.db.litellm_budgettable.create.assert_not_awaited() + assert [policy_request.operation for policy_request in received] == ["generate"] + assert received[0].existing_key is None + assert received[0].effective_key.team_id == "team-1" + assert received[0].effective_key.user_id is None + _assert_expires_in(received[0].effective_key, "3000d") + + +@pytest.mark.asyncio +async def test_effective_key_after_update_decodes_json_string_columns_and_keeps_omitted_fields(): + existing_key = LiteLLM_VerificationToken(token="tok", user_id="u1", team_id="team-a") + non_default_values = await prepare_key_update_data( + data=UpdateKeyRequest( + key="tok", router_settings={"num_retries": 3}, budget_limits=[{"budget_duration": "1d", "max_budget": 2.0}] + ), + existing_key_row=existing_key, + ) + assert isinstance(non_default_values["router_settings"], str) + assert isinstance(non_default_values["budget_limits"], str) + + effective_key = _effective_key_after_update(existing_key_row=existing_key, non_default_values=non_default_values) + + assert effective_key.router_settings == {"num_retries": 3} + assert effective_key.budget_limits is not None + assert effective_key.budget_limits[0]["max_budget"] == 2.0 + assert effective_key.budget_limits[0]["budget_duration"] == "1d" + assert effective_key.budget_limits[0]["reset_at"] is not None + assert effective_key.team_id == "team-a" + assert effective_key.user_id == "u1" + + +@pytest.mark.asyncio +async def test_effective_key_after_update_clears_expiry_for_a_minus_one_duration(): + existing_key = LiteLLM_VerificationToken(token="tok", expires=datetime(2027, 1, 1, tzinfo=timezone.utc)) + non_default_values = await prepare_key_update_data( + data=UpdateKeyRequest(key="tok", duration="-1"), existing_key_row=existing_key + ) + + effective_key = _effective_key_after_update(existing_key_row=existing_key, non_default_values=non_default_values) + + assert effective_key.expires is None + + +def test_effective_key_after_update_swaps_the_object_permission_id_and_drops_the_stale_relation(): + existing_key = LiteLLM_VerificationToken( + token="tok", + object_permission_id="op-old", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-old", mcp_servers=["old"]), + ) + + effective_key = _effective_key_after_update( + existing_key_row=existing_key, non_default_values={"object_permission_id": "op-new"} + ) + + assert effective_key.object_permission_id == "op-new" + assert effective_key.object_permission is None + assert existing_key.object_permission is not None + assert existing_key.object_permission.mcp_servers == ["old"] + + +def test_effective_key_for_generate_reflects_the_processed_request_without_mutating_it(): + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + data = GenerateKeyRequest( + duration="5d", + organization_id="org-1", + metadata={"a": 1}, + guardrails=["g1"], + tags=["t1"], + budget_duration="1d", + max_budget=3.0, + budget_limits=[{"budget_duration": "1d", "max_budget": 5.0}], + auto_rotate=True, + rotation_interval="30d", + object_permission={"mcp_servers": ["srv"]}, + key_type=LiteLLMKeyType.LLM_API, + ) + + effective_key = _effective_key_for_generate(data=data, now=now) + + assert effective_key.expires == now + timedelta(days=5) + assert effective_key.key_rotation_at == now + timedelta(days=30) + assert effective_key.budget_limits is not None + assert effective_key.budget_limits[0]["max_budget"] == 5.0 + assert effective_key.budget_limits[0]["reset_at"] is not None + assert effective_key.object_permission is None + assert effective_key.org_id == "org-1" + assert effective_key.metadata == {"a": 1, "guardrails": ["g1"], "tags": ["t1"]} + assert effective_key.max_budget == 3.0 + assert effective_key.budget_duration == "1d" + assert effective_key.budget_reset_at is not None + assert effective_key.key_type == "llm_api" + assert effective_key.allowed_routes == ["llm_api_routes"] + assert data.metadata == {"a": 1} + assert data.guardrails == ["g1"] + assert data.tags == ["t1"] + assert data.duration == "5d" + assert data.budget_limits is not None + assert data.budget_limits[0].reset_at is None + assert data.object_permission is not None + assert data.object_permission.mcp_servers == ["srv"] + + +def test_effective_key_for_generate_stores_no_budget_windows_for_an_empty_list(): + effective_key = _effective_key_for_generate( + data=GenerateKeyRequest(budget_limits=[]), now=datetime(2026, 1, 1, tzinfo=timezone.utc) + ) + + assert effective_key.budget_limits is None + + +def test_effective_key_for_generate_without_duration_never_expires(): + effective_key = _effective_key_for_generate( + data=GenerateKeyRequest(), now=datetime(2026, 1, 1, tzinfo=timezone.utc) + ) + + assert effective_key.expires is None + assert effective_key.budget_reset_at is None + assert effective_key.key_rotation_at is None + assert effective_key.key_type == "default" + + +def _policy_request_for_generate() -> CustomKeyPolicyRequest: + return CustomKeyPolicyRequest( + operation="generate", + existing_key=None, + effective_key=LiteLLM_VerificationToken(token="tok"), + request=GenerateKeyRequest(), + ) + + +@pytest.mark.asyncio +async def test_enforce_custom_key_policy_rejects_a_sync_hook(): + def sync_hook(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + return {"decision": True} + + with pytest.raises(ValueError, match="user_custom_key_policy must be a coroutine"): + await _enforce_custom_key_policy(hook=sync_hook, build_policy_request=_policy_request_for_generate) + + +@pytest.mark.asyncio +async def test_enforce_custom_key_policy_uses_the_default_denial_message(): + async def deny(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + return {"decision": False} + + with pytest.raises(HTTPException) as exc_info: + await _enforce_custom_key_policy(hook=deny, build_policy_request=_policy_request_for_generate) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == "Authentication Failed - Custom Auth Rule" + + +@pytest.mark.asyncio +async def test_enforce_custom_key_policy_allows_when_the_decision_is_missing(): + received: list[CustomKeyPolicyRequest] = [] + + async def no_decision(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + received.append(policy_request) + return {} + + await _enforce_custom_key_policy(hook=no_decision, build_policy_request=_policy_request_for_generate) + + assert len(received) == 1 + assert received[0].operation == "generate" + + +@pytest.mark.asyncio +async def test_enforce_custom_key_policy_never_builds_the_request_without_a_hook(): + await _enforce_custom_key_policy( + hook=None, build_policy_request=lambda: pytest.fail("policy request built without a hook") + ) + + @pytest.mark.asyncio async def test_regenerate_evicts_jwt_key_mapping_cache_so_next_jwt_call_gets_new_token(): """ @@ -13775,10 +14839,6 @@ async def test_regenerate_applies_normalized_mcp_object_permission(): "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_vector_stores_against_team", new_callable=AsyncMock, ), - patch( - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", - new_callable=AsyncMock, - ), patch( "litellm.proxy.management_endpoints.key_management_endpoints._execute_virtual_key_regeneration", execute_mock, @@ -18290,3 +19350,38 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): ) assert exc.value.status_code == 403 assert "Only proxy admins, team admins, or org admins" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch): + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateTeamKeysRequest, + KeyUpdateFields, + ) + + keys = [_make_team_key("tok-a"), _make_team_key("tok-b")] + mock = _setup_team_keys_mocks( + monkeypatch, find_many=keys, update_data=AsyncMock(return_value={"data": _updated({"max_budget": 50.0})}) + ) + received: list[CustomKeyPolicyRequest] = [] + + async def freeze_tok_b(policy_request: CustomKeyPolicyRequest) -> dict[str, object]: + received.append(policy_request) + if policy_request.existing_key is not None and policy_request.existing_key.token == "tok-b": + return {"decision": False, "message": "tok-b is frozen"} + return {"decision": True} + + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_policy", freeze_tok_b) + + response = await _call_as_admin( + BulkUpdateTeamKeysRequest( + team_id="team-abc", key_ids=["tok-a", "tok-b"], update_fields=KeyUpdateFields(max_budget=50.0) + ) + ) + + assert [update.key for update in response.successful_updates] == ["tok-a"] + assert [(failed.key, failed.failed_reason) for failed in response.failed_updates] == [("tok-b", "tok-b is frozen")] + mock.update_data.assert_awaited_once() + assert [policy_request.operation for policy_request in received] == ["update", "update"] + assert [policy_request.effective_key.max_budget for policy_request in received] == [50.0, 50.0] + assert [policy_request.effective_key.team_id for policy_request in received] == ["team-abc", "team-abc"] diff --git a/tests/test_litellm/proxy/management_endpoints/test_prompt_cache_prediction.py b/tests/test_litellm/proxy/management_endpoints/test_prompt_cache_prediction.py new file mode 100644 index 00000000000..0ec277be884 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_prompt_cache_prediction.py @@ -0,0 +1,698 @@ +import asyncio +import time +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from typing import Final, Literal + +import httpx +import pytest +from fastapi import FastAPI, Request +from pydantic import JsonValue + +import litellm +from litellm.caching.caching import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, _safe_set_request_parsed_body +from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, cache_scope, parse_prompt +from litellm.proxy.hooks.prompt_cache_prediction import ( + CacheObservation, + _cache_key, +) +from litellm.proxy.management_endpoints import prompt_cache_prediction as endpoint +from litellm.proxy.utils import InternalUsageCache +from litellm.types.management_endpoints.prompt_cache_prediction import CachePredictionResponse +from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + +_PROVIDER_KEY: Final = "cache-prediction-test-provider-key" +_CALLER: Final = "cache-prediction-test-caller-hash" + + +@pytest.fixture(autouse=True) +def anthropic_endpoint_environment(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + + +def _body(ttl: str = "5m", *, extended: bool = False) -> dict[str, JsonValue]: + blocks: Final[list[JsonValue]] = [ + {"type": "text", "text": "Stable context"}, + *([{"type": "text", "text": "Appended context"}] if extended else []), + ] + return { + "max_tokens": 10, + "system": "Follow the project conventions", + "messages": [ + { + "role": "user", + "content": [ + *blocks[:-1], + {**blocks[-1], "cache_control": {"type": "ephemeral", "ttl": ttl}}, + {"type": "text", "text": "Follow-up question"}, + ], + } + ], + } + + +def _prefix(body: Mapping[str, JsonValue]) -> PromptPrefix: + prefix: Final = parse_prompt(body) + assert prefix is not None + return prefix + + +def _deployment( + deployment_id: str = "sonnet", + model: str = "claude-sonnet-5", + *, + team_id: str | None = None, + api_base: str | None = None, +) -> Deployment: + return Deployment( + model_name=deployment_id, + litellm_params=LiteLLM_Params(model=f"anthropic/{model}", api_key=_PROVIDER_KEY, api_base=api_base), + model_info=ModelInfo(id=deployment_id, team_id=team_id), + ) + + +@dataclass(frozen=True) +class Counts: + total: int | None = 6_000 + prefix: int | None = 5_000 + + async def __call__(self, model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + assert api_key == _PROVIDER_KEY + assert model.startswith("claude-") + return self.total if "max_tokens" in body else self.prefix + + +async def _observe( + cache: DualCache, + body: Mapping[str, JsonValue], + *, + deployment_id: str = "sonnet", + model: str = "claude-sonnet-5", + cached_tokens: int = 5_000, + expired: bool = False, + caller: str = _CALLER, +) -> None: + prefix: Final = _prefix(body) + now: Final = time.time() + observation: Final = CacheObservation( + fingerprint=prefix.fingerprint, + cached_tokens=cached_tokens, + observed_at=now - 400 if expired else now - 10, + expires_at=now - 100 if expired else now + 290, + ) + scope: Final = cache_scope(caller, deployment_id, _PROVIDER_KEY, model) + await cache.async_set_cache(_cache_key(scope, prefix.fingerprint), observation.model_dump_json(), ttl=3_600) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("ttl", "cold_cost"), [("5m", 0.0145), ("1h", 0.022)]) +async def test_unobserved_cache_prices_cold_and_warm_bounds(ttl: str, cold_cost: float) -> None: + body: Final = _body(ttl) + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, DualCache(), Counts()) + + assert arm.cache_state == "unknown" + assert arm.reason == "no_compatible_observation" + assert arm.evidence is None + assert arm.estimate is not None and arm.cold is not None and arm.warm is not None + assert arm.estimate.input_cost == pytest.approx(cold_cost) + assert arm.cold.input_cost == pytest.approx(cold_cost) + assert arm.warm.input_cost == pytest.approx(0.003) + assert arm.cold.tokens.uncached_input_tokens == 1_000 + assert arm.cold.tokens.cache_read_input_tokens == 0 + assert arm.cold.tokens.cache_creation_5m_input_tokens == (5_000 if ttl == "5m" else 0) + assert arm.cold.tokens.cache_creation_1h_input_tokens == (5_000 if ttl == "1h" else 0) + assert arm.warm.tokens.cache_read_input_tokens == 5_000 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("cached_tokens", "warm_cost", "cold_cost"), [(5_400, 0.00228, 0.0147), (4_600, 0.00372, 0.0143)] +) +@pytest.mark.parametrize("expired", [False, True]) +async def test_exact_prefix_conserves_total_with_observed_count_in_all_scenarios( + cached_tokens: int, warm_cost: float, cold_cost: float, expired: bool +) -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body, cached_tokens=cached_tokens, expired=expired) + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, cache, Counts()) + + assert arm.cache_state == ("stale" if expired else "warm") + assert arm.evidence is not None + assert arm.estimate is not None and arm.warm is not None and arm.cold is not None + assert arm.warm.tokens.cache_read_input_tokens == cached_tokens + assert arm.warm.tokens.cache_creation_5m_input_tokens == 0 + assert arm.cold.tokens.cache_creation_5m_input_tokens == cached_tokens + assert arm.cold.tokens.cache_read_input_tokens == 0 + for scenario in (arm.estimate, arm.cold, arm.warm): + assert scenario.tokens.total_tokens == 6_000 + assert scenario.tokens.uncached_input_tokens == 6_000 - cached_tokens + assert arm.warm.input_cost == pytest.approx(warm_cost) + assert arm.cold.input_cost == pytest.approx(cold_cost) + assert arm.estimate.input_cost == pytest.approx(cold_cost if expired else warm_cost) + + +@pytest.mark.asyncio +async def test_observed_prefix_larger_than_full_request_returns_unknown() -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body, cached_tokens=6_001) + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, cache, Counts()) + + assert arm.cache_state == "unknown" + assert arm.reason == "inconsistent_prefix_token_count" + assert arm.estimate is None and arm.cold is None and arm.warm is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("ttl", "expected"), [("5m", 0.0053), ("1h", 0.0068)]) +async def test_append_only_prefix_reads_old_tokens_and_writes_extension(ttl: str, expected: float) -> None: + cache: Final = DualCache() + await _observe(cache, _body(ttl), cached_tokens=4_000) + body: Final = _body(ttl, extended=True) + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, cache, Counts()) + + assert arm.cache_state == "partial" + assert arm.estimate is not None + assert arm.estimate.tokens.cache_read_input_tokens == 4_000 + assert arm.estimate.tokens.cache_creation_5m_input_tokens == (1_000 if ttl == "5m" else 0) + assert arm.estimate.tokens.cache_creation_1h_input_tokens == (1_000 if ttl == "1h" else 0) + assert arm.estimate.input_cost == pytest.approx(expected) + + +@pytest.mark.asyncio +async def test_expired_observation_estimates_a_cold_rebuild() -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body, expired=True) + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, cache, Counts()) + + assert arm.cache_state == "stale" + assert arm.reason == "observation_expired" + assert arm.evidence is not None and arm.evidence.expires_at < time.time() + assert arm.estimate is not None and arm.cold is not None + assert arm.estimate.tokens.cache_read_input_tokens == 0 + assert arm.estimate.tokens.cache_creation_5m_input_tokens == 5_000 + assert arm.estimate.input_cost == arm.cold.input_cost + + +@pytest.mark.asyncio +async def test_below_model_minimum_prices_all_input_as_uncached() -> None: + body: Final = _body() + arm: Final = await endpoint.predict_arm( + _deployment(), body, _prefix(body), _CALLER, DualCache(), Counts(total=1_500, prefix=1_000) + ) + + assert arm.cache_state == "disabled" + assert arm.reason == "below_cache_minimum" + assert arm.estimate is not None + assert arm.estimate.tokens.uncached_input_tokens == 1_500 + assert arm.estimate.tokens.cache_read_input_tokens == 0 + assert arm.estimate.tokens.cache_creation_5m_input_tokens == 0 + assert arm.estimate.input_cost == pytest.approx(0.003) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("counts", [Counts(total=None), Counts(prefix=None), Counts(total=4_000)]) +async def test_unavailable_or_inconsistent_token_counts_return_null_estimates(counts: Counts) -> None: + body: Final = _body() + arm: Final = await endpoint.predict_arm(_deployment(), body, _prefix(body), _CALLER, DualCache(), counts) + + assert arm.cache_state == "unknown" + assert arm.reason == "token_count_unavailable" + assert arm.estimate is None and arm.cold is None and arm.warm is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("counts", [Counts(), Counts(total=1_500, prefix=1_000)]) +async def test_missing_prices_return_unknown_and_null_estimates( + monkeypatch: pytest.MonkeyPatch, counts: Counts +) -> None: + monkeypatch.setitem( + litellm.model_cost, + "claude-cache-unpriced-5", + {"litellm_provider": "anthropic", "mode": "chat"}, + ) + body: Final = _body() + arm: Final = await endpoint.predict_arm( + _deployment("cache-prediction-unpriced", "claude-cache-unpriced-5"), + body, + _prefix(body), + _CALLER, + DualCache(), + counts, + ) + + assert arm.cache_state == "unknown" + assert arm.reason == "pricing_unavailable" + assert arm.estimate is None and arm.cold is None and arm.warm is None + + +@pytest.mark.asyncio +async def test_custom_api_base_from_environment_returns_unknown_before_counting( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://custom.invalid") + body: Final = _body() + arm: Final = await endpoint.predict_arm( + _deployment(), body, _prefix(body), _CALLER, DualCache(), _unexpected_count + ) + + assert arm.cache_state == "unknown" + assert arm.reason == "unsupported_provider_endpoint" + assert arm.estimate is None and arm.cold is None and arm.warm is None + + +@pytest.mark.asyncio +async def test_explicit_official_api_base_overrides_custom_environment(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://custom.invalid") + body: Final = _body() + arm: Final = await endpoint.predict_arm( + _deployment(api_base="https://api.anthropic.com"), body, _prefix(body), _CALLER, DualCache(), Counts() + ) + + assert arm.cache_state == "unknown" + assert arm.reason == "no_compatible_observation" + assert arm.estimate is not None + assert arm.estimate.input_cost == pytest.approx(0.0145) + + +@dataclass(frozen=True) +class _ProxyLogging: + internal_usage_cache: InternalUsageCache + parallel_limiter: CustomLogger | None + + def get_proxy_hook(self, hook: str) -> CustomLogger | None: + return self.parallel_limiter if hook == "parallel_request_limiter" else None + + +def _app( + monkeypatch: pytest.MonkeyPatch, + cache: DualCache, + *, + caller: UserAPIKeyAuth | None = None, + current_team: str | None = None, + candidate_team: str | None = None, + counts: endpoint.TokenCounter = Counts(), + limiter: CustomLogger | Literal["default"] | None = "default", +) -> FastAPI: + import litellm.proxy.proxy_server as proxy_server + + model_list: Final = [ + _deployment("opus", "claude-opus-5", team_id=current_team).model_dump(exclude_unset=True), + _deployment("sonnet", team_id=candidate_team).model_dump(exclude_unset=True), + ] + router: Final = litellm.Router(model_list=model_list) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", model_list) + monkeypatch.setattr(endpoint, "count_prompt_tokens", counts) + app: Final = FastAPI() + app.include_router(endpoint.router) + app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler) + if caller is not None: + usage_cache: Final = InternalUsageCache(cache) + configured_limiter: Final = ( + _PROXY_MaxParallelRequestsHandler_v3(usage_cache) if isinstance(limiter, str) else limiter + ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", _ProxyLogging(usage_cache, configured_limiter)) + app.dependency_overrides[endpoint.user_api_key_auth] = lambda: caller + return app + + +async def _post( + app: FastAPI, + body: Mapping[str, JsonValue], + *, + current_deployment_id: str = "opus", + candidate_deployment_id: str = "sonnet", +) -> httpx.Response: + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + return await client.post( + "/cost/predict-cache", + json={ + "current_deployment_id": current_deployment_id, + "candidate_deployment_id": candidate_deployment_id, + "request": body, + }, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("warm_deployment", "warm_model", "expected_delta", "expected_penalty"), + [("sonnet", "claude-sonnet-5", -0.03325, 0.0), ("opus", "claude-opus-5", 0.007, 0.0115)], +) +async def test_switch_delta_accounts_for_each_deployment_cache( + monkeypatch: pytest.MonkeyPatch, + warm_deployment: str, + warm_model: str, + expected_delta: float, + expected_penalty: float, +) -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body, deployment_id=warm_deployment, model=warm_model) + app: Final = _app(monkeypatch, cache, caller=UserAPIKeyAuth(api_key=_CALLER)) + response: Final = await _post(app, body) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.switch_delta == pytest.approx(expected_delta) + assert result.cache_rebuild_penalty == pytest.approx(expected_penalty) + assert result.cache_guarantee is False + assert result.pricing_basis == "input_before_discounts_and_margins" + if warm_deployment == "sonnet": + assert result.switch.cache_state == "warm" + assert result.stay.cache_state == "unknown" + else: + assert result.stay.cache_state == "warm" + assert result.switch.cache_state == "unknown" + + +@pytest.mark.asyncio +async def test_missing_caller_identity_cannot_reuse_observations(monkeypatch: pytest.MonkeyPatch) -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body) + response: Final = await _post( + _app(monkeypatch, cache, caller=UserAPIKeyAuth(api_key=None), counts=_unexpected_count), body + ) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.reason == result.switch.reason == "caller_identity_unavailable" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None + + +@pytest.mark.asyncio +async def test_unauthenticated_request_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm.proxy.proxy_server as proxy_server + + monkeypatch.setattr(proxy_server, "master_key", "cache-prediction-test-master-key") + response: Final = await _post(_app(monkeypatch, DualCache()), _body()) + assert response.status_code == 401, response.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("arm", ["current", "candidate"]) +@pytest.mark.parametrize("caller_team", [None, "own-team"]) +@pytest.mark.parametrize("restricted", [False, True]) +async def test_foreign_and_missing_deployments_have_identical_authenticated_responses( + monkeypatch: pytest.MonkeyPatch, arm: str, caller_team: str | None, restricted: bool +) -> None: + allowed: Final = ("sonnet",) if arm == "current" else ("opus",) + app: Final = _app( + monkeypatch, + DualCache(), + caller=UserAPIKeyAuth(api_key=_CALLER, team_id=caller_team, models=list(allowed) if restricted else []), + current_team="foreign-team" if arm == "current" else None, + candidate_team="foreign-team" if arm == "candidate" else None, + counts=_unexpected_count, + ) + foreign: Final = await _post(app, _body()) + missing: Final = await _post( + app, + _body(), + current_deployment_id="missing-deployment" if arm == "current" else "opus", + candidate_deployment_id="missing-deployment" if arm == "candidate" else "sonnet", + ) + + assert foreign.status_code == missing.status_code == 404 + assert foreign.json() == missing.json() == {"detail": "Deployment not found"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("deployment_team", [None, "own-team"]) +async def test_visible_public_and_own_team_deployments_remain_available( + monkeypatch: pytest.MonkeyPatch, deployment_team: str | None +) -> None: + app: Final = _app( + monkeypatch, + DualCache(), + caller=UserAPIKeyAuth(api_key=_CALLER, team_id="own-team"), + current_team=deployment_team, + candidate_team=deployment_team, + ) + response: Final = await _post(app, _body()) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.estimate is not None and result.switch.estimate is not None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("arm", ["current", "candidate"]) +async def test_visible_deployment_outside_key_model_permissions_is_forbidden( + monkeypatch: pytest.MonkeyPatch, arm: str +) -> None: + allowed: Final = "sonnet" if arm == "current" else "opus" + denied: Final = "opus" if arm == "current" else "sonnet" + app: Final = _app(monkeypatch, DualCache(), caller=UserAPIKeyAuth(api_key=_CALLER, models=[allowed])) + response: Final = await _post(app, _body()) + assert response.status_code == 403, response.text + assert denied in response.text + + +@pytest.mark.asyncio +async def test_other_callers_warm_cache_is_not_prediction_evidence(monkeypatch: pytest.MonkeyPatch) -> None: + cache: Final = DualCache() + body: Final = _body() + await _observe(cache, body, caller="other-caller") + response: Final = await _post(_app(monkeypatch, cache, caller=UserAPIKeyAuth(api_key=_CALLER)), body) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.switch.cache_state == "unknown" + assert result.switch.reason == "no_compatible_observation" + assert result.switch.evidence is None + assert result.switch.estimate is not None + assert result.switch.estimate.tokens.cache_read_input_tokens == 0 + + +@pytest.mark.asyncio +async def test_count_failure_nulls_switch_comparison(monkeypatch: pytest.MonkeyPatch) -> None: + app: Final = _app( + monkeypatch, DualCache(), caller=UserAPIKeyAuth(api_key=_CALLER), counts=Counts(total=None) + ) + response: Final = await _post(app, _body()) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.reason == result.switch.reason == "token_count_unavailable" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("limiter", [None, CustomLogger()]) +async def test_missing_or_unsupported_limiter_returns_unknown_before_counting( + monkeypatch: pytest.MonkeyPatch, limiter: CustomLogger | None +) -> None: + app: Final = _app( + monkeypatch, DualCache(), caller=UserAPIKeyAuth(api_key=_CALLER), counts=_unexpected_count, limiter=limiter + ) + response: Final = await _post(app, _body()) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.reason == result.switch.reason == "limiter_unavailable" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None + + +@pytest.mark.asyncio +async def test_occupied_parallel_capacity_rejects_before_provider_count(monkeypatch: pytest.MonkeyPatch) -> None: + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + caller: Final = UserAPIKeyAuth(api_key=_CALLER, max_parallel_requests=1) + app: Final = _app(monkeypatch, cache, caller=caller, counts=_unexpected_count, limiter=limiter) + async with limiter.request_capacity(caller, "opus"): + response: Final = await _post(app, _body()) + + assert response.status_code == 429, response.text + assert "max_parallel_requests" in response.text + recovered: Final = await _post(_app(monkeypatch, cache, caller=caller, limiter=limiter), _body()) + assert recovered.status_code == 200, recovered.text + + +@pytest.mark.asyncio +async def test_each_count_consumes_the_deployment_group_rpm_limit(monkeypatch: pytest.MonkeyPatch) -> None: + calls: Final = asyncio.Queue[str]() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + calls.put_nowait(model) + return await Counts()(model, api_key, body) + + caller: Final = UserAPIKeyAuth(api_key=_CALLER, metadata={"model_rpm_limit": {"sonnet": 1}}) + app: Final = _app(monkeypatch, DualCache(), caller=caller, counts=count) + response: Final = await _post(app, _body()) + + assert response.status_code == 429, response.text + assert calls.qsize() == 3 + assert tuple(calls.get_nowait() for _ in range(3)) == ( + "claude-opus-5", "claude-opus-5", "claude-sonnet-5" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) +async def test_each_count_preserves_auth_cached_request_tag_limits( + monkeypatch: pytest.MonkeyPatch, metadata_key: str +) -> None: + calls: Final = asyncio.Queue[str]() + caller: Final = UserAPIKeyAuth(api_key=_CALLER, metadata={"tag_rpm_limit": {"cache-cost": 1}}) + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + calls.put_nowait(model) + return await Counts()(model, api_key, body) + + async def authenticated_request(request: Request) -> UserAPIKeyAuth: + data: Final = await _read_request_body(request) + _safe_set_request_parsed_body(request, {**data, metadata_key: {"tags": ["cache-cost"]}}) + return caller + + app: Final = _app(monkeypatch, DualCache(), caller=caller, counts=count) + app.dependency_overrides[endpoint.user_api_key_auth] = authenticated_request + response: Final = await _post(app, _body()) + + assert response.status_code == 429, response.text + assert "tag_per_key" in response.text + assert calls.qsize() == 1 + assert calls.get_nowait() == "claude-opus-5" + + +@pytest.mark.asyncio +async def test_provider_counter_failure_releases_parallel_capacity(monkeypatch: pytest.MonkeyPatch) -> None: + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + caller: Final = UserAPIKeyAuth(api_key=_CALLER, max_parallel_requests=1) + + async def fail_count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + raise RuntimeError("provider counter failed") + + app: Final = _app(monkeypatch, cache, caller=caller, counts=fail_count, limiter=limiter) + with pytest.raises(RuntimeError, match="provider counter failed"): + await _post(app, _body()) + recovered: Final = await _post(_app(monkeypatch, cache, caller=caller, limiter=limiter), _body()) + assert recovered.status_code == 200, recovered.text + assert recovered.json()["switch"]["estimate"]["input_cost"] == pytest.approx(0.0145) + + +@pytest.mark.asyncio +async def test_cancelled_provider_counter_releases_parallel_capacity(monkeypatch: pytest.MonkeyPatch) -> None: + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + caller: Final = UserAPIKeyAuth(api_key=_CALLER, max_parallel_requests=1) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def wait_count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + started.set() + await release.wait() + return await Counts()(model, api_key, body) + + app: Final = _app(monkeypatch, cache, caller=caller, counts=wait_count, limiter=limiter) + pending: Final = asyncio.create_task(_post(app, _body())) + try: + await asyncio.wait_for(started.wait(), timeout=5) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + release.set() + recovered: Final = await asyncio.wait_for(_post(app, _body()), timeout=5) + assert recovered.status_code == 200, recovered.text + assert recovered.json()["switch"]["estimate"]["input_cost"] == pytest.approx(0.0145) + finally: + pending.cancel() + release.set() + await asyncio.gather(pending, return_exceptions=True) + + +async def _unexpected_count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + pytest.fail("Unsupported prediction must return before contacting the token counter") + + +class RequestMutator(CustomLogger): + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str + ) -> dict[str, object]: + return {**data, "system": "Injected policy"} + + +@pytest.fixture +def request_mutator() -> Iterator[RequestMutator]: + callback: Final = RequestMutator() + litellm.logging_callback_manager.add_litellm_callback(callback) + try: + yield callback + finally: + litellm.logging_callback_manager.remove_callback_from_all_lists(callback) + + +@pytest.mark.asyncio +async def test_request_transform_callback_returns_unknown_before_token_counting( + monkeypatch: pytest.MonkeyPatch, request_mutator: RequestMutator +) -> None: + app: Final = _app( + monkeypatch, DualCache(), caller=UserAPIKeyAuth(api_key=_CALLER), counts=_unexpected_count + ) + response: Final = await _post(app, _body()) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.cache_state == result.switch.cache_state == "unknown" + assert result.stay.reason == result.switch.reason == "unsupported_request_transform" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None + + +@pytest.mark.asyncio +async def test_key_config_returns_unknown_before_token_counting(monkeypatch: pytest.MonkeyPatch) -> None: + app: Final = _app( + monkeypatch, + DualCache(), + caller=UserAPIKeyAuth(api_key=_CALLER, config={"model_list": []}), + counts=_unexpected_count, + ) + response: Final = await _post(app, _body()) + + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.cache_state == result.switch.cache_state == "unknown" + assert result.stay.reason == result.switch.reason == "unsupported_request_transform" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None + + +@pytest.mark.parametrize("headers", [ + {"anthropic-version": "2099-01-01"}, + {"anthropic-beta": "future-feature"}, +]) +@pytest.mark.asyncio +async def test_unsupported_provider_headers_cannot_reuse_default_version_evidence( + monkeypatch: pytest.MonkeyPatch, headers: dict[str, str] +) -> None: + cache: Final = DualCache() + await _observe(cache, _body(), deployment_id="sonnet") + app: Final = _app( + monkeypatch, cache, caller=UserAPIKeyAuth(api_key=_CALLER), counts=_unexpected_count + ) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response: Final = await client.post( + "/cost/predict-cache", + headers=headers, + json={"current_deployment_id": "opus", "candidate_deployment_id": "sonnet", "request": _body()}, + ) + assert response.status_code == 200, response.text + result: Final = CachePredictionResponse.model_validate(response.json()) + assert result.stay.cache_state == result.switch.cache_state == "unknown" + assert result.stay.reason == result.switch.reason == "unsupported_provider_headers" + assert result.stay.estimate is None and result.switch.estimate is None + assert result.switch_delta is None and result.cache_rebuild_penalty is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index ccb66fd9534..a2f534fbe4d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -9476,6 +9476,9 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[ + SimpleNamespace(model_id="weighted-id", model_name="group", model_info={}) + ]) # Mock model table creation mock_db_client.db.litellm_modeltable = MagicMock() @@ -9511,6 +9514,7 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): # Test router_settings with sample data router_settings_data = { + "weights": {"group": {"weighted-id": 1}}, "routing_strategy": "usage-based", "num_retries": 3, "retry_policy": {"max_retries": 5}, @@ -9544,6 +9548,12 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): deserialized_settings = json.loads(team_data["router_settings"]) assert deserialized_settings == router_settings_data + mock_team_create.reset_mock() + team_request.router_settings = {"weights": {"group": {"unknown-id": 1}}} + with pytest.raises(ProxyException, match="Unknown deployment ID"): + await new_team(data=team_request, http_request=dummy_request, user_api_key_dict=mock_admin_auth) + mock_team_create.assert_not_awaited() + @pytest.mark.asyncio async def test_get_team_daily_activity_member_with_permission_sees_all_spend( @@ -9739,6 +9749,9 @@ async def test_update_team_with_router_settings( # Configure mocked prisma client mock_db_client.jsonify_team_object = lambda db_data: db_data mock_db_client.db = MagicMock() + mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[ + SimpleNamespace(model_id="weighted-id", model_name="group", model_info={}) + ]) # Mock existing team row existing_team_mock = MagicMock() @@ -9773,6 +9786,7 @@ async def test_update_team_with_router_settings( # Test router_settings with updated data router_settings_data = { + "weights": {"group": {"weighted-id": 1}}, "routing_strategy": "latency-based", "num_retries": 2, } @@ -9805,6 +9819,12 @@ async def test_update_team_with_router_settings( deserialized_settings = json.loads(team_data["router_settings"]) assert deserialized_settings == router_settings_data + mock_team_update.reset_mock() + team_update_request.router_settings = {"weights": {"group": {"unknown-id": 1}}} + with pytest.raises(ProxyException, match="Unknown deployment ID"): + await update_team(data=team_update_request, http_request=dummy_request, user_api_key_dict=mock_admin_auth) + mock_team_update.assert_not_awaited() + @pytest.mark.asyncio async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( @@ -14366,3 +14386,123 @@ async def test_get_team_spend_by_user_rejects_bad_input(mock_db_client, team_ids assert exc_info.value.status_code == 400 assert expected_error in str(exc_info.value.detail) mock_db_client.db.query_raw.assert_not_called() + + +class _TeamRowWithOrganization(LiteLLM_TeamTable): + litellm_organization_table: LiteLLM_OrganizationTable | None = None + + +@pytest.mark.parametrize( + "organization, expected_models", + [ + ( + LiteLLM_OrganizationTable( + organization_id="org-1", + budget_id="budget-1", + models=["all-proxy-models"], + created_by="admin", + updated_by="admin", + ), + ["all-proxy-models"], + ), + ( + LiteLLM_OrganizationTable( + organization_id="org-1", + budget_id="budget-1", + models=["gpt-4o"], + created_by="admin", + updated_by="admin", + ), + ["gpt-4o"], + ), + (None, None), + ], +) +@pytest.mark.asyncio +async def test_team_info_returns_parent_organization_models(organization, expected_models): + """/team/info must report the parent org's model ceiling. + + A team admin who is not an org admin gets a 403 from /organization/info, so this + is the only read that can tell the Admin UI whether the org allows all proxy + models. Without it the team edit form hides the "All Proxy Models" option and a + team admin cannot grant their team everything on the proxy. + """ + from fastapi import Request + + from litellm.proxy.management_endpoints import team_endpoints + + team_row = _TeamRowWithOrganization( + team_id="team-1", + organization_id="org-1" if organization is not None else None, + litellm_organization_table=organization, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.get_data = AsyncMock(return_value=[]) + + memberships = AsyncMock(return_value=[]) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: no seam on team_info + patch.object(team_endpoints, "get_all_team_memberships", memberships), # test-quality-ok: no seam on team_info + ): + response = await team_endpoints.team_info( + http_request=MagicMock(spec=Request), + team_id="team-1", + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response["team_info"].organization_models == expected_models + + +@pytest.mark.parametrize( + "caller, expected_models", + [ + (UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.INTERNAL_USER), ["gpt-4o"]), + (UserAPIKeyAuth(user_id="member-1", user_role=LitellmUserRoles.INTERNAL_USER), None), + (UserAPIKeyAuth(team_id="team-1"), None), + ], +) +@pytest.mark.asyncio +async def test_team_info_reports_parent_organization_models_only_to_team_managers(caller, expected_models): + """Plain members and team keys can read their team, but not the org's wider allow-list.""" + from fastapi import Request + + from litellm.proxy.management_endpoints import team_endpoints + + team_row = _TeamRowWithOrganization( + team_id="team-1", + organization_id="org-1", + members_with_roles=[ + Member(user_id="admin-1", role="admin"), + Member(user_id="member-1", role="user"), + ], + litellm_organization_table=LiteLLM_OrganizationTable( + organization_id="org-1", + budget_id="budget-1", + models=["gpt-4o"], + created_by="admin", + updated_by="admin", + ), + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + mock_prisma.get_data = AsyncMock(return_value=[]) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: no seam on team_info + patch.object(team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[])), # test-quality-ok: no seam on team_info + patch.object( # test-quality-ok: no seam on team_info + team_endpoints, "_is_user_org_admin_for_team", AsyncMock(return_value=False) + ), + ): + response = await team_endpoints.team_info( + http_request=MagicMock(spec=Request), + team_id="team-1", + user_api_key_dict=caller, + ) + + assert response["team_info"].organization_models == expected_models diff --git a/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py new file mode 100644 index 00000000000..b5349fc2387 --- /dev/null +++ b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py @@ -0,0 +1,431 @@ +import json +from contextlib import asynccontextmanager +from typing import Final + +import httpx +import pytest +from prisma.errors import UniqueViolationError +from pydantic import BaseModel, ConfigDict, ValidationError + +from litellm.caching.caching import DualCache +from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, UserAPIKeyAuth +from litellm.proxy.list_api.common import ManagementProblem +from litellm.proxy.management_helpers.bulk_user_creation import bulk_create_users +from litellm.types.proxy.management_endpoints.internal_user_endpoints import ( + BulkNewUserItem, + BulkNewUserRequest, +) + +ADMIN: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) +INTERNAL: Final = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER) + + +class _UserRow(BaseModel): + model_config = ConfigDict(extra="allow") + + user_id: str + user_email: str | None = None + user_role: str | None = None + teams: list[str] = [] + max_budget: float | None = None + + +class _UserTable: + """Enough of the Prisma user table for the bulk path: set lookups, one create_many and per-row fallbacks.""" + + def __init__( + self, + fail_ids: frozenset[str] = frozenset(), + commit_then_drop: bool = False, + raced_ids: frozenset[str] = frozenset(), + ) -> None: + self.rows: dict[str, _UserRow] = {} + self.fail_ids = fail_ids + self.commit_then_drop = commit_then_drop + self.raced_ids = raced_ids + self.create_many_calls = 0 + + async def count(self, where: object = None) -> int: + return 0 if where is not None else len(self.rows) + + async def find_many(self, where: dict[str, dict[str, object]]) -> list[_UserRow]: + if "user_id" in where: + wanted = where["user_id"]["in"] + return [row for row in self.rows.values() if row.user_id in wanted] + wanted_emails = {str(e).lower() for e in where["user_email"]["in"]} + return [row for row in self.rows.values() if (row.user_email or "").lower() in wanted_emails] + + async def create(self, data: dict[str, object]) -> _UserRow: + row = _UserRow.model_validate(data) + if row.user_id in self.fail_ids or row.user_id in self.rows: + raise RuntimeError(f"insert failed for {row.user_id}") + self.rows[row.user_id] = row + return row + + async def create_many(self, data: list[dict[str, object]]) -> int: + self.create_many_calls += 1 + rows = [_UserRow.model_validate(d) for d in data] + if any(row.user_id in self.fail_ids for row in rows): + raise RuntimeError("batch insert failed") + raced = [row.user_id for row in rows if row.user_id in self.raced_ids] + if raced: + for user_id in raced: + self.rows[user_id] = _UserRow(user_id=user_id, user_email=f"{user_id}@other-request.example") + raise UniqueViolationError({}, message="Unique constraint failed on the fields: (`user_id`)") + for row in rows: + self.rows[row.user_id] = row + if self.commit_then_drop: + raise httpx.ReadError("connection reset after commit") + return len(rows) + + async def update(self, where: dict[str, str], data: dict[str, object]) -> _UserRow: + row = self.rows[where["user_id"]] + updated = _UserRow.model_validate({**row.model_dump(), **data}) + self.rows[row.user_id] = updated + return updated + + +class _TeamTable: + def __init__(self, teams: list[LiteLLM_TeamTable]) -> None: + self.rows = {team.team_id: team for team in teams} + self.update_calls = 0 + + async def find_many(self, where: dict[str, dict[str, list[str]]]) -> list[LiteLLM_TeamTable]: + return [self.rows[team_id] for team_id in where["team_id"]["in"] if team_id in self.rows] + + async def update(self, where: dict[str, str], data: dict[str, str]) -> LiteLLM_TeamTable: + self.update_calls += 1 + team = self.rows[where["team_id"]] + team.members_with_roles = [Member(**m) for m in json.loads(data["members_with_roles"])] + return team + + +class _MembershipTable: + def __init__(self) -> None: + self.rows: list[dict[str, object]] = [] + + async def create_many(self, data: list[dict[str, object]], skip_duplicates: bool = False) -> int: + self.rows.extend(data) + return len(data) + + +class _Tx: + def __init__(self, db: "_Db") -> None: + self.litellm_teamtable = db.litellm_teamtable + self.litellm_teammembership = db.litellm_teammembership + self.locks: list[str] = [] + + async def query_raw(self, sql: str, *args: object) -> list[dict[str, object]]: + if "pg_advisory_xact_lock" in sql: + self.locks.append(str(args[0])) + return [] + team = self.litellm_teamtable.rows.get(str(args[0])) + if team is None: + return [] + return [{"members_with_roles": [m.model_dump() for m in team.members_with_roles]}] + + +class _Db: + def __init__( + self, + teams: list[LiteLLM_TeamTable], + fail_ids: frozenset[str] = frozenset(), + commit_then_drop: bool = False, + raced_ids: frozenset[str] = frozenset(), + ) -> None: + self.litellm_usertable = _UserTable(fail_ids, commit_then_drop, raced_ids) + self.litellm_teamtable = _TeamTable(teams) + self.litellm_teammembership = _MembershipTable() + + +class _FakePrisma: + def __init__( + self, + teams: list[LiteLLM_TeamTable] | None = None, + fail_ids: frozenset[str] = frozenset(), + commit_then_drop: bool = False, + raced_ids: frozenset[str] = frozenset(), + ) -> None: + self.db = _Db(teams or [], fail_ids, commit_then_drop, raced_ids) + self.tx_count = 0 + self.locks: list[str] = [] + + def jsonify_object(self, data: dict[str, object]) -> dict[str, object]: + return data + + @asynccontextmanager + async def tx(self): + self.tx_count += 1 + tx = _Tx(self.db) + yield tx + self.locks.extend(tx.locks) + + +class _License: + def __init__(self, max_users: int | None = None) -> None: + self.max_users = max_users + self.seen: list[int] = [] + + def is_over_limit(self, total_users: int) -> bool: + self.seen.append(total_users) + return self.max_users is not None and total_users > self.max_users + + +def _team(team_id: str, members: list[Member] | None = None) -> LiteLLM_TeamTable: + return LiteLLM_TeamTable(team_id=team_id, members_with_roles=members or []) + + +async def _no_keys(**kwargs: object) -> dict[str, object]: + raise AssertionError(f"key generation was not requested: {kwargs}") + + +async def _run(prisma, users, caller=ADMIN, license=None, generate_key=_no_keys): + return await bulk_create_users( + users=[BulkNewUserItem(**u) for u in users], + user_api_key_dict=caller, + prisma_client=prisma, + license_check=license or _License(), + litellm_proxy_admin_name="default_user_id", + user_api_key_cache=DualCache(), + generate_key=generate_key, + ) + + +@pytest.mark.asyncio +async def test_creates_users_and_team_membership_in_every_store(): + prisma = _FakePrisma(teams=[_team("t1", [Member(user_id="existing", role="admin")]), _team("t2")]) + response = await _run( + prisma, + [ + {"user_id": "u1", "user_email": "a@example.com", "teams": ["t1", "t2"], "max_budget": 50}, + {"user_id": "u2", "user_email": "b@example.com", "teams": ["t1"]}, + {"user_id": "u3", "user_email": "c@example.com"}, + ], + ) + + assert (response.meta.total_requested, response.meta.created, response.meta.failed) == (3, 3, 0) + assert [r.user_id for r in response.data] == ["u1", "u2", "u3"] + assert all(r.success and r.key is None and r.error is None for r in response.data) + assert [r.teams for r in response.data] == [("t1", "t2"), ("t1",), ()] + + users = prisma.db.litellm_usertable.rows + assert users["u1"].teams == ["t1", "t2"] and users["u1"].max_budget == 50 + assert users["u2"].teams == ["t1"] and users["u3"].teams == [] + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["existing", "u1", "u2"] + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t2"].members_with_roles] == ["u1"] + assert sorted((m["team_id"], m["user_id"]) for m in prisma.db.litellm_teammembership.rows) == [ + ("t1", "u1"), + ("t1", "u2"), + ("t2", "u1"), + ] + + +@pytest.mark.asyncio +async def test_user_id_already_on_the_roster_keeps_the_team_and_is_not_added_twice(): + prisma = _FakePrisma(teams=[_team("t1", [Member(user_id="u1", role="user")])]) + response = await _run(prisma, [{"user_id": "u1", "teams": ["t1"]}, {"user_id": "u2", "teams": ["t1"]}]) + + assert [r.success for r in response.data] == [True, True] + assert [r.teams for r in response.data] == [("t1",), ("t1",)] + assert [r.error for r in response.data] == [None, None] + assert prisma.db.litellm_usertable.rows["u1"].teams == ["t1"] + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1", "u2"] + + +@pytest.mark.asyncio +async def test_one_insert_and_one_locked_write_per_team(): + prisma = _FakePrisma(teams=[_team("t1"), _team("t2")]) + await _run( + prisma, + [{"user_id": f"u{i}", "teams": ["t1"] if i % 2 else ["t1", "t2"]} for i in range(20)], + ) + + assert prisma.db.litellm_usertable.create_many_calls == 1 + assert prisma.tx_count == 2 + assert sorted(prisma.locks) == ["t1", "t2"] + assert prisma.db.litellm_teamtable.update_calls == 2 + assert len(prisma.db.litellm_teamtable.rows["t1"].members_with_roles) == 20 + assert len(prisma.db.litellm_teamtable.rows["t2"].members_with_roles) == 10 + + +@pytest.mark.asyncio +async def test_bad_rows_fail_alone_and_good_rows_still_land(): + prisma = _FakePrisma(teams=[_team("t1")]) + prisma.db.litellm_usertable.rows["taken"] = _UserRow(user_id="taken", user_email="Taken@Example.com") + response = await _run( + prisma, + [ + {"user_id": "u1", "user_email": "a@example.com", "teams": ["t1"]}, + {"user_id": "u2", "user_email": "A@EXAMPLE.COM"}, + {"user_id": "u1", "user_email": "z@example.com"}, + {"user_id": "u3", "user_email": "taken@example.com"}, + {"user_id": "taken"}, + {"user_id": "u4", "teams": ["missing"]}, + {"user_id": "u5", "teams": ["t1", "missing"]}, + {"user_id": "u6", "budget_duration": "not-a-duration"}, + {"user_id": "u7", "user_email": "ok@example.com", "teams": ["t1"]}, + ], + ) + + assert [r.success for r in response.data] == [True, False, False, False, False, False, False, False, True] + assert (response.meta.created, response.meta.failed) == (2, 7) + errors = [r.error for r in response.data] + assert "Duplicate user_email" in errors[1] + assert "Duplicate user_id" in errors[2] + assert "already exists" in errors[3] and "already exists" in errors[4] + assert "missing" in errors[5] and "does not exist" in errors[5] + assert "missing" in errors[6] + assert errors[7] is not None + + assert set(prisma.db.litellm_usertable.rows) == {"taken", "u1", "u7"} + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1", "u7"] + + +@pytest.mark.asyncio +async def test_insert_failure_falls_back_to_per_row_and_reports_only_that_row(): + prisma = _FakePrisma(teams=[_team("t1")], fail_ids=frozenset({"u2"})) + response = await _run( + prisma, + [{"user_id": "u1", "teams": ["t1"]}, {"user_id": "u2", "teams": ["t1"]}, {"user_id": "u3"}], + ) + + assert [r.success for r in response.data] == [True, False, True] + assert "insert failed for u2" in (response.data[1].error or "") + assert set(prisma.db.litellm_usertable.rows) == {"u1", "u3"} + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1"] + + +@pytest.mark.asyncio +async def test_insert_that_committed_but_lost_its_response_still_counts_as_created(): + prisma = _FakePrisma(teams=[_team("t1")], commit_then_drop=True) + response = await _run(prisma, [{"user_id": "u1", "teams": ["t1"]}, {"user_id": "u2"}]) + + assert [r.success for r in response.data] == [True, True] + assert [r.error for r in response.data] == [None, None] + assert set(prisma.db.litellm_usertable.rows) == {"u1", "u2"} + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u1"] + + +@pytest.mark.asyncio +async def test_user_id_taken_by_a_concurrent_request_is_not_claimed_by_this_batch(): + prisma = _FakePrisma(teams=[_team("t1")], raced_ids=frozenset({"u1"})) + response = await _run(prisma, [{"user_id": "u1", "teams": ["t1"]}, {"user_id": "u2", "teams": ["t1"]}]) + + assert [r.success for r in response.data] == [False, True] + assert "User id=u1 already exists" in (response.data[0].error or "") + assert prisma.db.litellm_usertable.rows["u1"].user_email == "u1@other-request.example" + assert [m.user_id for m in prisma.db.litellm_teamtable.rows["t1"].members_with_roles] == ["u2"] + + +@pytest.mark.asyncio +async def test_team_write_failure_keeps_user_and_reports_it_on_the_row(): + prisma = _FakePrisma(teams=[_team("t1"), _team("t2")]) + + async def explode(where, data): + raise RuntimeError("roster write failed") + + prisma.db.litellm_teamtable.update = explode + response = await _run(prisma, [{"user_id": "u1", "teams": ["t1", "t2"]}]) + + result = response.data[0] + assert result.success is True + assert result.teams == () + assert "t1" in (result.error or "") and "roster write failed" in (result.error or "") + assert prisma.db.litellm_usertable.rows["u1"].teams == [] + assert (response.meta.created, response.meta.failed) == (1, 0) + + +@pytest.mark.asyncio +async def test_keys_are_opt_in_per_row(): + prisma = _FakePrisma() + calls: list[dict[str, object]] = [] + + async def generate_key(**kwargs: object) -> dict[str, object]: + calls.append(kwargs) + return {"token": f"sk-{kwargs['user_id']}"} + + response = await _run( + prisma, + [ + {"user_id": "u1"}, + { + "user_id": "u2", + "auto_create_key": True, + "models": ["gpt-4o"], + "key_alias": "u2-key", + "blocked": True, + "permissions": {"get_spend_routes": True}, + "aliases": {"fast": "gpt-4o"}, + "config": {"tier": "gold"}, + "budget_fallbacks": {"gpt-4o": ["gpt-4o-mini"]}, + }, + {"user_id": "u3", "auto_create_key": False}, + ], + generate_key=generate_key, + ) + + assert [r.key for r in response.data] == [None, "sk-u2", None] + assert len(calls) == 1 + assert calls[0]["user_id"] == "u2" and calls[0]["table_name"] == "key" + assert calls[0]["models"] == ("gpt-4o",) and calls[0]["key_alias"] == "u2-key" + assert calls[0]["blocked"] is True + assert calls[0]["permissions"] == {"get_spend_routes": True} + assert calls[0]["aliases"] == {"fast": "gpt-4o"} + assert calls[0]["config"] == {"tier": "gold"} + assert calls[0]["budget_fallbacks"] == {"gpt-4o": ("gpt-4o-mini",)} + assert set(prisma.db.litellm_usertable.rows) == {"u1", "u2", "u3"} + + +@pytest.mark.asyncio +async def test_non_admin_cannot_create_admin_users_but_other_rows_proceed(): + prisma = _FakePrisma() + response = await _run( + prisma, + [{"user_id": "u1", "user_role": "proxy_admin"}, {"user_id": "u2", "user_role": "internal_user"}], + caller=INTERNAL, + ) + + assert [r.success for r in response.data] == [False, True] + assert "Only proxy admins" in (response.data[0].error or "") + assert set(prisma.db.litellm_usertable.rows) == {"u2"} + + +@pytest.mark.asyncio +async def test_license_is_checked_once_against_the_whole_batch(): + prisma = _FakePrisma() + prisma.db.litellm_usertable.rows["existing"] = _UserRow(user_id="existing") + license = _License(max_users=3) + + with pytest.raises(ManagementProblem) as exc: + await _run(prisma, [{"user_id": f"u{i}"} for i in range(3)], license=license) + + assert (exc.value.problem.status, exc.value.problem.type) == (403, "urn:litellm:error:license-limit-exceeded") + assert license.seen == [4] + assert set(prisma.db.litellm_usertable.rows) == {"existing"} + + ok = await _run(prisma, [{"user_id": f"u{i}"} for i in range(2)], license=license) + assert ok.meta.created == 2 + + resend = await _run(prisma, [{"user_id": f"u{i}"} for i in range(2)], license=license) + assert [r.success for r in resend.data] == [False, False] + assert all("already exists" in (r.error or "") for r in resend.data) + assert license.seen == [4, 3] + assert set(prisma.db.litellm_usertable.rows) == {"existing", "u0", "u1"} + + +def test_request_rejects_empty_oversized_and_invite_rows(): + with pytest.raises(ValidationError): + BulkNewUserRequest(users=[]) + with pytest.raises(ValidationError): + BulkNewUserRequest(users=[{"user_email": f"{i}@example.com"} for i in range(501)]) + with pytest.raises(ValidationError, match="send_invite_email"): + BulkNewUserItem(user_email="a@example.com", send_invite_email=True) + assert len(BulkNewUserRequest(users=[{"user_email": f"{i}@example.com"} for i in range(500)]).users) == 500 + assert BulkNewUserItem(user_email="a@example.com").auto_create_key is False + + +def test_request_rejects_unknown_fields_at_both_levels(): + with pytest.raises(ValidationError, match="extra_forbidden"): + BulkNewUserRequest(users=[{"user_email": "a@example.com", "user_emial": "typo"}]) + with pytest.raises(ValidationError, match="extra_forbidden"): + BulkNewUserRequest(users=[{"user_email": "a@example.com"}], dry_run=True) 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 d57bed430c1..11066d4ed38 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 @@ -5030,6 +5030,100 @@ async def test_websocket_passthrough_rewrites_gateway_alias_setup_model(): assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash" +@pytest.mark.parametrize( + "setup_model", + ["gemini-live-2.5-flash", "models/gemini-live-2.5-flash", "publishers/google/models/gemini-live-2.5-flash"], +) +def test_vertex_live_setup_model_resolves_before_extraction(setup_model): + """A bare gateway alias left the session logged as ``unknown`` at zero cost. + + The model was read off the raw client frame, and the extractor only yields a name when the string + already contains ``/models/``. The rewriter qualifies it a few lines later for the upstream, so a + client that addressed the gateway the documented way, by alias, logged no model and therefore + resolved no cost-map entry. Resolving first is what puts the real name on the logging object. + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _build_vertex_live_setup_model_rewriter, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _extract_model_from_vertex_ai_setup, + _resolved_vertex_live_setup, + ) + + rewriter = _build_vertex_live_setup_model_rewriter( + vertex_project="proj-db", vertex_location="global", llm_router=None + ) + setup_data = {"model": setup_model} + + resolved = _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, rewriter)) + + assert resolved == "gemini-live-2.5-flash", "an unresolved setup model logs the session as 'unknown'" + + +@pytest.mark.asyncio +async def test_websocket_passthrough_logs_a_bare_alias_setup_model(): + """End to end through the relay: a bare alias must reach the logging object as a real model name. + + This is the call-site half of the fix. The helper tests above pass even if extraction moves back + before the rewrite, so this one drives the real websocket relay and asserts on what got logged, + which is the name the cost map is looked up by. An unbilled session logs ``unknown``. + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _build_vertex_live_setup_model_rewriter, + ) + + upstream_ws = RecordingUpstreamWebSocket() + setup_frame = json.dumps({"setup": {"model": "gemini-live-2.5-flash"}}) + websocket = _client_websocket( + AsyncMock( + side_effect=[ + {"type": "websocket.receive", "text": setup_frame}, + {"type": "websocket.disconnect"}, + ] + ) + ) + built = [] + real_logging = litellm.litellm_core_utils.litellm_logging.Logging + + def _capture(*args, **kwargs): + obj = real_logging(*args, **kwargs) + built.append(obj) + return obj + + with _patched_websocket_passthrough_environment(upstream_ws): + with patch("litellm.litellm_core_utils.litellm_logging.Logging", side_effect=_capture): + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + setup_model_rewriter=_build_vertex_live_setup_model_rewriter( + vertex_project="proj-db", vertex_location="global", llm_router=None + ), + ) + + assert built, "the relay should have built a logging object" + assert built[0].model == "gemini-live-2.5-flash", "a bare alias must not log as 'unknown'" + + +def test_vertex_live_setup_resolution_is_inert_without_a_rewriter(): + """Non-Live passthrough routes pass no rewriter, so the frame must be handed over untouched.""" + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _extract_model_from_vertex_ai_setup, + _resolved_vertex_live_setup, + ) + + setup_data = {"model": "projects/p/locations/global/publishers/google/models/gemini-live-2.5-flash"} + + assert _resolved_vertex_live_setup(setup_data, None) is setup_data + assert _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, None)) == ( + "gemini-live-2.5-flash" + ) + + @pytest.mark.asyncio @pytest.mark.parametrize("rcvd_close", [None, "abnormal", "no_status"]) async def test_websocket_passthrough_does_not_relay_unsendable_upstream_close(rcvd_close): diff --git a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py index 4aea2e16364..53ea761daa7 100644 --- a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py +++ b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py @@ -227,7 +227,9 @@ async def test_otel_request_validation_exception_handler_empty_errors_invalid_pa async def test_otel_request_validation_exception_handler_returns_a_problem_on_the_control_plane(): """`/management/v1` answers validation errors as RFC 9457, so a caller there gets a 400 problem document rather than the proxy-wide 422 `{"detail": [...]}` shape.""" - errors = [{"loc": ["query", "page_size"], "msg": "Input should be less than or equal to 100", "type": "less_than_equal"}] + errors = [ + {"loc": ["query", "page_size"], "msg": "Input should be less than or equal to 100", "type": "less_than_equal"} + ] exc = RequestValidationError(errors) request = _make_request(path="/management/v1/spend_logs/end_users") @@ -242,6 +244,26 @@ async def test_otel_request_validation_exception_handler_returns_a_problem_on_th assert "detail" in body and not isinstance(body["detail"], list) +@pytest.mark.asyncio +async def test_otel_request_validation_exception_handler_answers_a_bad_control_plane_body_with_422(): + """A request body that fails validation, an unknown field included, is 422 on + `/management/v1`; only query parameter problems are 400.""" + errors = [ + {"loc": ["body", "users", 0, "user_emial"], "msg": "Extra inputs are not permitted", "type": "extra_forbidden"} + ] + exc = RequestValidationError(errors) + request = _make_request(path="/management/v1/users/bulk") + + response = await otel_request_validation_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 422 + assert response.media_type == "application/problem+json" + assert body["type"] == "urn:litellm:error:invalid-request-body" + assert body["status"] == 422 + assert "users.0.user_emial: Extra inputs are not permitted" in body["detail"] + + @pytest.mark.asyncio async def test_otel_request_validation_exception_handler_leaves_other_routes_on_422(): """The problem+json branch is scoped by path prefix. A route that merely contains @@ -249,9 +271,7 @@ async def test_otel_request_validation_exception_handler_leaves_other_routes_on_ exc = RequestValidationError([]) for path in ("/management", "/v1/management/foo", "/customer/list"): - response = await otel_request_validation_exception_handler( - request=_make_request(path=path), exc=exc - ) + response = await otel_request_validation_exception_handler(request=_make_request(path=path), exc=exc) assert response.status_code == 422, path assert json.loads(response.body) == {"detail": []}, path @@ -294,6 +314,4 @@ async def test_otel_unhandled_exception_handler_reraises_proxy_exception_error() async def test_otel_unhandled_exception_handler_reraises_http_exception_invalid(): request = _make_request() with pytest.raises(HTTPException): - await otel_unhandled_exception_handler( - request=request, exc=HTTPException(status_code=418, detail="teapot") - ) + await otel_unhandled_exception_handler(request=request, exc=HTTPException(status_code=418, detail="teapot")) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 3854beb9370..cabfcc9918f 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1402,6 +1402,51 @@ class TestProxyBaseLLMRequestProcessing: assert "x-litellm-key-spend" in headers_7 assert float(headers_7["x-litellm-key-spend"]) == 0.001 # Should use original spend on error + @pytest.mark.parametrize( + ("hidden_params", "request_data", "expected_call_id"), + [ + ( + {"litellm_call_id": "call-from-hidden-params"}, + {"litellm_call_id": "call-from-request"}, + "call-from-hidden-params", + ), + ({}, {"litellm_call_id": "call-from-request"}, "call-from-request"), + ({"model_id": "m-1"}, {"litellm_call_id": "call-from-request"}, "call-from-request"), + ], + ) + def test_get_custom_headers_call_id_falls_back_to_hidden_params_then_request_data( + self, hidden_params, request_data, expected_call_id + ): + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0.0 + + headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + hidden_params=hidden_params, + request_data=request_data, + ) + + assert headers["x-litellm-call-id"] == expected_call_id + + def test_get_custom_headers_explicit_call_id_wins_over_fallbacks(self): + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0.0 + + headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + call_id="explicit-call-id", + hidden_params={"litellm_call_id": "call-from-hidden-params"}, + request_data={"litellm_call_id": "call-from-request"}, + ) + + assert headers["x-litellm-call-id"] == "explicit-call-id" + @pytest.mark.asyncio async def test_queue_time_seconds_is_set_in_metadata(self, monkeypatch): """ @@ -6917,6 +6962,45 @@ class TestModelDeploymentsSupportStreamOptions: assert self._support(None, None) is False +@pytest.mark.asyncio +@pytest.mark.parametrize("key_settings, expected", [ + (None, {"group": {"team": 100}}), + ({"weights": {"group": {"key": 100}}}, {"group": {"key": 100}}), + ({"timeout": 30}, None), + ({"weights": {"group": {"key": "legacy"}}}, None), +]) +async def test_saved_weights_override_caller_input_and_preserve_key_precedence( + monkeypatch: pytest.MonkeyPatch, + key_settings: dict[str, int | dict[str, dict[str, int | str]]] | None, + expected: dict[str, dict[str, int]] | None, +) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "get_team_object", AsyncMock( + return_value=SimpleNamespace(router_settings={"weights": {"group": {"team": 100}}}) + )) + forged = {"group": {"caller": 100}} + processor = ProxyBaseLLMRequestProcessing(data={ + "model": "group", "weights": forged, "_router_weights": forged, + "router_settings_override": {"weights": forged}, + }) + logging = MagicMock(spec=ProxyLogging) + logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"]) + data, _ = await processor.common_processing_pre_call_logic( + request=Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []}), + general_settings={}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", team_id="team-a", router_settings=key_settings), + proxy_logging_obj=logging, + proxy_config=proxy_server.ProxyConfig(), + route_type="acompletion", + llm_router=litellm.Router(model_list=[]), + ) + assert "weights" not in data + assert data.get("_router_weights") == expected + assert logging.pre_call_hook.call_args.kwargs["data"].get("_router_weights") == expected + + class TestPerRequestModelGroupAlias: """``router_settings.model_group_alias`` on a key or team has to be resolved by the proxy: the Router resolves aliases from its own shared instance diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index c0c853ae2c5..7d6c5d3cebe 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -374,7 +374,7 @@ async def test_save_background_health_checks_to_db(): """Test the main background health check save function""" mock_prisma = MagicMock() mock_prisma.save_health_check_result = AsyncMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) model_list = [ { @@ -398,9 +398,9 @@ async def test_save_background_health_checks_to_db(): "background_health_check", ) - # Should call get_all_latest_health_checks and save_health_check_result, and report completion + # Should read the latest rows and save_health_check_result, and report completion assert persisted is True - mock_prisma.get_all_latest_health_checks.assert_called_once() + mock_prisma.db.query_raw.assert_awaited_once() mock_prisma.save_health_check_result.assert_called_once() call_kwargs = mock_prisma.save_health_check_result.call_args[1] @@ -493,7 +493,7 @@ def _one_model_setup(): @pytest.mark.asyncio async def test_save_background_health_checks_to_db_returns_false_when_a_write_fails(): mock_prisma = MagicMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.save_health_check_result = AsyncMock(return_value=None) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() @@ -504,6 +504,23 @@ async def test_save_background_health_checks_to_db_returns_false_when_a_write_fa assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 1) +@pytest.mark.asyncio +async def test_save_background_health_checks_to_db_writes_nothing_when_the_latest_row_read_fails(mock_prisma): + """ + A failed dedup read must not read as an empty table. Treated that way, every model was written on every + cycle by every pod while the read kept failing, which is what filled the table in production. + """ + mock_prisma.db.query_raw = AsyncMock(side_effect=RuntimeError("db down")) + mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) + model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() + + persisted = await _save_background_health_checks_to_db( + mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 0) + + @pytest.mark.asyncio async def test_save_background_health_checks_to_db_no_prisma(): """Test graceful handling when no prisma client""" @@ -515,7 +532,7 @@ async def test_save_background_health_checks_to_db_no_prisma(): async def test_save_background_health_checks_to_db_exception_handling(): """Test exception handling in background health check save""" mock_prisma = MagicMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(side_effect=Exception("DB Error")) + mock_prisma.db.query_raw = AsyncMock(side_effect=Exception("DB Error")) model_list = [ { diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 37b983d709a..099afa57eec 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -957,6 +957,8 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "litellm_gateway_injected_cache": "forged-deployment-id", "metadata": copy.deepcopy(malicious_metadata), "litellm_metadata": copy.deepcopy(malicious_metadata), + "weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}}, + "_router_weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}}, } updated = await add_litellm_data_to_request( @@ -974,6 +976,10 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): assert "enable_prompt_caching" not in updated assert "routing_decision" not in updated assert "litellm_gateway_injected_cache" not in updated + assert "weights" not in updated + assert "_router_weights" not in updated + assert "weights" not in updated["proxy_server_request"]["body"] + assert "_router_weights" not in updated["proxy_server_request"]["body"] stripped_keys = { "disable_global_guardrails", @@ -7346,13 +7352,14 @@ def _reserved_stamp_key(key_metadata: dict | None = None) -> UserAPIKeyAuth: _PLANTED_STAMPS = { "attempted_fallbacks": 99, "original_model_group": "spoofed-group", + "request_retry_count": -100, "_client_output_ceiling": {"api_base": "https://attacker.example"}, "client_key": "client_value", } @pytest.mark.asyncio -async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets(): +async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_both_buckets() -> None: """attempted_fallbacks and original_model_group are router-written facts the spend row reads back; a client planting them in either bucket is dropped at the boundary so the router never sees a reserved key it did not write.""" @@ -7378,11 +7385,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_bo assert "attempted_fallbacks" not in updated["metadata"] assert "original_model_group" not in updated["metadata"] assert "_client_output_ceiling" not in updated["metadata"] + assert "request_retry_count" not in updated["metadata"] assert updated["metadata"]["client_key"] == "client_value" @pytest.mark.asyncio -async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata(): +async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_json_string_litellm_metadata() -> None: from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request data = { @@ -7403,11 +7411,12 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_js assert "litellm_metadata" not in updated assert "attempted_fallbacks" not in updated["metadata"] assert "original_model_group" not in updated["metadata"] + assert "request_retry_count" not in updated["metadata"] assert updated["metadata"]["client_key"] == "client_value" @pytest.mark.asyncio -async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in(): +async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite_pricing_override_opt_in() -> None: """The pricing strip is gated on allow_client_pricing_override; the reserved-stamp strip is not, because no key or team setting makes a client-written fallback count valid.""" from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request @@ -7431,6 +7440,7 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_despite assert updated["metadata"]["model_info"] == {"input_cost_per_token": 0.0} assert "attempted_fallbacks" not in updated["metadata"] assert "original_model_group" not in updated["metadata"] + assert "request_retry_count" not in updated["metadata"] @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e09dddfec5b..04173ced776 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -12795,6 +12795,45 @@ async def test_moderations_reraises_proxy_exception_unwrapped(): mock_logging.post_call_failure_hook.assert_awaited_once() +@pytest.mark.asyncio +async def test_moderations_response_carries_litellm_call_id_header(): + from fastapi import Response + + from litellm.types.utils import ModerationCreateResponse + + call_id = "moderation-call-id-123" + moderation_response = ModerationCreateResponse(id="modr-1", model="omni-moderation-latest", results=[]) + moderation_response._hidden_params = {"litellm_call_id": call_id, "model_id": "mod-deployment-1"} + + async def fake_llm_call(): + return moderation_response + + async def passthrough_add_litellm_data(data, **kwargs): + return {**data, "litellm_call_id": call_id} + + request = MagicMock() + request.body = AsyncMock(return_value=b'{"input": "hi"}') + fastapi_response = Response() + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", spend=0.0) + + with ( + patch.object(proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data), # test-quality-ok: the route reads this module global, no injection point + patch.object(proxy_server_module, "route_request", new=AsyncMock(return_value=fake_llm_call())), # test-quality-ok: fakes the provider call so the response headers assembled by the real route are observable + patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global, no injection point + ): + mock_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data) + mock_logging.update_request_status = AsyncMock() + result = await proxy_server_module.moderations( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) + + assert result is moderation_response + assert fastapi_response.headers["x-litellm-call-id"] == call_id + assert fastapi_response.headers["x-litellm-model-id"] == "mod-deployment-1" + + @pytest.mark.asyncio async def test_init_agents_in_db_rebuilds_registry_under_agent_reconcile_lock(monkeypatch): from litellm.proxy.agent_endpoints.agent_registry import ( diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index 634b90e445a..5d5273be243 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -276,3 +276,13 @@ def test_a_server_only_marker_is_not_taken_from_the_caller(field, forged, defaul auth = UserAPIKeyAuth(api_key="sk-1234", **{field: forged}) assert getattr(auth, field) == default + + +@pytest.mark.parametrize("weight", [True, "1", -1, 0, float("inf")]) +def test_key_and_team_weights_reject_invalid_numeric_values(weight: bool | str | int | float) -> None: + from pydantic import ValidationError + from litellm.proxy._types import GenerateKeyRequest, NewTeamRequest + + for request_type in (GenerateKeyRequest, NewTeamRequest): + with pytest.raises(ValidationError): + request_type(router_settings={"weights": {"group": {"id": weight}}}) diff --git a/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py b/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py index 100ba653f3a..0072997a0d9 100644 --- a/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py +++ b/tests/test_litellm/proxy/types_utils/test_db_overlay_remote_module_scrub.py @@ -43,6 +43,7 @@ def test_litellm_settings_callback_list_strips_remote_urls(field): "custom_auth", "custom_key_generate", "custom_key_update", + "custom_key_policy", "custom_sso", "custom_ui_sso_sign_in_handler", ], diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index d3d41c5b54b..643e65af47c 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch import pytest import litellm +from litellm.models.credentials import CredentialItem from litellm.realtime_api import main as realtime_main from litellm.realtime_api.main import _with_resolved_session_model @@ -224,9 +225,11 @@ def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch class _CapturingConnect: def __init__(self) -> None: self.url: str | None = None + self.kwargs: dict[str, object] = {} def __call__(self, url: str, **kwargs: object) -> "_CapturingConnect": self.url = url + self.kwargs = kwargs return self async def __aenter__(self) -> MagicMock: @@ -241,6 +244,72 @@ class _CapturingConnect: return None +@pytest.mark.asyncio +async def test_azure_health_check_resolves_stored_credentials(monkeypatch): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="azure-rt", + credential_values={ + "api_key": "sk-from-credential", + "api_base": "https://example.openai.azure.com", + "api_version": "2025-04-01-preview", + }, + credential_info={}, + ) + ], + ) + connect = _CapturingConnect() + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gpt-realtime", + custom_llm_provider="azure", + api_key=None, + realtime_protocol="beta", + model_params={"model": "azure/gpt-realtime", "litellm_credential_name": "azure-rt"}, + ) + assert connect.kwargs["additional_headers"] == {"api-key": "sk-from-credential"} + assert connect.url is not None + assert connect.url.startswith("wss://example.openai.azure.com") + assert "api-version=2025-04-01-preview" in connect.url + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("custom_llm_provider", "model", "expected_url"), + [ + ("xai", "grok-voice-latest", "wss://api.x.ai/v1/realtime?model=grok-voice-latest"), + ("openai", "gpt-realtime", "wss://api.openai.com/v1/realtime?model=gpt-realtime"), + ], +) +async def test_bearer_health_check_sends_stored_credential_as_bearer_token( + monkeypatch, custom_llm_provider: str, model: str, expected_url: str +): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="voice-key", + credential_values={"api_key": "sk-from-credential"}, + credential_info={}, + ) + ], + ) + connect = _CapturingConnect() + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model=model, + custom_llm_provider=custom_llm_provider, + api_key=None, + model_params={"model": f"{custom_llm_provider}/{model}", "litellm_credential_name": "voice-key"}, + ) + assert connect.kwargs["additional_headers"] == {"Authorization": "Bearer sk-from-credential"} + assert connect.url == expected_url + + @pytest.mark.asyncio async def test_azure_health_check_probes_ga_transcription_url_for_transcription_model(local_model_cost_map): """Regression for LIT-6240: transcription-only models (mode audio_transcription diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py b/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py index 2cfec6a1844..b78dabbfe48 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py @@ -68,3 +68,105 @@ async def test_async_fallback_tags_skip_responses_api_bridge(): await coro assert captured.get("_skip_responses_api_bridge") is True + + +_CODEX_ADDITIONAL_TOOLS_ITEM = { + "type": "additional_tools", + "id": "at_codex", + "role": "developer", + "tools": [ + { + "type": "namespace", + "name": "functions", + "description": "", + "tools": [ + { + "type": "custom", + "name": "exec", + "description": "Runs a shell command.", + "format": {"type": "grammar", "syntax": "lark", "definition": "start: /.+/"}, + }, + { + "type": "function", + "name": "wait", + "description": "Waits for a background command.", + "parameters": {"type": "object", "properties": {"id": {"type": "string"}}}, + }, + ], + } + ], +} +_CODEX_INPUT = [_CODEX_ADDITIONAL_TOOLS_ITEM, {"type": "message", "role": "user", "content": "Run ls"}] + + +def test_sync_fallback_hoists_additional_tools_input_items_into_chat_tools(): + handler = LiteLLMCompletionTransformationHandler() + captured: dict = {} + + def fake_completion(**kwargs): + captured.update(kwargs) + raise _StopForwarding() + + with patch("litellm.completion", fake_completion): # test-quality-ok: no DI seam; the file stubs this same boundary + with pytest.raises(_StopForwarding): + handler.response_api_handler( + model="bedrock/us.openai.gpt-5.6", + input=_CODEX_INPUT, + responses_api_request={}, + custom_llm_provider="bedrock", + _is_async=False, + ) + + assert [message["role"] for message in captured["messages"]] == ["user"] + functions_by_name = {tool["function"]["name"]: tool["function"] for tool in captured["tools"]} + assert set(functions_by_name) == {"exec", "functions__wait"} + assert set(functions_by_name["exec"]["parameters"]["properties"]) == {"content"} + + +@pytest.mark.asyncio +async def test_async_fallback_returns_hoisted_nested_custom_tool_call_as_custom_tool_call(): + from litellm.responses.litellm_completion_transformation.transformation import TOOL_CALLS_CACHE + from litellm.types.utils import ChatCompletionMessageToolCall, Choices, Function, Message, ModelResponse + + handler = LiteLLMCompletionTransformationHandler() + tool_call_id = "call_exec_hoisted" + + async def fake_acompletion(**kwargs): + return ModelResponse( + id="chatcmpl-exec", + created=1, + model="us.openai.gpt-5.6", + object="chat.completion", + choices=[ + Choices( + finish_reason="tool_calls", + index=0, + message=Message( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionMessageToolCall( + id=tool_call_id, + type="function", + function=Function(name="exec", arguments='{"content": "ls"}'), + ) + ], + ), + ) + ], + ) + + try: + with patch("litellm.acompletion", fake_acompletion): # test-quality-ok: no DI seam; file stubs this boundary + response = await handler.response_api_handler( + model="bedrock/us.openai.gpt-5.6", + input=_CODEX_INPUT, + responses_api_request={}, + custom_llm_provider="bedrock", + _is_async=True, + ) + finally: + TOOL_CALLS_CACHE.delete_cache(key=tool_call_id) + + tool_calls = [(item.type, item.name, item.input) for item in response.output if item.type == "custom_tool_call"] + assert tool_calls == [("custom_tool_call", "exec", "ls")] diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index d5a21bccad7..0ed101952be 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -2506,6 +2506,7 @@ class TestToolTransformation: "tools": [ "ignored", {"type": "namespace", "name": "ignored"}, + {"type": "web_search", "name": "ignored"}, { "type": "function", "name": "spawn_agent", @@ -2527,6 +2528,36 @@ class TestToolTransformation: "type": "object", } + def test_transform_nested_namespace_custom_tool_becomes_a_content_function_under_its_short_name(self): + namespace_tool = { + "type": "namespace", + "name": "functions", + "description": "Codex shell tools.", + "tools": [ + { + "type": "custom", + "name": "exec", + "description": "Runs a shell command.", + "format": {"type": "grammar", "syntax": "lark", "definition": "start: /.+/"}, + }, + ], + } + + result_tools, _ = ( + LiteLLMCompletionResponsesConfig.transform_responses_api_tools_to_chat_completion_tools( + tools=[namespace_tool] + ) + ) + + assert len(result_tools) == 1 + function = result_tools[0]["function"] + assert function["name"] == "exec" + assert function["description"].startswith("Codex shell tools.") + assert "Runs a shell command." in function["description"] + assert "start: /.+/" in function["description"] + assert function["parameters"]["required"] == ["content"] + assert function["parameters"]["properties"]["content"]["type"] == "string" + @pytest.mark.parametrize( "model, custom_llm_provider", [ @@ -3788,6 +3819,143 @@ class TestEnsureOutputItemContentPartAdded: assert added.item.name == "spawn_agent" assert added.item.namespace == "collaboration" + def test_streaming_nested_custom_tool_call_comes_back_as_custom_tool_call(self): + from litellm.responses.litellm_completion_transformation.custom_tools import extract_custom_tool_names + + iterator = self._make_iterator() + iterator.responses_api_request = { + "tools": [ + { + "type": "namespace", + "name": "functions", + "tools": [ + { + "type": "custom", + "name": "exec", + "format": {"type": "grammar", "syntax": "lark", "definition": "start: /.+/"}, + } + ], + } + ] + } + iterator._custom_tool_names = extract_custom_tool_names(iterator.responses_api_request.get("tools")) + iterator._namespace_tool_names = LiteLLMCompletionResponsesConfig.namespace_tool_name_map( + iterator.responses_api_request.get("tools") + ) + + iterator._queue_tool_call_delta_events( + [{"index": 0, "id": "call_exec", "function": {"name": "exec", "arguments": '{"content":"ls"}'}}] + ) + iterator._queue_final_tool_call_done_events( + ModelResponse( + id="chatcmpl-exec", + created=1, + model="us.openai.gpt-5.6", + object="chat.completion", + choices=[ + Choices( + finish_reason="tool_calls", + index=0, + message=Message( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionMessageToolCall( + id="call_exec", + type="function", + function=Function(name="exec", arguments='{"content":"ls"}'), + ) + ], + ), + ) + ], + ) + ) + + added = iterator._pending_tool_events[0] + assert added.item.type == "custom_tool_call" + assert added.item.name == "exec" + done = iterator._pending_tool_events[-1] + assert done.item.type == "custom_tool_call" + assert done.item.input == "ls" + + def test_streaming_namespaced_function_sharing_a_nested_custom_short_name_stays_a_function_call(self): + from litellm.responses.litellm_completion_transformation.custom_tools import extract_custom_tool_names + + iterator = self._make_iterator() + iterator.responses_api_request = { + "tools": [ + { + "type": "namespace", + "name": "alpha", + "tools": [ + { + "type": "custom", + "name": "run", + "format": {"type": "grammar", "syntax": "lark", "definition": "start: /.+/"}, + } + ], + }, + { + "type": "namespace", + "name": "beta", + "tools": [ + { + "type": "function", + "name": "run", + "parameters": {"type": "object", "properties": {"job_id": {"type": "string"}}}, + } + ], + }, + ] + } + iterator._custom_tool_names = extract_custom_tool_names(iterator.responses_api_request.get("tools")) + iterator._namespace_tool_names = LiteLLMCompletionResponsesConfig.namespace_tool_name_map( + iterator.responses_api_request.get("tools") + ) + function_call = {"id": "call_fn", "function": {"name": "beta__run", "arguments": '{"job_id":"42"}'}} + custom_call = {"id": "call_custom", "function": {"name": "run", "arguments": '{"content":"echo hi"}'}} + + iterator._queue_tool_call_delta_events([{"index": 0, **function_call}, {"index": 1, **custom_call}]) + iterator._queue_final_tool_call_done_events( + ModelResponse( + id="chatcmpl-run", + created=1, + model="us.openai.gpt-5.6", + object="chat.completion", + choices=[ + Choices( + finish_reason="tool_calls", + index=0, + message=Message( + content=None, + role="assistant", + tool_calls=[ + ChatCompletionMessageToolCall( + id=call["id"], type="function", function=Function(**call["function"]) + ) + for call in (function_call, custom_call) + ], + ), + ) + ], + ) + ) + + items = [ + event.item + for event in iterator._pending_tool_events + if event.type in ("response.output_item.added", "response.output_item.done") + ] + function_items = [item for item in items if item.call_id == "call_fn"] + custom_items = [item for item in items if item.call_id == "call_custom"] + assert len(function_items) == 2 and len(custom_items) == 2 + assert all((item.type, item.name, item.namespace) == ("function_call", "run", "beta") for item in function_items) + assert function_items[-1].arguments == '{"job_id":"42"}' + assert all(item.type == "custom_tool_call" and item.name == "run" for item in custom_items) + assert all(getattr(item, "namespace", None) is None for item in custom_items) + assert custom_items[-1].input == "echo hi" + def test_streaming_unqualified_namespace_tool_calls_restore_namespace(self): """A unique nested tool name without the namespace still maps back.""" iterator = self._make_iterator() diff --git a/tests/test_litellm/responses/test_additional_tools.py b/tests/test_litellm/responses/test_additional_tools.py new file mode 100644 index 00000000000..bef3b27eacd --- /dev/null +++ b/tests/test_litellm/responses/test_additional_tools.py @@ -0,0 +1,48 @@ +from litellm.responses.additional_tools import hoist_additional_tools + +_EXEC_TOOL = {"type": "custom", "name": "exec", "format": {"type": "grammar", "syntax": "lark", "definition": "start: /.+/"}} +_WAIT_TOOL = {"type": "function", "name": "wait", "parameters": {"type": "object", "properties": {}}} +_TOP_LEVEL_TOOL = {"type": "function", "name": "top_level", "parameters": {"type": "object", "properties": {}}} +_USER_MESSAGE = {"type": "message", "role": "user", "content": "Run ls"} + + +def test_string_input_passes_through_with_existing_tools(): + hoisted = hoist_additional_tools("hello", [_TOP_LEVEL_TOOL]) + + assert hoisted.input == "hello" + assert hoisted.tools == (_TOP_LEVEL_TOOL,) + assert hoisted.hoisted == () + + +def test_input_without_additional_tools_items_is_returned_untouched(): + request_input = [_USER_MESSAGE] + + hoisted = hoist_additional_tools(request_input, None) + + assert hoisted.input is request_input + assert hoisted.tools == () + assert hoisted.hoisted == () + + +def test_additional_tools_items_are_stripped_and_appended_after_top_level_tools_in_item_order(): + request_input = [ + {"type": "additional_tools", "id": "at_1", "role": "developer", "tools": [_EXEC_TOOL]}, + _USER_MESSAGE, + {"type": "additional_tools", "id": "at_2", "role": "developer", "tools": [_WAIT_TOOL]}, + ] + + hoisted = hoist_additional_tools(request_input, [_TOP_LEVEL_TOOL]) + + assert hoisted.input == [_USER_MESSAGE] + assert hoisted.tools == (_TOP_LEVEL_TOOL, _EXEC_TOOL, _WAIT_TOOL) + assert hoisted.hoisted == (_EXEC_TOOL, _WAIT_TOOL) + + +def test_additional_tools_item_without_a_tools_list_is_stripped_and_contributes_nothing(): + request_input = [{"type": "additional_tools", "id": "at_1", "role": "developer", "tools": "exec"}, _USER_MESSAGE] + + hoisted = hoist_additional_tools(request_input, None) + + assert hoisted.input == [_USER_MESSAGE] + assert hoisted.tools == () + assert hoisted.hoisted == () diff --git a/tests/test_litellm/responses/test_custom_tool_call.py b/tests/test_litellm/responses/test_custom_tool_call.py index e80301c3b2f..2ed71ee3ecf 100644 --- a/tests/test_litellm/responses/test_custom_tool_call.py +++ b/tests/test_litellm/responses/test_custom_tool_call.py @@ -55,6 +55,24 @@ class TestCustomToolUtilities: names = extract_custom_tool_names(tools) assert names == set() + def test_extract_custom_tool_names_walks_namespace_tools(self): + tools = [ + {"type": "function", "name": "regular_tool"}, + { + "type": "namespace", + "name": "functions", + "tools": [ + {"type": "custom", "name": "exec"}, + {"type": "function", "name": "wait"}, + "ignored", + ], + }, + {"type": "namespace", "name": "empty", "tools": "not-a-list"}, + ] + + names = extract_custom_tool_names(tools) + assert names == {"exec"} + def test_extract_custom_tool_names_none(self): """Test extraction with None input.""" names = extract_custom_tool_names(None) diff --git a/tests/test_litellm/responses/test_streaming_iterator_error_events.py b/tests/test_litellm/responses/test_streaming_iterator_error_events.py index ad74861c096..3d7c220804a 100644 --- a/tests/test_litellm/responses/test_streaming_iterator_error_events.py +++ b/tests/test_litellm/responses/test_streaming_iterator_error_events.py @@ -1,9 +1,14 @@ """ Regression: in-stream error events (type="error", type="response.failed") must raise instead of being returned as benign chunks, mirroring chat streaming -semantics (_handle_stream_fallback_error): non-retriable 4xx (except 429) -raise litellm.APIError directly; 429 and 5xx are wrapped in -MidStreamFallbackError so the Router's mid-stream fallback machinery fires. +semantics (_handle_stream_fallback_error). The event's code, type and status go +through litellm.exception_type, so each event raises the same typed exception +the non-streaming path raises for that provider error: non-retriable 4xx +(except 429) raise that typed exception directly, so a context-length event +surfaces as ContextWindowExceededError(400) with no MidStreamFallbackError +wrapping, while 429, 5xx and ContentPolicyViolationError are wrapped in +MidStreamFallbackError so the Router's mid-stream fallback machinery fires and +its content_policy_fallbacks dispatch sees the trigger it matches on. Status mapping must consider both the OpenAI error `type` (e.g. "invalid_request_error") and `code` (e.g. "invalid_prompt", @@ -66,12 +71,12 @@ def test_maybe_raise_for_error_event_wraps_unknown_error_in_mid_stream_fallback( with pytest.raises(MidStreamFallbackError) as exc_info: iterator._maybe_raise_for_error_event(chunk) assert exc_info.value.status_code == 500 - assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert isinstance(exc_info.value.original_exception, litellm.InternalServerError) assert exc_info.value.original_exception.status_code == 500 def test_maybe_raise_for_error_event_maps_rate_limit_code_to_429_mid_stream_fallback(): - """429 is retriable: it must be wrapped so the Router can fall back, carrying the mapped APIError.""" + """429 is retriable: it must be wrapped so the Router can fall back, carrying the mapped RateLimitError.""" iterator = _make_iterator() chunk = _make_error_chunk("tokens", "rate_limit_exceeded", "Too many requests") with pytest.raises(MidStreamFallbackError) as exc_info: @@ -79,15 +84,15 @@ def test_maybe_raise_for_error_event_maps_rate_limit_code_to_429_mid_stream_fall assert exc_info.value.status_code == 429 assert exc_info.value.generated_content == "" assert exc_info.value.is_pre_first_chunk is True - assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert isinstance(exc_info.value.original_exception, litellm.RateLimitError) assert exc_info.value.original_exception.status_code == 429 def test_maybe_raise_for_error_event_maps_invalid_request_type_to_400(): - """Client errors classified via the `type` field must raise APIError directly (no fallback).""" + """Client errors classified via the `type` field must raise BadRequestError directly (no fallback).""" iterator = _make_iterator() chunk = _make_error_chunk("invalid_request_error", "invalid_prompt", "bad request") - with pytest.raises(litellm.APIError) as exc_info: + with pytest.raises(litellm.BadRequestError) as exc_info: iterator._maybe_raise_for_error_event(chunk) assert exc_info.value.status_code == 400 assert not isinstance(exc_info.value, MidStreamFallbackError) @@ -99,12 +104,86 @@ def test_maybe_raise_for_error_event_maps_context_length_code_to_400(): chunk = Mock() chunk.type = "error" chunk.error = {"code": "context_length_exceeded", "message": "too long"} - with pytest.raises(litellm.APIError) as exc_info: + with pytest.raises(litellm.BadRequestError) as exc_info: iterator._maybe_raise_for_error_event(chunk) assert exc_info.value.status_code == 400 assert not isinstance(exc_info.value, MidStreamFallbackError) +def test_maybe_raise_for_error_event_raises_context_window_exceeded_directly(): + """A context-length error event maps to ContextWindowExceededError exactly like the non-streaming + path and, being a non-retriable client error, is raised directly rather than wrapped for mid-stream + fallback, preserving the direct-SDK 400 contract from issue #15785.""" + iterator = _make_iterator() + chunk = _make_error_chunk( + "invalid_request_error", + "context_length_exceeded", + "This model's maximum context length is 128000 tokens. However, your messages resulted in 130000 tokens.", + ) + with pytest.raises(litellm.ContextWindowExceededError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert exc_info.value.status_code == 400 + assert not isinstance(exc_info.value, MidStreamFallbackError) + assert "maximum context length" in str(exc_info.value) + + +CONTENT_POLICY_MESSAGE = "This content was flagged for possible cybersecurity risk. The response was halted mid-stream." + + +@pytest.mark.parametrize("custom_llm_provider", ["openai", "azure"]) +def test_maybe_raise_for_error_event_wraps_content_policy_violation_for_content_policy_fallbacks( + custom_llm_provider: str, +): + """Regression: a content_policy_violation error event used to raise a bare APIError, so the Router's + content_policy_fallbacks never fired. It must map to ContentPolicyViolationError (the same exception the + non-streaming path raises) and be wrapped so the Router's mid-stream fallback catches it.""" + iterator = _make_iterator() + iterator.custom_llm_provider = custom_llm_provider + chunk = _make_error_chunk("invalid_request_error", "content_policy_violation", CONTENT_POLICY_MESSAGE) + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert isinstance(exc_info.value.original_exception, litellm.ContentPolicyViolationError) + assert exc_info.value.original_exception.status_code == 400 + assert exc_info.value.status_code == 400 + assert exc_info.value.is_pre_first_chunk is True + assert CONTENT_POLICY_MESSAGE in str(exc_info.value.original_exception) + + +def test_maybe_raise_for_response_failed_event_wraps_content_policy_violation(): + iterator = _make_iterator() + chunk = _make_failed_chunk( + {"type": "invalid_request_error", "code": "content_policy_violation", "message": CONTENT_POLICY_MESSAGE} + ) + with pytest.raises(MidStreamFallbackError) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + assert isinstance(exc_info.value.original_exception, litellm.ContentPolicyViolationError) + + +@pytest.mark.parametrize( + "error_type,error_code,expected_exception", + [ + ("invalid_request_error", "content_policy_violation", litellm.ContentPolicyViolationError), + ("tokens", "rate_limit_exceeded", litellm.RateLimitError), + ("invalid_request_error", "insufficient_quota", litellm.RateLimitError), + ("server_error", "internal_error", litellm.InternalServerError), + ("invalid_request_error", "invalid_prompt", litellm.BadRequestError), + ("invalid_request_error", "model_not_found", litellm.NotFoundError), + ("server_error", "vector_store_timeout", litellm.Timeout), + ], +) +def test_error_event_raises_the_same_typed_exception_as_the_non_streaming_path( + error_type: str, error_code: str, expected_exception: type[Exception] +): + iterator = _make_iterator() + chunk = _make_error_chunk(error_type, error_code, "provider message") + with pytest.raises((MidStreamFallbackError, expected_exception)) as exc_info: + iterator._maybe_raise_for_error_event(chunk) + raised = exc_info.value + typed_exception = raised.original_exception if isinstance(raised, MidStreamFallbackError) else raised + assert type(typed_exception) is expected_exception + assert "provider message" in str(typed_exception) + + def test_maybe_raise_for_error_event_maps_insufficient_quota_to_429(): """OpenAI returns HTTP 429 for insufficient_quota; it must not map to 400 even though its type is invalid_request_error-adjacent, and it must be wrapped for fallback.""" @@ -113,6 +192,7 @@ def test_maybe_raise_for_error_event_maps_insufficient_quota_to_429(): with pytest.raises(MidStreamFallbackError) as exc_info: iterator._maybe_raise_for_error_event(chunk) assert exc_info.value.status_code == 429 + assert isinstance(exc_info.value.original_exception, litellm.RateLimitError) def test_maybe_raise_for_error_event_passes_through_normal_chunk(): @@ -186,10 +266,40 @@ async def test_async_iterator_raises_mid_stream_fallback_on_rate_limit_error_eve assert exc_info.value.status_code == 429 assert exc_info.value.is_pre_first_chunk is True assert exc_info.value.generated_content == "" - assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert isinstance(exc_info.value.original_exception, litellm.RateLimitError) assert exc_info.value.original_exception.status_code == 429 +@pytest.mark.asyncio +async def test_async_iterator_content_policy_violation_after_first_chunk_carries_generated_content(): + """The customer's case: text streams, then the provider halts the stream with a + content_policy_violation error event. The iterator must surface ContentPolicyViolationError + inside MidStreamFallbackError, together with the text already streamed.""" + iterator = _make_async_iterator_with_events( + [ + {"type": "response.output_text.delta", "delta": "partial "}, + { + "type": "error", + "error": { + "type": "invalid_request_error", + "code": "content_policy_violation", + "message": CONTENT_POLICY_MESSAGE, + }, + }, + ] + ) + + stream = aiter(iterator) + first_chunk = await anext(stream) + assert first_chunk is not None + + with pytest.raises(MidStreamFallbackError) as exc_info: + await anext(stream) + assert isinstance(exc_info.value.original_exception, litellm.ContentPolicyViolationError) + assert exc_info.value.is_pre_first_chunk is False + assert exc_info.value.generated_content == "partial " + + @pytest.mark.asyncio async def test_async_iterator_error_after_first_chunk_carries_generated_content(): """An error after streamed output must expose the accumulated text so the router's @@ -205,14 +315,13 @@ async def test_async_iterator_error_after_first_chunk_carries_generated_content( ] ) - chunks = [] - async def _drain(): - async for chunk in iterator: - chunks.append(chunk) + stream = aiter(iterator) + first_chunk = await anext(stream) + second_chunk = await anext(stream) + assert first_chunk is not None and second_chunk is not None with pytest.raises(MidStreamFallbackError) as exc_info: - await _drain() - assert len(chunks) == 2 + await anext(stream) assert exc_info.value.status_code == 500 assert exc_info.value.is_pre_first_chunk is False assert exc_info.value.generated_content == "hello world" @@ -265,7 +374,7 @@ def test_handle_logging_failed_response_maps_rate_limit_to_429(): ): iterator._handle_logging_failed_response() logged_exception = mock_run_async.call_args.kwargs["exception"] - assert isinstance(logged_exception, litellm.APIError) + assert isinstance(logged_exception, litellm.RateLimitError) assert logged_exception.status_code == 429 assert "throttled" in str(logged_exception) @@ -282,10 +391,28 @@ def test_handle_logging_failed_response_maps_type_field_to_400(): ): iterator._handle_logging_failed_response() logged_exception = mock_run_async.call_args.kwargs["exception"] - assert isinstance(logged_exception, litellm.APIError) + assert isinstance(logged_exception, litellm.BadRequestError) assert logged_exception.status_code == 400 +def test_handle_logging_failed_response_logs_content_policy_violation(): + """Failure logging must record the same typed exception the stream raises, so logging + integrations see a content policy violation instead of a generic APIError.""" + iterator = _make_iterator() + iterator.completed_response = _make_failed_chunk( + {"type": "invalid_request_error", "code": "content_policy_violation", "message": CONTENT_POLICY_MESSAGE} + ) + with ( + patch.object(import_module("litellm.responses.streaming_iterator"), "run_async_function") as mock_run_async, + patch.object(import_module("litellm.responses.streaming_iterator"), "executor"), + ): + iterator._handle_logging_failed_response() + logged_exception = mock_run_async.call_args.kwargs["exception"] + assert isinstance(logged_exception, litellm.ContentPolicyViolationError) + assert logged_exception.status_code == 400 + assert CONTENT_POLICY_MESSAGE in str(logged_exception) + + def test_handle_logging_failed_response_records_usage_and_cost(): """Usage on a response.failed event must reach failure spend accounting via combined_usage_object.""" iterator = _make_iterator() @@ -357,7 +484,7 @@ def test_sync_iterator_raises_mid_stream_fallback_on_rate_limit_error_event(): for _ in iterator: pass assert exc_info.value.status_code == 429 - assert isinstance(exc_info.value.original_exception, litellm.APIError) + assert isinstance(exc_info.value.original_exception, litellm.RateLimitError) def test_every_openai_sdk_response_error_code_has_explicit_status_mapping(): @@ -413,7 +540,7 @@ def test_maybe_raise_for_response_failed_event_maps_image_code_to_400(): chunk = Mock() chunk.type = "response.failed" chunk.response = mock_response_obj - with pytest.raises(litellm.APIError) as exc_info: + with pytest.raises(litellm.BadRequestError) as exc_info: iterator._maybe_raise_for_error_event(chunk) assert exc_info.value.status_code == 400 assert not isinstance(exc_info.value, MidStreamFallbackError) diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index dba44d1e2e8..cd72388ea21 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -9,7 +9,7 @@ import json import logging import sys import time -from collections.abc import AsyncIterator, Mapping +from collections.abc import AsyncIterator, Mapping, Sequence from copy import deepcopy from functools import partial from typing import Dict, Final, List, Literal @@ -31,6 +31,7 @@ from litellm.router_utils.auto_router_model_naming import ( ) from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import ( OUTPUT_TOKEN_CEILING_PARAMS, RETURN_RAW_MODEL_NAME_METADATA_KEY, @@ -70,6 +71,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import ( from litellm.types.router import ( Deployment, LiteLLM_Params, + PreRoutingHookResponse, RouterErrors, TaggedPreRoutingStrategy, ) @@ -5574,6 +5576,387 @@ class TestRoutingDecisionCauseLogging: assert "cause=semantic_keyword_match" not in router_log_capture.text +class TestTierModelAffinity: + @staticmethod + async def _route( + router: ComplexityRouter, + metadata: Mapping[str, object], + proposed_model: str, + prompt: str = "compact", + messages: list[dict[str, object]] | None = None, + ) -> PreRoutingHookResponse: + def choose(candidates: Sequence[str]) -> str: + return proposed_model if proposed_model in candidates else candidates[0] + + request_metadata: Final = dict(metadata) + with patch( # test-quality-ok: [TQ008] alternate proposals make affinity reuse deterministic + "litellm.router_strategy.complexity_router.complexity_router.random.choice", + side_effect=choose, + ): + result: Final = await router.async_pre_routing_hook( + model="affinity-router", + request_kwargs={"metadata": request_metadata}, + messages=messages if messages is not None else [{"role": "user", "content": prompt}], + ) + assert result is not None + if router.config.adaptive: + assert request_metadata["adaptive_router_chosen_model"] == result.model + return result + + @staticmethod + def _router( + mock_router_instance: MagicMock, + adaptive: bool = False, + deployment_affinity: bool = True, + plugins: bool = False, + ) -> ComplexityRouter: + mock_router_instance.cache = DualCache() + mock_router_instance.model_list = [] + mock_router_instance.model_name_to_deployment_indices = {} + return ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": { + tier: [ + {"model_name": model, "litellm_params": {"temperature": temperature}} + for model in ("model-a", "model-b") + ] + for tier, temperature in (("SIMPLE", 0.1), ("REASONING", 0.9)) + }, + "adaptive": adaptive, + "deployment_affinity": deployment_affinity, + "session_affinity": False, + **({"plugins": [_DummyPlugin()]} if plugins else {}), + }, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("adaptive", [False, True]) + async def test_reuses_model_per_tier_without_pinning_classification( + self, mock_router_instance: MagicMock, adaptive: bool + ) -> None: + router: Final = self._router(mock_router_instance, adaptive=adaptive) + metadata: Final = {"session_id": "same-session"} + first: Final = await self._route(router, metadata, "model-a") + if adaptive: + from litellm.router_strategy.adaptive_router.bandit import BanditCell + from litellm.router_strategy.adaptive_router.classifier import classify_prompt + + bandit: Final = router._ensure_adaptive_router() + assert bandit is not None + bandit._cells[(classify_prompt("compact"), "model-a")] = BanditCell(alpha=5.0, beta=5.0) + repeated: Final = await self._route(router, metadata, "model-b") + reasoning: Final = await self._route( + router, metadata, "model-b", "Let's think step by step and reason through this problem carefully." + ) + returned: Final = await self._route(router, metadata, "model-b") + + assert (first.model, repeated.model, reasoning.model, returned.model) == ( + "model-a", "model-a", "model-b", "model-a" + ) + assert tuple(result.routing_decision["tier"] for result in (first, repeated, reasoning, returned)) == ( + "SIMPLE", "SIMPLE", "REASONING", "SIMPLE" + ) + assert returned.litellm_params == {"temperature": 0.1} + assert reasoning.litellm_params == {"temperature": 0.9} + + @pytest.mark.asyncio + @pytest.mark.parametrize("identity_key", ["user_api_key_hash", "user_api_key_user_id"]) + async def test_isolates_sessions_and_authenticated_callers( + self, mock_router_instance: MagicMock, identity_key: str + ) -> None: + router: Final = self._router(mock_router_instance) + first_caller: Final = {"session_id": "shared", identity_key: "caller-a"} + other_caller: Final = {"session_id": "shared", identity_key: "caller-b"} + other_session: Final = {"session_id": "separate", identity_key: "caller-a"} + + assert (await self._route(router, first_caller, "model-a")).model == "model-a" + assert (await self._route(router, other_caller, "model-b")).model == "model-b" + assert (await self._route(router, other_session, "model-b")).model == "model-b" + assert (await self._route(router, first_caller, "model-b")).model == "model-a" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "metadata,deployment_affinity,plugins", + [ + ({}, True, False), + ({"session_id": "generated", SESSION_ID_GENERATED_METADATA_KEY: True}, True, False), + ({"session_id": "provided"}, False, False), + ({"session_id": "provided"}, True, True), + ], + ids=["absent-session", "generated-session", "disabled", "plugin-policy"], + ) + async def test_does_not_pin_without_eligible_session( + self, + mock_router_instance: MagicMock, + metadata: Mapping[str, object], + deployment_affinity: bool, + plugins: bool, + ) -> None: + router: Final = self._router( + mock_router_instance, deployment_affinity=deployment_affinity, plugins=plugins + ) + assert (await self._route(router, metadata, "model-a")).model == "model-a" + assert (await self._route(router, metadata, "model-b")).model == "model-b" + + @pytest.mark.asyncio + @pytest.mark.parametrize("adaptive", [False, True]) + async def test_replaces_pin_outside_the_context_candidate_domain(self, adaptive: bool) -> None: + router: Final = ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=_windowed_router(_SMALL, _BIG), + complexity_router_config={ + "tiers": {"SIMPLE": ["small-model", "big-model"]}, + "adaptive": adaptive, + "deployment_affinity": True, + "session_affinity": False, + }, + ) + metadata: Final = {"session_id": "growing-context"} + assert (await self._route(router, metadata, "small-model")).model == "small-model" + oversized: Final = await router.async_pre_routing_hook( + model="affinity-router", + request_kwargs={"metadata": dict(metadata)}, + messages=_OVERSIZED_TURNS, + ) + assert oversized is not None + assert oversized.model == "big-model" + assert oversized.routing_decision["tier"] == "SIMPLE" + assert (await self._route(router, metadata, "small-model")).model == "big-model" + + @pytest.mark.asyncio + @pytest.mark.parametrize("session_affinity", [False, True], ids=["user-turn", "session-affinity"]) + @pytest.mark.parametrize("gate", ["image", "health"]) + async def test_temporary_replay_gate_keeps_the_held_tiers_model_preference( + self, mock_router_instance: MagicMock, session_affinity: bool, gate: Literal["image", "health"] + ) -> None: + async def get_healthy_deployments( + model: str, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None = None, + input: object = None, + parent_otel_span: object = None, + health_check_probe: bool = False, + ) -> list[dict[str, object]]: + unavailable: Final = ( + gate == "health" + and model == "model-a" + and messages is not None + and bool(messages) + and messages[-1].get("role") == "tool" + ) + return [] if unavailable else [{"model_name": model, "model_info": {"id": f"deployment-{model}"}}] + + cache: Final = DualCache() + mock_router_instance.cache = cache + mock_router_instance.async_get_healthy_deployments = get_healthy_deployments + router: Final = TestModalityRouting._router( + mock_router_instance, + { + "tiers": {"SIMPLE": ["model-a", "model-b"]}, + "deployment_affinity": True, + "session_affinity": session_affinity, + "classification_mode": "every_request" if session_affinity else "user_turn", + "modality_routing": True, + "modality_pin_override": True, + }, + {"model-a": False, "model-b": True}, + ) + metadata: Final = {"session_id": "replay-session"} + continuation: Final[list[dict[str, object]]] = [ + {"role": "user", "content": "compact"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [IMG_PART] if gate == "image" else "done"}, + ] + assert (await self._route(router, metadata, "model-a")).model == "model-a" + + replayed: Final = await self._route(router, metadata, "model-b", messages=continuation) + assert replayed.model == "model-b" + assert replayed.routing_decision["tier"] == "SIMPLE" + assert replayed.routing_decision["cause"] == ( + "health_failover" + if gate == "health" + else ("modality_pin_override" if session_affinity else "user_turn_continuation") + ) + cache_key: Final = router._get_session_affinity_cache_key("replay-session", {"metadata": metadata}) + assert await cache.async_get_cache(cache_key) == {"model": "model-a", "tier": "SIMPLE"} + + next_ask: Final = await self._route(router, metadata, "model-b") + assert next_ask.model == "model-a" + assert next_ask.routing_decision["tier"] == "SIMPLE" + assert next_ask.routing_decision["cause"] == ( + "session_affinity_pin" if session_affinity else "heuristic_scorer" + ) + + @pytest.mark.asyncio + async def test_user_turn_replay_refreshes_the_model_used_within_its_tier( + self, mock_router_instance: MagicMock + ) -> None: + clock: Final = MagicMock(return_value=100.0) + mock_router_instance.cache = DualCache(in_memory_cache=InMemoryCache(clock=clock)) + router: Final = ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": {"SIMPLE": ["model-a", "model-b"]}, + "classification_mode": "user_turn", + "session_affinity_ttl_seconds": 10, + }, + ) + metadata: Final = {"session_id": "same-session"} + continuation: Final[list[dict[str, object]]] = [ + {"role": "user", "content": "compact"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "done"}, + ] + assert (await self._route(router, metadata, "model-a")).model == "model-a" + clock.return_value = 105.0 + replayed: Final = await self._route(router, metadata, "model-b", messages=continuation) + assert replayed.model == "model-a" + assert replayed.routing_decision["cause"] == "user_turn_continuation" + + clock.return_value = 111.0 + next_ask: Final = await self._route(router, metadata, "model-b") + assert next_ask.model == "model-a" + assert next_ask.routing_decision["tier"] == "SIMPLE" + assert next_ask.routing_decision["cause"] == "heuristic_scorer" + + @pytest.mark.asyncio + async def test_session_escalation_keeps_the_selected_tier_when_models_overlap( + self, mock_router_instance: MagicMock + ) -> None: + cache: Final = DualCache() + mock_router_instance.cache = cache + router: Final = ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": { + "SIMPLE": "base", + **{ + tier: [ + {"model_name": model, "litellm_params": {"temperature": temperature}} + for model in models + ] + for tier, models, temperature in ( + ("MEDIUM", ("shared", "middle"), 0.4), + ("COMPLEX", ("shared", "higher"), 0.8), + ) + }, + }, + "session_affinity": True, + "keyword_tier_rules": [{"keywords": ["visit_complex"], "tier": "COMPLEX"}], + }, + ) + metadata: Final = {"session_id": "same-session"} + assert (await self._route(router, metadata, "higher", "visit_complex")).model == "higher" + cache_key: Final = router._get_session_affinity_cache_key("same-session", {"metadata": metadata}) + await cache.async_set_cache(cache_key, {"model": "base", "tier": "SIMPLE"}, ttl=600) + + result: Final = await self._route(router, metadata, "shared", "LITELLM ESCALATE") + assert result.model == "shared" + assert result.routing_decision["tier"] == "MEDIUM" + assert result.routing_decision["cause"] == "session_affinity_escalation" + assert result.litellm_params == {"temperature": 0.4} + assert await cache.async_get_cache(cache_key) == {"model": "shared", "tier": "MEDIUM"} + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "stale_tier", + ["NON_REASONING", "REMOVED_TIER", 7, []], + ids=["inactive-tier", "unknown-tier", "integer-tier", "list-tier"], + ) + @pytest.mark.parametrize( + "prompt,expected_model,expected_tier", + [("compact", "model-a", "SIMPLE"), ("LITELLM ESCALATE", "model-b", "MEDIUM")], + ids=["ordinary-replay", "escalation"], + ) + async def test_reclassifies_session_pin_outside_the_active_tier_ladder( + self, + mock_router_instance: MagicMock, + stale_tier: object, + prompt: str, + expected_model: str, + expected_tier: str, + ) -> None: + cache: Final = DualCache() + mock_router_instance.cache = cache + router: Final = ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": {"SIMPLE": "model-a", "MEDIUM": "model-b"}, + "session_affinity": True, + }, + ) + metadata: Final = {"session_id": "same-session"} + cache_key: Final = router._get_session_affinity_cache_key("same-session", {"metadata": metadata}) + await cache.async_set_cache(cache_key, {"model": "model-a", "tier": stale_tier}, ttl=600) + + result: Final = await self._route(router, metadata, expected_model, prompt) + + assert result.model == expected_model + assert result.routing_decision["tier"] == expected_tier + assert result.routing_decision["cause"] == "heuristic_scorer" + assert await cache.async_get_cache(cache_key) == {"model": expected_model, "tier": expected_tier} + + @pytest.mark.asyncio + @pytest.mark.parametrize("classification_mode", ["every_request", "user_turn"]) + async def test_custom_tier_keeps_its_own_model( + self, mock_router_instance: MagicMock, classification_mode: Literal["every_request", "user_turn"] + ) -> None: + mock_router_instance.cache = DualCache() + router: Final = ComplexityRouter( + model_name="affinity-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=_custom_tier_config( + tiers={"SIMPLE": ["model-a", "model-b"], "SECURITY_REVIEW": ["model-a", "model-b"], "COMPLEX": "model-a"}, + deployment_affinity=True, + classification_mode=classification_mode, + keyword_tier_rules=[ + {"keywords": ["compact"], "tier": "SIMPLE"}, + {"keywords": ["audit"], "tier": "SECURITY_REVIEW"}, + ], + ), + ) + metadata: Final = {"session_id": "custom-session"} + assert (await self._route(router, metadata, "model-a")).model == "model-a" + assert (await self._route(router, metadata, "model-b", "audit")).model == "model-b" + assert (await self._route(router, metadata, "model-b")).model == "model-a" + retained: Final = await self._route(router, metadata, "model-a", "audit") + assert retained.model == "model-b" + assert retained.routing_decision["tier"] == "SECURITY_REVIEW" + if classification_mode == "user_turn": + continuation: Final[list[dict[str, object]]] = [ + {"role": "user", "content": "audit"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "done"}, + ] + replayed: Final = await self._route(router, metadata, "model-a", messages=continuation) + assert replayed.model == "model-b" + assert replayed.routing_decision["tier"] == "SECURITY_REVIEW" + assert replayed.routing_decision["cause"] == "user_turn_continuation" + + class TestSessionAffinity: """Test the session_affinity sticky-routing behavior (off by default).""" @@ -5638,11 +6021,8 @@ class TestSessionAffinity: tier_pinned, deployment_pinned, ): - """deployment_affinity pins the deployment inside each routed group without pinning which - group the session routes to, so with session_affinity off the tier must still reclassify - on every turn while the marker the Router stamps is still emitted. Turn 1 classifies - REASONING and turn 2 SIMPLE, so a reclassified turn 2 moves model while a tier-pinned one - does not. plugins suppress both pins, since a stale pin would bypass the plugin pipeline.""" + """Deployment affinity retains a model per tier while classification continues. + Session affinity keeps the first tier too; plugins suppress both affinity policies.""" mock_router_instance.cache = DualCache() router = ComplexityRouter( model_name="test-router", @@ -5692,8 +6072,7 @@ class TestSessionAffinity: @pytest.mark.asyncio async def test_disabled_by_default_reclassifies_every_turn(self, mock_router_instance, basic_config): - """Regression: session_affinity defaults to False, so a shared session_id must NOT - pin the first turn's model; every turn is classified on its own merits.""" + """With session_affinity off, a shared session can move from REASONING to SIMPLE.""" assert "session_affinity" not in basic_config mock_router_instance.cache = DualCache() router = ComplexityRouter( @@ -5848,7 +6227,7 @@ class TestSessionAffinity: @pytest.mark.asyncio async def test_respects_ttl_seconds(self, mock_router_instance, basic_config): - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value=None) mock_router_instance.cache = cache router = ComplexityRouter( @@ -5872,7 +6251,7 @@ class TestSessionAffinity: async def test_ttl_refreshed_on_cache_hit(self, mock_router_instance, basic_config): """Regression: a pinned turn must refresh the TTL, not just the first write -- otherwise a session outliving session_affinity_ttl_seconds silently loses its pin.""" - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value="o1-preview") mock_router_instance.cache = cache router = ComplexityRouter( @@ -7112,7 +7491,8 @@ class TestEscalationKeywords: complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": ["o1-a", "o1-b", "o1-c"]}}, ) for pinned in ("o1-a", "o1-b", "o1-c"): - assert router._escalated_pin(pinned) == pinned + escalated: Final = router._escalated_pin(pinned) + assert (escalated.model, escalated.tier) == (pinned, "REASONING") @pytest.mark.asyncio async def test_session_escalation_at_ceiling_keeps_multi_model_pin(self, mock_router_instance): @@ -12009,7 +12389,7 @@ async def test_session_pin_uses_recorded_tier_when_model_is_in_multiple_tiers(mo @pytest.mark.asyncio async def test_session_pin_survives_json_list_round_trip(mock_router_instance): - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value=["shared", "SIMPLE"]) mock_router_instance.cache = cache router = ComplexityRouter( @@ -12988,7 +13368,7 @@ class TestModalityRouting: {"role": "user", "content": [{"type": "text", "text": "quick lookup: what is this?"}, IMG_PART]} ] elif path.startswith(("pin_kept", "pin_replacement", "pin_override")): - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value={"model": "text-cheap", "tier": "SIMPLE"}) mock_router_instance.cache = cache config["session_affinity"] = True @@ -13178,7 +13558,7 @@ class TestModalityRouting: @pytest.mark.asyncio async def test_pin_override_serves_the_image_turn_without_repinning(self, mock_router_instance): """The override is for one request: the session keeps the model it was pinned to.""" - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value={"model": "text-cheap", "tier": "SIMPLE"}) mock_router_instance.cache = cache router = self._router( @@ -13211,7 +13591,7 @@ class TestModalityRouting: @pytest.mark.asyncio async def test_pin_override_with_no_capable_model_rejects_and_keeps_the_pin(self, mock_router_instance): """The clear 400 replaces the provider's, and a rejected turn must not cost the session its pin.""" - cache = AsyncMock() + cache: Final = AsyncMock(in_memory_cache=DualCache().in_memory_cache, redis_cache=None) cache.async_get_cache = AsyncMock(return_value={"model": "text-cheap", "tier": "SIMPLE"}) mock_router_instance.cache = cache router = self._router( @@ -14322,18 +14702,37 @@ class TestTierHealthFailover: cooling=("id-a1",), raises_for={"exhausted-b": raised}, ) - key = router._get_session_affinity_cache_key("sess-exhausted", {}) - await router.litellm_router_instance.cache.async_set_cache( - key=key, value={"model": "dead-a", "tier": "SIMPLE"}, ttl=600 - ) - results = [ - await router.async_pre_routing_hook( - model="m", request_kwargs={"metadata": {"session_id": "sess-exhausted"}}, messages=self.SIMPLE_MESSAGE + sessions: Final = tuple(f"sess-exhausted-{sample}" for sample in range(20)) + await asyncio.gather( + *( + router.litellm_router_instance.cache.async_set_cache( + key=router._get_session_affinity_cache_key(session_id, {}), + value={"model": "dead-a", "tier": "SIMPLE"}, + ttl=600, + ) + for session_id in sessions ) - for _ in range(20) + ) + results: Final = [ + await router.async_pre_routing_hook( + model="m", request_kwargs={"metadata": {"session_id": session_id}}, messages=self.SIMPLE_MESSAGE + ) + for session_id in sessions ] assert {r.model for r in results} == expected + def choose_other(candidates: Sequence[str]) -> str: + return next((model for model in candidates if model != results[0].model), candidates[0]) + + with patch( # test-quality-ok: [TQ008] an alternate healthy proposal proves retained affinity across failover + "litellm.router_strategy.complexity_router.complexity_router.random.choice", + side_effect=choose_other, + ): + retained: Final = await router.async_pre_routing_hook( + model="m", request_kwargs={"metadata": {"session_id": sessions[0]}}, messages=self.SIMPLE_MESSAGE + ) + assert retained.model == results[0].model + @pytest.mark.asyncio async def test_a_group_the_router_has_no_deployment_for_is_not_a_failover_target(self, mock_router_instance): """The owner answers an unconfigured group with BadRequestError. Reading that as live diff --git a/tests/test_litellm/router_strategy/test_simple_shuffle.py b/tests/test_litellm/router_strategy/test_simple_shuffle.py index 165c1751f63..abf02860a50 100644 --- a/tests/test_litellm/router_strategy/test_simple_shuffle.py +++ b/tests/test_litellm/router_strategy/test_simple_shuffle.py @@ -1,4 +1,5 @@ from collections import Counter +from inspect import isawaitable import pytest @@ -52,3 +53,52 @@ async def test_uniform_pick_when_every_configured_weight_is_zero(): assert counts["unweighted"] > 0 assert counts["standby"] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selector", [ + "get_available_deployment", "async_get_available_deployment", + "get_available_deployment_for_pass_through", "async_get_available_deployment_for_pass_through", +]) +async def test_scoped_weights_are_request_local_and_respect_eligibility(selector: str) -> None: + router = Router(model_list=[ + { + **_deployment(deployment_id, { + "weight": 100 if deployment_id == "global" else 0, "use_in_pass_through": True, + }), + "model_name": f"model_name_{team_id}_{deployment_id}", + "model_info": { + "id": deployment_id, "team_id": team_id, "team_public_model_name": "test-model", "blocked": blocked, + }, + } + for deployment_id, team_id, blocked in ( + ("global", "team-a", False), ("scoped", "team-a", False), + ("blocked", "team-a", True), ("foreign", "other-team", False), + ) + ], num_retries=0) + + for weights, expected in ( + ({"test-model": {"global": 0, "scoped": 100, "blocked": 100, "foreign": 100}}, "scoped"), + ({"test-model": {"global": 100, "scoped": 0}}, "global"), + ({"test-model": {"foreign": 100}}, "global"), + ({"test-model": {"blocked": 100}}, "global"), + (None, "global"), + ): + result = getattr(router, selector)( + model="test-model", + request_kwargs={"metadata": {"user_api_key_team_id": "team-a"}, "_router_weights": weights}, + ) + deployment = await result if isawaitable(result) else result + assert deployment["model_info"]["id"] == expected + + +def test_scoped_weights_approximate_the_configured_split() -> None: + router = Router(model_list=[_deployment("primary"), _deployment("secondary")], num_retries=0) + counts = Counter( + router.get_available_deployment( + model="test-model", + request_kwargs={"_router_weights": {"test-model": {"primary": 80, "secondary": 20}}}, + )["model_info"]["id"] + for _ in range(1000) + ) + assert 700 < counts["primary"] < 900 diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py index ea8e2eacaa6..b93b8c1cdfc 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py @@ -21,6 +21,7 @@ from unittest.mock import AsyncMock, patch import pytest import litellm +from litellm.models.credentials import CredentialItem from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ResponsesAPIResponse @@ -1082,6 +1083,173 @@ def test_boundary_key_accepts_pydantic_litellm_params_instance(): ) +def test_boundary_key_resolves_missing_values_from_named_credential(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + with ( + patch.object( # test-quality-ok: credential registry is the direct dependency under test + litellm, + "credential_list", + [ + CredentialItem( + credential_name="account-a", + credential_values={ + "api_base": "https://account-a.example.com", + "api_key": "credential-key-a", + }, + credential_info={}, + ) + ], + ) + ): + boundary = EncryptedContentAffinityCheck._encryption_boundary_key({"litellm_credential_name": "account-a"}) + + assert boundary == ("https://account-a.example.com", "credential-key-a") + + +def test_boundary_key_matches_named_credential_precedence(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + with ( + patch.object( # test-quality-ok: credential registry is the direct dependency under test + litellm, + "credential_list", + [ + CredentialItem( + credential_name="account-a", + credential_values={ + "api_base": "https://credential.example.com", + "api_key": "credential-key-a", + }, + credential_info={}, + ) + ], + ) + ): + boundary = EncryptedContentAffinityCheck._encryption_boundary_key( + { + "api_base": "https://deployment.example.com", + "api_key": "deployment-key", + "litellm_credential_name": "account-a", + } + ) + + assert boundary == ("https://credential.example.com", "credential-key-a") + + +def test_boundary_key_resolves_credential_when_explicit_values_are_empty(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + with ( + patch.object( # test-quality-ok: credential registry is the direct dependency under test + litellm, + "credential_list", + [ + CredentialItem( + credential_name="account-a", + credential_values={ + "api_base": "https://credential.example.com", + "api_key": "credential-key-a", + }, + credential_info={}, + ) + ], + ) + ): + boundary = EncryptedContentAffinityCheck._encryption_boundary_key( + { + "api_base": "", + "api_key": "", + "litellm_credential_name": "account-a", + } + ) + + assert boundary == ("https://credential.example.com", "credential-key-a") + + +def test_boundary_fallback_matches_deployments_with_same_named_credential_values(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + with ( + patch.object( # test-quality-ok: credential registry is the direct dependency under test + litellm, + "credential_list", + [ + CredentialItem( + credential_name="account-a", + credential_values={ + "api_base": "https://account-a.example.com", + "api_key": "credential-key-a", + }, + credential_info={}, + ), + CredentialItem( + credential_name="account-a-peer", + credential_values={ + "api_base": "https://account-a.example.com", + "api_key": "credential-key-a", + }, + credential_info={}, + ), + CredentialItem( + credential_name="account-b", + credential_values={ + "api_base": "https://account-b.example.com", + "api_key": "credential-key-b", + }, + credential_info={}, + ), + ], + ) + ): + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5.3-codex", + "litellm_params": { + "model": "azure/gpt-5.3-codex", + "litellm_credential_name": "account-a", + }, + "model_info": {"id": "origin"}, + } + ], + num_retries=0, + ) + check = EncryptedContentAffinityCheck(router=router) + healthy_deployments = [ + { + "model_info": {"id": "peer-same-boundary"}, + "litellm_params": { + "model": "azure/gpt-5.4", + "litellm_credential_name": "account-a-peer", + }, + }, + { + "model_info": {"id": "peer-different-boundary"}, + "litellm_params": { + "model": "azure/gpt-5.4", + "litellm_credential_name": "account-b", + }, + }, + ] + + matches, originating = check._find_deployments_on_same_encryption_boundary( + healthy_deployments=healthy_deployments, + model_id="origin", + ) + + assert originating is not None + assert [deployment["model_info"]["id"] for deployment in matches] == ["peer-same-boundary"] + + def test_boundary_key_rejects_non_dict_like_inputs(): """ Inputs that don't expose ``.get()`` (None, lists, strings, ints) -> None. diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py index cf48888600e..780300bf9e1 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_session_id_affinity.py @@ -1,3 +1,5 @@ +import asyncio +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -6,7 +8,9 @@ import pytest import json import litellm +from litellm.caching.affinity_cache import claim_affinity_pin from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( DeploymentAffinityCheck, @@ -558,6 +562,124 @@ async def test_claim_pin_falls_back_to_pod_local_when_redis_is_down(): assert second == "our-deployment" +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("stored", "expected"), + [ + ({"model": "first"}, {"model": "first"}), + ('{ "model" : "first" }', {"model": "first"}), + ({"model": "removed"}, {"model": "second"}), + ({"model": "first", "extra": "stale"}, {"model": "second"}), + ({"model_id": "first"}, {"model": "second"}), + ("first", {"model": "second"}), + (None, {"model": "second"}), + ], +) +async def test_eligible_affinity_claim_replaces_stale_pins_and_slides_ttl( + stored: object, expected: object +) -> None: + clock: Final = MagicMock(return_value=100.0) + cache: Final = DualCache(in_memory_cache=InMemoryCache(clock=clock)) + cache.in_memory_cache.set_cache("tier-pin", stored, ttl=10) + clock.return_value = 105.0 + + winner: Final = await claim_affinity_pin( + cache, "tier-pin", {"model": "second"}, 30, + eligible_values=({"model": "first"}, {"model": "second"}), + ) + + assert winner == expected + assert cache.in_memory_cache.ttl_dict["tier-pin"] == 135.0 + clock.return_value = 111.0 + assert cache.in_memory_cache.get_cache("tier-pin") == expected + clock.return_value = 136.0 + assert cache.in_memory_cache.get_cache("tier-pin") is None + + +@pytest.mark.asyncio +async def test_concurrent_eligible_claims_return_one_winner() -> None: + cache: Final = DualCache() + candidates: Final = ({"model": "first"}, {"model": "second"}) + winners: Final = await asyncio.gather(*( + claim_affinity_pin( + cache, "tier-pin", candidates[index % 2], 30, + eligible_values=candidates, + ) + for index in range(20) + )) + + assert winners == [{"model": "first"}] * 20 + assert cache.in_memory_cache.get_cache("tier-pin") == {"model": "first"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("stored", "expected", "refresh"), + [ + ({"model_id": 7}, "7", True), + ({"model_id": "other"}, "other", False), + ({"model": "7"}, None, False), + (["7"], None, False), + ], +) +async def test_legacy_deployment_claim_retains_decoder_and_keepalive( + stored: object, expected: str | None, refresh: bool +) -> None: + clock: Final = MagicMock(return_value=100.0) + cache: Final = DualCache(in_memory_cache=InMemoryCache(clock=clock)) + callback: Final = DeploymentAffinityCheck( + cache=cache, ttl_seconds=30, + enable_user_key_affinity=False, enable_responses_api_affinity=False, + ) + cache.in_memory_cache.set_cache("deployment-pin", stored, ttl=10) + clock.return_value = 105.0 + + winner: Final = await callback._claim_pin( + "deployment-pin", {"model_id": "7"}, 30 + ) + + assert winner == expected + assert cache.in_memory_cache.ttl_dict["deployment-pin"] == ( + 135.0 if refresh else 110.0 + ) + assert cache.in_memory_cache.get_cache("deployment-pin") == ( + {"model_id": "7"} if refresh else stored + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("raw", "expected", "stored"), + [ + (b'{"model_id": "winner"}', "winner", {"model_id": "winner"}), + ('"winner"', "winner", "winner"), + ("winner", "winner", "winner"), + (b"winner", "winner", "winner"), + ('{"model": "winner"}', None, {"model": "winner"}), + (None, "candidate", None), + (123, "candidate", None), + ({"model_id": "winner"}, "candidate", None), + ], +) +async def test_redis_deployment_claim_preserves_legacy_result_decoding( + raw: object, expected: str | None, stored: object +) -> None: + redis: Final = MagicMock() + redis.async_register_script.return_value = AsyncMock(return_value=raw) + cache: Final = DualCache(redis_cache=redis) + callback: Final = DeploymentAffinityCheck( + cache=cache, ttl_seconds=30, + enable_user_key_affinity=False, enable_responses_api_affinity=False, + ) + + winner: Final = await callback._claim_pin( + "deployment-pin", {"model_id": "candidate"}, 30 + ) + + assert winner == expected + assert cache.in_memory_cache.get_cache("deployment-pin") == stored + + @pytest.mark.asyncio async def test_marker_session_affinity_read_and_write_agree_for_wildcard_groups(): """Wildcard deployments keep the literal pattern as model_name on both the read diff --git a/tests/test_litellm/rust_bridge/test_lifecycle.py b/tests/test_litellm/rust_bridge/test_lifecycle.py index 1f0b5591c2b..d73385621d5 100644 --- a/tests/test_litellm/rust_bridge/test_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_lifecycle.py @@ -8,7 +8,7 @@ from litellm.rust_bridge.lifecycle import check_limits @pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) @pytest.mark.parametrize( - "cap, attempted_retries, refused", + "cap, request_retry_count, refused", [(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)], ids=[ "cap-above-four-reached", @@ -17,12 +17,15 @@ from litellm.rust_bridge.lifecycle import check_limits "cap-of-zero-refuses-first-retry", ], ) -def test_check_limits_reads_attempted_retries( - monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, attempted_retries: int, refused: bool +def test_check_limits_reads_request_retry_count( + monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool ) -> None: monkeypatch.setattr(litellm, "num_retries_per_request", cap) monkeypatch.setattr(litellm, "max_budget", None) - kwargs: Final = {"model": "mistral/mistral-ocr-latest", metadata_key: {"attempted_retries": attempted_retries}} + kwargs: Final = { + "model": "mistral/mistral-ocr-latest", + metadata_key: {"request_retry_count": request_retry_count}, + } if refused: with pytest.raises(RuntimeError, match="Max retries per request hit!"): check_limits(kwargs) diff --git a/tests/test_litellm/test_assert_ci_coverage.py b/tests/test_litellm/test_assert_ci_coverage.py index d948c1a4155..69db2411742 100644 --- a/tests/test_litellm/test_assert_ci_coverage.py +++ b/tests/test_litellm/test_assert_ci_coverage.py @@ -9,8 +9,12 @@ the question neither covers: whether the job that globs a file then deselects it """ import importlib.util +import json import sys from pathlib import Path +from typing import Final + +import yaml _REPO_ROOT = Path(__file__).resolve().parents[2] _MODULE_PATH = _REPO_ROOT / ".github" / "scripts" / "assert_ci_coverage.py" @@ -20,6 +24,56 @@ sys.modules[_spec.name] = coverage # @dataclass(slots=True) rebuilds via sys.mo _spec.loader.exec_module(coverage) +def test_integration_manifest_requires_exclusive_scheduled_circleci_owner(tmp_path: Path) -> None: + test_path: Final = "tests/integration/management/test_contract.py" + test_file: Final = tmp_path / test_path + test_file.parent.mkdir(parents=True) + test_file.write_text("def test_contract(): pass\n") + (tmp_path / "tests/integration/contracts.json").write_text( + json.dumps({"groups": {"management": ["management"]}, "tests": {f"{test_path}::test_contract": ["mgmt.test"]}}) + ) + paths, findings = coverage._integration_ownership(tmp_path) + assert not paths + assert [finding.detail for finding in findings] == ["dedicated CircleCI runner is missing"] + circle: Final = tmp_path / ".circleci/config.yml" + circle.parent.mkdir() + circle.write_text( + yaml.safe_dump( + { + "jobs": { + "integration_contracts": { + "steps": [{"run": {"command": "bash .circleci/scripts/run_integration.sh management"}}] + } + }, + "workflows": {"integration": {"jobs": [{"integration_contracts": {"suite": "management"}}]}}, + } + ) + ) + paths, findings = coverage._integration_ownership(tmp_path) + assert paths == frozenset({test_path}) + assert findings == () + configured: Final = yaml.safe_load(circle.read_text()) + configured["workflows"]["integration"]["jobs"] = [ + {"integration_contracts": {"matrix": {"parameters": {"suite": ["providers"]}}}} + ] + circle.write_text(yaml.safe_dump(configured)) + _, findings = coverage._integration_ownership(tmp_path) + assert [(finding.subject, finding.detail) for finding in findings] == [ + ("management", "canonical integration group is not scheduled by CircleCI") + ] + configured["workflows"]["integration"]["jobs"][0]["integration_contracts"]["matrix"]["parameters"]["suite"] = [ + "management" + ] + circle.write_text(yaml.safe_dump(configured)) + workflow: Final = tmp_path / ".github/workflows/test.yml" + workflow.parent.mkdir(parents=True) + workflow.write_text(yaml.safe_dump({"jobs": {"tests": {"steps": [{"run": "pytest tests/integration"}]}}})) + _, findings = coverage._integration_ownership(tmp_path) + assert [(finding.subject, finding.detail) for finding in findings] == [ + (test_path, "integration contract is also selected by GitHub Actions") + ] + + def test_an_ancestor_directory_covers_a_file_but_does_not_name_it(): # The whole point of the split: `tests/x` answers "does it run?" but not # "which shard owns it?" — accepting it for the latter is how a new child diff --git a/tests/test_litellm/test_bedrock_usgov_pricing.py b/tests/test_litellm/test_bedrock_usgov_pricing.py index 3dfd7350a06..4d5b27a8668 100644 --- a/tests/test_litellm/test_bedrock_usgov_pricing.py +++ b/tests/test_litellm/test_bedrock_usgov_pricing.py @@ -112,6 +112,11 @@ GOV_ROW_SOURCES = { } +BEDROCK_PRICE_LIST_URL = ( + "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" +) + + def _non_pricing_fields(info): return {k: v for k, v in info.items() if "cost" not in k and k not in ("litellm_provider", "source")} @@ -121,8 +126,10 @@ def test_usgov_rows_keep_commercial_limits_and_capabilities(model_data, gov_key) """A gov row differs from the commercial row it mirrors only in price and provider: context limits, mode, and capability flags stay identical, so a hand-copied row cannot silently drop tool calling or shrink the context window. + The only source a gov row may cite is the AWS price list, which prices the + us-gov regions itself; a commercial doc URL copied along with the row is not. """ gov = model_data[gov_key] assert _non_pricing_fields(gov) == _non_pricing_fields(model_data[GOV_ROW_SOURCES[gov_key]]) assert "search_context_cost_per_query" not in gov - assert "source" not in gov + assert gov.get("source", BEDROCK_PRICE_LIST_URL) == BEDROCK_PRICE_LIST_URL diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 68e9b6143a0..8c3436d3108 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -3909,6 +3909,57 @@ def _batch_cache_usage() -> Usage: ) +def test_batch_cost_calculator_prices_multimodal_tokens_at_modality_rates(): + from litellm.cost_calculator import batch_cost_calculator + + model_info: ModelInfo = { + "input_cost_per_token_batches": 1e-7, + "input_cost_per_audio_token_batches": 3.25e-6, + "input_cost_per_image_token_batches": 2.25e-7, + "input_cost_per_video_token_batches": 6e-6, + } + usage = Usage( + prompt_tokens=100, + completion_tokens=0, + total_tokens=100, + prompt_tokens_details=PromptTokensDetailsWrapper( + audio_tokens=64, + image_tokens=10, + video_tokens=6, + ), + ) + + prompt_cost, _ = batch_cost_calculator( + usage=usage, + model="gemini-embedding-2", + custom_llm_provider="vertex_ai", + model_info=model_info, + ) + + assert prompt_cost == pytest.approx(20 * 1e-7 + 64 * 3.25e-6 + 10 * 2.25e-7 + 6 * 6e-6) + + +def test_batch_cost_calculator_falls_back_to_text_batch_rate_for_modalities(): + from litellm.cost_calculator import batch_cost_calculator + + model_info: ModelInfo = {"input_cost_per_token_batches": 1e-7} + usage = Usage( + prompt_tokens=100, + completion_tokens=0, + total_tokens=100, + prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=64), + ) + + prompt_cost, _ = batch_cost_calculator( + usage=usage, + model="gemini-embedding-2", + custom_llm_provider="vertex_ai", + model_info=model_info, + ) + + assert prompt_cost == pytest.approx(100 * 1e-7) + + def test_batch_cost_calculator_prices_cache_creation_tokens_at_cache_write_rate(): """ LIT-4008 regression: anthropic batch usage is dominated by cache tokens. diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 4f7a51eb531..3dccb2b35bf 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -3795,3 +3795,58 @@ def test_azure_ai_speech_on_a_foundry_host_uses_the_azure_openai_deployment_rout 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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index f1b445fb1bd..0485da3eba9 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3694,6 +3694,111 @@ async def test_aresponses_streaming_iterator_fallback(): assert call_kwargs["disable_fallbacks"] is False +@pytest.mark.asyncio +async def test_aresponses_streaming_content_policy_error_event_routes_to_content_policy_fallback(): + """Regression: a mid-stream content_policy_violation error event never reached + content_policy_fallbacks. The iterator raised a bare APIError the wrapper does not + catch, and even once wrapped, the MidStreamFallbackError envelope was handed to the + fallback dispatch, whose isinstance branch on ContentPolicyViolationError never matched. + The stream below is the customer's shape: a raw OpenAI error event with code + content_policy_violation, transformed by the real OpenAI config, and the router must + call the content_policy_fallbacks target, not the general fallbacks one.""" + from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + + router = litellm.Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/gpt-5.4", "api_key": "k1"}}, + { + "model_name": "content-fallback", + "litellm_params": {"model": "gemini/gemini-2.5-flash", "api_key": "k2"}, + }, + {"model_name": "general-fallback", "litellm_params": {"model": "openai/gpt-5-mini", "api_key": "k3"}}, + ], + fallbacks=[{"primary": ["general-fallback"]}], + content_policy_fallbacks=[{"primary": ["content-fallback"]}], + ) + error_event = { + "type": "error", + "sequence_number": 2, + "error": { + "type": "invalid_request_error", + "code": "content_policy_violation", + "message": "This content was flagged for possible cybersecurity risk. The response was halted mid-stream.", + "param": None, + }, + } + + async def aiter_bytes(): + yield f"data: {json.dumps(error_event)}\n\n".encode() + + raw_response = MagicMock() + raw_response.headers = {} + raw_response.aiter_bytes = aiter_bytes + logging_obj = MagicMock(spec=LiteLLMLogging) + logging_obj.model_call_details = {"litellm_params": {}} + logging_obj.completion_start_time = None + source = ResponsesAPIStreamingIterator( + response=raw_response, + model="gpt-5.4", + responses_api_provider_config=OpenAIResponsesAPIConfig(), + logging_obj=logging_obj, + custom_llm_provider="openai", + ) + fallback_chunks = [MagicMock(type="response.output_text.delta"), MagicMock(type="response.completed")] + fallback_call = AsyncMock(return_value=_AsyncList(fallback_chunks)) + + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={ + "model": "primary", + "stream": True, + "input": "Hi", + "original_generic_function": fallback_call, + }, + ) + collected = [chunk async for chunk in wrapped] + + assert collected == fallback_chunks + fallback_call.assert_awaited_once() + assert fallback_call.await_args.kwargs["model"] == "gemini/gemini-2.5-flash" + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_unwraps_content_policy_trigger_for_fallback_dispatch(): + """The fallback dispatch matches on the trigger's own type, so the wrapper must hand it the + ContentPolicyViolationError carried inside MidStreamFallbackError, not the envelope.""" + router = _make_router_with_fallback("openai/gpt-5.4", "openai/gpt-5-mini") + content_policy_error = litellm.ContentPolicyViolationError( + message="flagged mid-stream", llm_provider="openai", model="openai/gpt-5.4" + ) + src = _make_responses_iterator( + chunks=[MagicMock(type="response.created")], + error=MidStreamFallbackError( + message=str(content_policy_error), + model="openai/gpt-5.4", + llm_provider="openai", + original_exception=content_policy_error, + is_pre_first_chunk=True, + ), + model="openai/gpt-5.4", + ) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=_AsyncList([MagicMock(type="response.completed")])), + ) as mock_fallback_utils: + wrapped = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={"model": "openai/gpt-5.4", "stream": True, "input": "Hi"}, + ) + [chunk async for chunk in wrapped] + + mock_fallback_utils.assert_awaited_once() + assert mock_fallback_utils.await_args.kwargs["e"] is content_policy_error + + @pytest.mark.asyncio @pytest.mark.parametrize( "fallback_headers", @@ -10961,6 +11066,66 @@ async def test_num_retries_per_request_stops_retries_at_caps_above_four(monkeypa ] +def _failing_group_with_healthy_fallback_router(num_retries: int) -> litellm.Router: + return litellm.Router( + model_list=[ + { + "model_name": "broken-group", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-fake", + "mock_response": "litellm.InternalServerError", + }, + }, + { + "model_name": "healthy-group", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake", "mock_response": "ok"}, + }, + ], + fallbacks=[{"broken-group": ["healthy-group"]}], + num_retries=num_retries, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "cap, planted_count, hop_refused", + [(2, None, True), (4, None, False), (2, -100, True)], + ids=["cap-spent-before-the-hop", "cap-not-reached-by-the-hop", "planted-negative-count-does-not-lift-the-cap"], +) +async def test_num_retries_per_request_counts_retries_across_fallback_hops( + monkeypatch: pytest.MonkeyPatch, cap: int, planted_count: int | None, hop_refused: bool +) -> None: + """num_retries_per_request caps the retries of one request, fallback hops included. Each hop starts a + fresh per-hop attempted_retries at zero, so a cap read from that counter let every hop retry from zero + and a request could spend far more retries than the cap allows. A caller who plants a negative count + in the request metadata must not push the cap further away either.""" + monkeypatch.setattr(litellm, "num_retries_per_request", cap) + router = _failing_group_with_healthy_fallback_router(num_retries=1) + recorder = _FallbackAttemptRecorder() + litellm.callbacks.append(recorder) + try: + metadata = {} if planted_count is None else {"request_retry_count": planted_count} + request = router.acompletion( + model="broken-group", messages=[{"role": "user", "content": "hi"}], metadata=metadata + ) + if not hop_refused: + assert (await request).choices[0].message.content == "ok" + return + with pytest.raises(litellm.InternalServerError): + await request + finally: + litellm.callbacks.remove(recorder) + + assert recorder.failed_targets == ["healthy-group"] + hop_refusals = [ + record["attempted_retries"] + for record in recorder.breadcrumbs_per_target[0] + if record["model_group"] == "healthy-group" and "Max retries per request hit!" in record["exception_string"] + ] + assert hop_refusals == [0, 1] + + @pytest.mark.asyncio async def test_fallback_traceback_stays_available_at_debug_level(): """Dropping the stack from the ERROR line is only safe because the fallback path still diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 2dfe00b0537..8bf8489fc52 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -29,6 +29,7 @@ from litellm._logging import ( verbose_logger, ) from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor from litellm.proxy.utils import is_valid_api_key from litellm.types.utils import ( @@ -891,7 +892,10 @@ def validate_model_cost_values(model_data, exceptions=None): "input_cost_per_video_per_second_above_8s_interval", "input_cost_per_video_per_second_above_15s_interval", "input_cost_per_video_per_second_above_128k_tokens", + "input_cost_per_audio_token_batches", + "input_cost_per_image_token_batches", "input_cost_per_token_batches", + "input_cost_per_video_token_batches", "output_cost_per_token_batches", "input_cost_per_token_cache_hit", "cache_creation_input_token_cost", @@ -1040,7 +1044,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_second": {"type": "number"}, "input_cost_per_token": {"type": "number"}, "input_cost_per_token_above_128k_tokens": {"type": "number"}, + "input_cost_per_audio_token_batches": {"type": "number"}, + "input_cost_per_image_token_batches": {"type": "number"}, "input_cost_per_token_batches": {"type": "number"}, + "input_cost_per_video_token_batches": {"type": "number"}, "input_cost_per_token_cache_hit": {"type": "number"}, "input_cost_per_video_per_second": {"type": "number"}, "input_cost_per_video_per_second_above_8s_interval": {"type": "number"}, @@ -2945,7 +2952,7 @@ def test_model_info_for_openrouter_kimi_k2_5(): def test_gemini_embedding_2_ga_in_cost_map(): - """GA and Vertex preview gemini-embedding-2 entries align with multimodal unit pricing.""" + """GA and Vertex preview gemini-embedding-2 entries align with multimodal token pricing.""" import json from pathlib import Path @@ -2967,9 +2974,15 @@ def test_gemini_embedding_2_ga_in_cost_map(): assert info.get("mode") == "embedding" assert info.get("supports_multimodal") is True assert info.get("input_cost_per_token") == 2e-07 - assert info.get("input_cost_per_image") == 0.00012 - assert info.get("input_cost_per_audio_per_second") == 0.00016 - assert info.get("input_cost_per_video_per_second") == 0.00079 + assert info.get("input_cost_per_audio_token") == 6.5e-06 + assert info.get("input_cost_per_image_token") == 4.5e-07 + assert info.get("input_cost_per_video_token") == 1.2e-05 + assert info.get("input_cost_per_audio_token_batches") == 3.25e-06 + assert info.get("input_cost_per_image_token_batches") == 2.25e-07 + assert info.get("input_cost_per_video_token_batches") == 6e-06 + assert "input_cost_per_image" not in info + assert "input_cost_per_audio_per_second" not in info + assert "input_cost_per_video_per_second" not in info if provider in ("vertex_ai-embedding-models", "vertex_ai"): assert ( info.get("uses_embed_content") is True @@ -4064,10 +4077,11 @@ class TestMetadataNoneHandling: _RETRY_CAP_CASES: Final = ( - pytest.param(5, {"attempted_retries": 5}, True, id="cap-above-four-reached"), - pytest.param(5, {"attempted_retries": 4}, False, id="cap-above-four-not-reached"), - pytest.param(0, {"attempted_retries": 0}, False, id="first-attempt-passes-cap-of-zero"), - pytest.param(0, {"attempted_retries": 1}, True, id="cap-of-zero-refuses-first-retry"), + pytest.param(5, {"request_retry_count": 5}, True, id="cap-above-four-reached"), + pytest.param(5, {"request_retry_count": 4}, False, id="cap-above-four-not-reached"), + pytest.param(0, {"request_retry_count": 0}, False, id="first-attempt-passes-cap-of-zero"), + pytest.param(0, {"request_retry_count": 1}, True, id="cap-of-zero-refuses-first-retry"), + pytest.param(0, {"attempted_retries": 1}, False, id="per-hop-attempted-retries-is-not-the-cap"), pytest.param(5, {"previous_models": ("a", "b", "c", "d", "e")}, False, id="breadcrumb-count-is-not-the-cap"), pytest.param(5, None, False, id="metadata-none"), ) @@ -4085,7 +4099,9 @@ def _capped_completion_kwargs(metadata_key: str, metadata: object) -> dict[str, @pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) @pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES) -def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metadata_key, cap, metadata, refused): +def test_num_retries_per_request_reads_request_retry_count_sync( + monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool +) -> None: monkeypatch.setattr(litellm, "num_retries_per_request", cap) kwargs: Final = _capped_completion_kwargs(metadata_key, metadata) if refused: @@ -4098,7 +4114,9 @@ def test_num_retries_per_request_reads_attempted_retries_sync(monkeypatch, metad @pytest.mark.asyncio @pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) @pytest.mark.parametrize("cap, metadata, refused", _RETRY_CAP_CASES) -async def test_num_retries_per_request_reads_attempted_retries_async(monkeypatch, metadata_key, cap, metadata, refused): +async def test_num_retries_per_request_reads_request_retry_count_async( + monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, metadata: object, refused: bool +) -> None: monkeypatch.setattr(litellm, "num_retries_per_request", cap) kwargs: Final = _capped_completion_kwargs(metadata_key, metadata) if refused: @@ -4630,6 +4648,16 @@ def test_aws_bedrock_project_id_excluded_from_bedrock_optional_params(): assert result["aws_region_name"] == "us-east-1" +@pytest.mark.parametrize("filter_name", [ + "get_non_default_completion_params", "get_non_default_transcription_params", "filter_out_litellm_params", +]) +def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> None: + filtered = getattr(litellm.utils, filter_name)( + {"provider_option": "kept", "_router_weights": {"group": {"deployment": 100}}} + ) + assert filtered == {"provider_option": "kept"} + + class TestGetOptionalParamsTencent: """Tests that tencent provider uses TencentChatConfig for parameter mapping.""" @@ -5389,6 +5417,26 @@ def test_websearch_interception_control_fields_never_reach_the_provider(): assert set(WEBSEARCH_INTERNAL_CONTROL_FIELDS) <= set(all_litellm_params) +def test_get_litellm_params_keys_never_reach_the_provider(): + """Bridges (chat <-> Responses, agentic loop follow-ups) forward litellm_params as + `completion()` kwargs. Any key the param builder does not recognize is swept into + extra_body, and OpenAI rejects the call with `Unknown parameter: 'model_alias_map'`. + """ + litellm_param_keys = frozenset(get_litellm_params()) - {"drop_params"} + kwargs = { + "a_real_provider_specific_param": 1, + "model_alias_map": {"alias": "gpt-5.4"}, + **{key: "configured-value" for key in litellm_param_keys - {"model_alias_map"}}, + } + + non_default = get_non_default_completion_params(kwargs) + + assert non_default == {"a_real_provider_specific_param": 1}, ( + "litellm params leaked into the provider params: " + f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}" + ) + + def test_bedrock_batch_params_never_reach_the_provider(): """A Bedrock managed-batch deployment carries aws_batch_role_arn / s3_* / bedrock_tags in its litellm_params, and the same deployment also serves chat. diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 77d9ef167d0..dfcd63d3019 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -806,7 +806,7 @@ async def test_shared_call_limits_still_reject_before_reading_ocr_file( monkeypatch.setattr(litellm, "_current_cost", 2) monkeypatch.setattr(litellm, "num_retries_per_request", 1 if limit == "retries" else None) expected: Final = litellm.BudgetExceededError if limit == "budget" else RuntimeError - arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"attempted_retries": 1}} + arguments: Final = {"document": {"type": "file", "file": File()}, "metadata": {"request_retry_count": 1}} with pytest.raises(expected, match=r"Budget has been exceeded|Max retries per request hit"): await call_aocr(ocr_server, **arguments) if asynchronous else call_ocr(ocr_server, **arguments) assert reads == [] diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index e2a08a40bcb..773854d29e6 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2322,7 +2322,7 @@ }, "src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx": { "no-nested-ternary": { - "count": 4 + "count": 3 } }, "src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx index 81c39258f67..86b596d4bcd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx @@ -57,7 +57,7 @@ export function GuardrailDetail({ guardrailId, onBack, accessToken = null, start return list.map((l: Record) => ({ id: l.id as string, timestamp: l.timestamp as string, - action: l.action as "blocked" | "passed" | "flagged", + action: l.action as LogEntry["action"], score: l.score as number | undefined, model: l.model as string | undefined, input_snippet: l.input_snippet as string | undefined, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts index eccd8a80748..56e45216ec5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts @@ -328,7 +328,8 @@ describe("useTeam", () => { }); it("should return team data when query is successful", async () => { - (teamInfoCall as any).mockResolvedValue(mockTeams[0]); + // /team/info answers with an envelope; the hook is typed as the team itself. + (teamInfoCall as any).mockResolvedValue({ team_id: "team-1", team_info: mockTeams[0], keys: [] }); const { result } = renderHook(() => useTeam("team-1"), { wrapper }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts index 14e95bcd543..05025adc5e6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts @@ -163,7 +163,8 @@ export const useTeam = (teamId?: string) => { throw new Error("Missing auth or teamId"); } - return teamInfoCall(accessToken, teamId); + const { team_info } = (await teamInfoCall(accessToken, teamId)) as { team_info: Team }; + return team_info; }, initialData: () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx index cf6576ce6e9..d22ac0d8742 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx @@ -242,13 +242,11 @@ describe("ProjectDetail", () => { it("should show team information when team data is available", () => { mockUseTeam.mockReturnValue({ data: { - team_info: { - team_id: "team-1", - team_alias: "Engineering", - models: ["gpt-4"], - spend: 50, - members_with_roles: [], - }, + team_id: "team-1", + team_alias: "Engineering", + models: ["gpt-4"], + spend: 50, + members_with_roles: [], }, isLoading: false, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx index f94240f3c9e..2585b67c81b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx @@ -14,16 +14,6 @@ import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { EditProjectModal } from "./ProjectModals/EditProjectModal"; import { ProjectKeysSection } from "./ProjectKeysSection"; -interface TeamInfoShape { - team_id: string; - team_alias?: string; - models?: string[]; - max_budget?: number | null; - budget_duration?: string | null; - spend?: number; - members_with_roles?: { user_id: string; role: string }[]; -} - interface ProjectDetailProps { projectId: string; onBack: () => void; @@ -33,10 +23,7 @@ const utilisationTone = (percent: number) => (percent >= 90 ? "over" : percent > export function ProjectDetail({ projectId, onBack }: ProjectDetailProps) { const { data: project, isLoading } = useProjectDetails(projectId); - const { data: teamData } = useTeam(project?.team_id ?? undefined); - // teamInfoCall returns { team_id, team_info: {...}, keys, team_memberships } - const teamInfo: TeamInfoShape | undefined = ((teamData as unknown as { team_info?: TeamInfoShape })?.team_info ?? - teamData) as TeamInfoShape | undefined; + const { data: teamInfo } = useTeam(project?.team_id ?? undefined); const [isEditModalVisible, setIsEditModalVisible] = useState(false); const spend = project?.spend ?? 0; diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx index ab91e10c2fd..083b1e5f3e2 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx @@ -2,7 +2,7 @@ import userEvent from "@testing-library/user-event"; import React from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { renderWithProviders, screen, testQueryClient, waitFor } from "../../../tests/test-utils"; +import { renderWithProviders, screen, testQueryClient, waitFor, within } from "../../../tests/test-utils"; import type { LogEntry as SpendLogEntry } from "@/components/view_logs/columns"; import { LogViewer } from "./LogViewer"; @@ -95,3 +95,16 @@ describe("GuardrailsMonitor LogViewer drawer", () => { }); }); }); + +describe("GuardrailsMonitor LogViewer not_run rows", () => { + it("renders a not_run log as a neutral Not run badge instead of a pass or failure", () => { + renderWithProviders( + , + ); + + const row = screen.getByRole("button", { name: /system prompt only/ }); + expect(within(row).getByText("Not run")).toHaveClass("text-muted-foreground"); + expect(within(row).queryByText("Passed")).not.toBeInTheDocument(); + expect(within(row).queryByText("Blocked")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx index 0703c94c2ed..2abd699ba86 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx @@ -1,4 +1,4 @@ -import { CircleCheck, ChevronDown, TriangleAlert, X } from "lucide-react"; +import { CircleCheck, ChevronDown, MinusCircle, TriangleAlert, X } from "lucide-react"; import { useQuery } from "@tanstack/react-query"; import moment from "moment"; import React, { useState } from "react"; @@ -10,9 +10,16 @@ import type { LogEntry as ViewLogsLogEntry } from "@/components/view_logs/column import type { LogEntry } from "./mockData"; const actionConfig: Record< - "blocked" | "passed" | "flagged", + "blocked" | "passed" | "flagged" | "not_run", { icon: React.ElementType; color: string; bg: string; border: string; label: string } > = { + not_run: { + icon: MinusCircle, + color: "text-muted-foreground", + bg: "bg-muted", + border: "border-border", + label: "Not run", + }, blocked: { icon: X, color: "text-destructive", diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts index 2b42f7907f1..591d5cd3edd 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts @@ -10,7 +10,7 @@ export interface LogEntry { input_snippet?: string; output_snippet?: string; score?: number; - action: "blocked" | "passed" | "flagged"; + action: "blocked" | "passed" | "flagged" | "not_run"; model?: string; reason?: string; latency_ms?: number; diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx index 34a21122027..93a43d5e533 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx @@ -415,6 +415,139 @@ describe("ModelSelect", () => { } }); + it("should take the org model ceiling from the team when /organization/info is not readable", async () => { + const testCases = [ + { + name: "org allows all proxy models", + organizationModels: ["all-proxy-models"], + shouldShowSentinel: true, + offered: ["gpt-4", "claude-3"], + notOffered: [] as string[], + }, + { + name: "org places no ceiling at all", + organizationModels: [], + shouldShowSentinel: true, + offered: ["gpt-4", "claude-3"], + notOffered: [] as string[], + }, + { + name: "org restricts the team to one model", + organizationModels: ["gpt-4"], + shouldShowSentinel: false, + offered: ["gpt-4"], + notOffered: ["claude-3"], + }, + ]; + + for (const testCase of testCases) { + const user = userEvent.setup(); + // A team admin gets a 403 from /organization/info, so the org query never resolves. + mockUseOrganization.mockReturnValue({ data: undefined, isLoading: false } as any); + mockUseTeam.mockReturnValue({ + data: { team_id: "team-1", organization_models: testCase.organizationModels }, + isLoading: false, + } as any); + + const { unmount } = renderWithProviders( + , + ); + + await openModelList(user); + if (testCase.shouldShowSentinel) { + expectOffered("All Proxy Models"); + } else { + expectNotOffered("All Proxy Models"); + } + expectOffered("No Default Models"); + testCase.offered.forEach(expectOffered); + testCase.notOffered.forEach(expectNotOffered); + + unmount(); + } + }); + + it("should stay in the loading state while a list-seeded team is still fetching its org ceiling", () => { + mockUseOrganization.mockReturnValue({ data: undefined, isLoading: false } as any); + mockUseTeam.mockReturnValue({ + data: { team_id: "team-1", models: [] }, + isLoading: false, + isFetching: true, + } as any); + + renderWithProviders( + , + ); + + expect(screen.queryAllByRole("combobox")).toHaveLength(0); + }); + + it("should not hold the loading state on a background refetch once the org ceiling is known", async () => { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ data: undefined, isLoading: false } as any); + mockUseTeam.mockReturnValue({ + data: { team_id: "team-1", organization_models: ["all-proxy-models"] }, + isLoading: false, + isFetching: true, + } as any); + + renderWithProviders( + , + ); + + await openModelList(user); + expectOffered("All Proxy Models"); + }); + + it("should offer no models for an org team when neither the team nor the org reports a ceiling", async () => { + const testCases = [ + { name: "/team/info withheld the ceiling", team: { team_id: "team-1", organization_models: null } }, + { name: "/team/info failed after the list seeded the team", team: { team_id: "team-1", models: [] } }, + ]; + + for (const testCase of testCases) { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ data: undefined, isLoading: false } as any); + mockUseTeam.mockReturnValue({ data: testCase.team, isLoading: false, isFetching: false } as any); + + const { unmount } = renderWithProviders( + , + ); + + await openModelList(user); + expectNotOffered("All Proxy Models"); + expectOffered("No Default Models"); + expectNotOffered("gpt-4"); + expectNotOffered("claude-3"); + + unmount(); + } + }); + it("should use custom dataTestId when provided", async () => { renderWithProviders( + organizationModels.length === 0 || organizationModels.includes(MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value); + +// useTeam seeds from the team list, which omits organization_models; /team/info is the only source of the org ceiling. +const isAwaitingOrganizationModels = (team: Team | undefined, isFetchingTeam: boolean) => + isFetchingTeam && team !== undefined && team.organization_models === undefined; + const contextFilters: Record string[]> = { user: ({ allProxyModels, userModels, options }) => { if (!userModels) return []; @@ -82,18 +89,10 @@ const contextFilters: Record { - if (selectedOrganization) { - if ( - selectedOrganization.models.includes(MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value) || - selectedOrganization.models.length === 0 - ) { - return allProxyModels; - } - return allProxyModels.filter((model) => selectedOrganization.models.includes(model)); - } - - return allProxyModels ?? []; + team: ({ allProxyModels, organizationID, organizationModels }) => { + if (organizationModels === undefined) return organizationID ? [] : allProxyModels; + if (isUncappedModelCeiling(organizationModels)) return allProxyModels; + return allProxyModels.filter((model) => organizationModels.includes(model)); }, organization: ({ allProxyModels }) => { @@ -108,7 +107,7 @@ const contextFilters: Record { const deduplicatedProxyModels = Array.from(new Map(allProxyModels.map((m) => [m.id, m])).values()).map( (model) => model.id, @@ -118,7 +117,13 @@ const filterModels = ( const filterFn = contextFilters[ctx.context]; if (!filterFn) return []; - return filterFn({ allProxyModels: deduplicatedProxyModels, ...extra, options: ctx.options }); + const filterArgs: FilterContextArgs = { + allProxyModels: deduplicatedProxyModels, + organizationID: ctx.organizationID, + ...extra, + options: ctx.options, + }; + return filterFn(filterArgs); }; export const ModelSelect = (props: ModelSelectProps) => { @@ -126,16 +131,17 @@ export const ModelSelect = (props: ModelSelectProps) => { const { id, teamID, organizationID, options, context, dataTestId, value = [], onChange, style } = props; const { showAllProxyModelsOverride, includeSpecialOptions } = options || {}; const { data: allProxyModels, isLoading: isLoadingAllProxyModels } = useAllProxyModels(); - const { data: team, isLoading: isLoadingTeam } = useTeam(teamID); + const { data: team, isLoading: isLoadingTeam, isFetching: isFetchingTeam } = useTeam(teamID); const { data: organization, isLoading: isLoadingOrganization } = useOrganization(organizationID); const { data: currentUser, isLoading: isCurrentUserLoading } = useCurrentUser(); const isSpecialOption = (value: string) => MODEL_SENTINEL_OPTIONS.some((sv) => sv.value === value); const hasSpecialOptionSelected = value.some(isSpecialOption); - const isLoading = isLoadingAllProxyModels || isLoadingTeam || isLoadingOrganization || isCurrentUserLoading; - const organizationHasAllProxyModels = - organization?.models.includes(MODEL_SELECT_ALL_PROXY_MODELS_SPECIAL_VALUE.value) || - organization?.models.length === 0; + const isTeamPending = isLoadingTeam || isAwaitingOrganizationModels(team, isFetchingTeam); + const isLoading = isLoadingAllProxyModels || isTeamPending || isLoadingOrganization || isCurrentUserLoading; + // The org's ceiling rides on /team/info, which a team admin may read; /organization/info 403s for them. + const organizationModels = team?.organization_models ?? organization?.models; + const organizationHasAllProxyModels = organizationModels !== undefined && isUncappedModelCeiling(organizationModels); const shouldShowAllProxyModels = showAllProxyModelsOverride || (organizationHasAllProxyModels && includeSpecialOptions) || context === "global"; @@ -159,8 +165,7 @@ export const ModelSelect = (props: ModelSelectProps) => { }; const filteredModels = filterModels(allProxyModels?.data ?? [], props, { - selectedTeam: team, - selectedOrganization: organization, + organizationModels, userModels: currentUser?.models, }); diff --git a/ui/litellm-dashboard/src/components/add_model/AffinityControls.tsx b/ui/litellm-dashboard/src/components/add_model/AffinityControls.tsx index 9022d424369..325362ea177 100644 --- a/ui/litellm-dashboard/src/components/add_model/AffinityControls.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AffinityControls.tsx @@ -28,13 +28,13 @@ export const AffinityControls: React.FC<{ onChange({ ...value, deployment_affinity: deploymentAffinity })} - aria-label="Pin a session to one deployment per model group" + aria-label="Pin one model deployment per tier" /> - Pin a session to one deployment per model group + Pin one model deployment per tier - Keeps a session on the same deployment within a group, so provider prompt caches stay warm. Turn off to - load-balance every turn. + Reuses the model chosen for each tier and its deployment when available. Requests can still move between tiers. + Turn off to select models and load-balance deployments every turn.