chore: merge main into litellm_remove_lit002_dict_ban
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-10-09 07:18:45 +00:00
commit d14859c4c7
135 changed files with 8726 additions and 2396 deletions

View file

@ -856,58 +856,6 @@ jobs:
- store_test_results:
path: test-results
litellm_router_unit_testing: # Runs all tests with the "router" keyword
docker:
- *python312_image
working_directory: ~/project
resource_class: large
steps:
- checkout
- skip_if_unrelated_changes
- setup_google_dns
- install_uv
- install_rust
- restore_cache:
keys:
- v1-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- save_cache:
paths:
- ~/.cache/uv
key: v1-uv-cache-{{ checksum "uv.lock" }}
# Run pytest and generate JUnit XML report
- setup_litellm_enterprise_pip
- run:
name: Run tests
command: |
mkdir -p test-results
TEST_FILES=$(circleci tests glob "tests/router_unit_tests/**/test_*.py")
echo "$TEST_FILES" | circleci tests run \
--verbose \
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
-v \
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
--junitxml=test-results/junit.xml \
--durations=5 \
-n 4"
no_output_timeout: 15m
- run:
name: Rename the coverage files
command: |
mv coverage.xml router_unit_tests_coverage.xml
mv .coverage router_unit_tests_coverage
# Store test results
- store_test_results:
path: test-results
- persist_to_workspace:
root: .
paths:
- router_unit_tests_coverage.xml
- router_unit_tests_coverage
llm_translation_testing:
docker:
- *python312_image
@ -1024,16 +972,13 @@ jobs:
command: |
mkdir -p test-results
export LITELLM_LOG=WARNING
TEST_FILES=$(circleci tests glob "tests/guardrails_tests/**/test_*.py")
echo "$TEST_FILES" | circleci tests run \
--verbose \
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
-vv \
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
--junitxml=test-results/junit.xml \
--durations=5 \
-n 2 \
--timeout=120 --timeout_method=thread"
uv run --no-sync python -m pytest --tb=short \
-vv \
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
--junitxml=test-results/junit.xml \
--durations=5 \
--timeout=120 --timeout_method=thread \
"tests/guardrails_tests/test_custom_guardrail.py::test_get_guardrail_dynamic_request_body_params"
no_output_timeout: 15m
- run:
name: Rename the coverage files
@ -2493,7 +2438,7 @@ jobs:
- run:
name: Combine Coverage
command: |
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage google_generate_content_endpoint_coverage litellm_utils_coverage router_unit_tests_coverage auth_ui_unit_tests_coverage
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage google_generate_content_endpoint_coverage litellm_utils_coverage
uv tool run --from 'coverage[toml]==7.10.6' coverage xml
- codecov/upload:
file: ./coverage.xml
@ -3433,15 +3378,38 @@ workflows:
- not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- using_litellm_on_windows
- windows_release_wheel
- provider_replay_harness
- base_sdk_install
- local_testing_part1
- local_testing_part2
- langfuse_logging_unit_tests
- litellm_router_testing
- litellm_router_unit_testing
- using_litellm_on_windows:
filters:
branches:
only: main
- windows_release_wheel:
filters:
branches:
only: main
- provider_replay_harness:
filters:
branches:
only: main
- base_sdk_install:
filters:
branches:
only: main
- local_testing_part1:
filters:
branches:
only: main
- local_testing_part2:
filters:
branches:
only: main
- langfuse_logging_unit_tests:
filters:
branches:
only: main
- litellm_router_testing:
filters:
branches:
only: main
- auth_ui_unit_tests
- build_docker_database_image
- e2e_ui_testing
@ -3461,30 +3429,68 @@ workflows:
- proxy_store_model_in_db_tests:
requires:
- build_docker_database_image
- proxy_build_from_pip_tests
- proxy_build_from_pip_tests:
filters:
branches:
only: main
- proxy_pass_through_endpoint_tests:
requires:
- build_docker_database_image
- proxy_e2e_anthropic_messages_tests:
requires:
- build_docker_database_image
- llm_translation_testing
- realtime_translation_testing
- llm_translation_testing:
filters:
branches:
only: main
- realtime_translation_testing:
filters:
branches:
only: main
- guardrails_testing
- google_generate_content_endpoint_testing
- ocr_testing
- search_testing
- batches_testing
- litellm_utils_testing
- pass_through_unit_testing
- image_gen_testing
- logging_testing
- audio_testing
- google_generate_content_endpoint_testing:
filters:
branches:
only: main
- ocr_testing:
filters:
branches:
only: main
- search_testing:
filters:
branches:
only: main
- batches_testing:
filters:
branches:
only: main
- litellm_utils_testing:
filters:
branches:
only: main
- pass_through_unit_testing:
filters:
branches:
only: main
- image_gen_testing:
filters:
branches:
only: main
- logging_testing:
filters:
branches:
only: main
- audio_testing:
filters:
branches:
only: main
- upload-coverage:
filters:
branches:
only: main
requires:
- realtime_translation_testing
- google_generate_content_endpoint_testing
- guardrails_testing
- ocr_testing
- search_testing
- batches_testing
@ -3496,17 +3502,30 @@ workflows:
- langfuse_logging_unit_tests
- local_testing_part1
- local_testing_part2
- litellm_router_unit_testing
- auth_ui_unit_tests
- db_migration_disable_update_check:
requires:
- build_docker_database_image
- installing_litellm_on_python
- installing_litellm_on_python_3_13
- installing_litellm_on_python_v2_migration_resolver
- installing_litellm_on_python:
filters:
branches:
only: main
- installing_litellm_on_python_3_13:
filters:
branches:
only: main
- installing_litellm_on_python_v2_migration_resolver:
filters:
branches:
only: main
- helm_chart_testing:
filters:
branches:
only: main
requires:
- build_docker_database_image
- test_bad_database_url:
filters:
branches:
only: main
requires:
- build_docker_database_image

View file

@ -1,15 +0,0 @@
{
"cases": {
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"MODEL-ALLOW": "tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py::test_can_object_call_model_allows_listed_model_for_key",
"MODEL-DENY": "tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
"CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
}
}

View file

@ -36,7 +36,15 @@ GLOB_CHARS = frozenset("*?")
# itself decomposed one level deeper and is checked through its own entry.
SHARDED_ROOTS: tuple[str, ...] = (
"tests/test_litellm",
"tests/unit",
"tests/unit/enterprise",
"tests/unit/enterprise/enterprise_callbacks",
"tests/unit/enterprise/proxy",
"tests/unit/llms",
"tests/unit/proxy",
"tests/unit/proxy/_experimental",
"tests/unit/responses",
"tests/unit/skills",
)

View file

@ -16,29 +16,11 @@ import socket
import subprocess
import sys
import time
from collections import Counter
from collections.abc import Sequence
from dataclasses import dataclass, field
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final, NoReturn, TextIO, cast
import pytest
EXPECTED_CASES: Final = (
"CHAT-JSON",
"CHAT-TEXT-STREAM",
"CHAT-TOOL-STREAM",
"MODEL-ALLOW",
"MODEL-DENY",
"COST-EXPLICIT",
"COST-ZERO",
"LOG-CONTENT-ON",
"LOG-CONTENT-OFF",
"CALLBACK-SUCCESS",
"CALLBACK-FAILURE",
)
@dataclass(frozen=True, slots=True)
class CheckResult:
@ -57,8 +39,6 @@ class _Args:
ready_deadline: float = 120.0
shutdown_deadline: float = 20.0
poll_interval: float = 0.5
manifest: str = ""
rootdir: str | None = None
def fail(reason: str) -> NoReturn:
@ -345,120 +325,6 @@ def _terminate(proc: subprocess.Popen[bytes], log_file: TextIO) -> None:
log_file.close()
def _load_manifest(path: Path) -> MappingProxyType[str, str]:
def no_duplicates(pairs: list[tuple[object, object]]) -> dict[object, object]:
seen: dict[object, object] = {}
for key, value in pairs:
if key in seen:
raise ValueError(f"duplicate key in manifest: {key}")
seen[key] = value
return seen
raw_value: object = cast(object, json.loads(path.read_text(), object_pairs_hook=no_duplicates))
if not isinstance(raw_value, dict):
raise ValueError("manifest must be an object")
loaded: Final = cast(dict[object, object], raw_value)
cases_value: object = loaded.get("cases")
if not isinstance(cases_value, dict):
raise ValueError("manifest must be an object with a 'cases' object")
cases_any: Final = cast(dict[object, object], cases_value)
cases: Final = {k: v for k, v in cases_any.items() if isinstance(k, str) and isinstance(v, str)}
if len(cases) != len(cases_any):
raise ValueError("manifest 'cases' must map string ids to string node ids")
return MappingProxyType(cases)
@dataclass(slots=True, eq=False)
class _Recorder:
collect_failed: list[str] = field(default_factory=list)
collected: tuple[str, ...] = ()
reports: dict[str, list[tuple[str, str, bool]]] = field(default_factory=dict)
def pytest_collectreport(self, report: pytest.CollectReport) -> None:
if report.failed:
self.collect_failed.append(report.nodeid)
def pytest_collection_finish(self, session: pytest.Session) -> None:
self.collected = tuple(item.nodeid for item in session.items)
def pytest_runtest_logreport(self, report: pytest.TestReport) -> None:
self.reports.setdefault(report.nodeid, []).append((report.when, report.outcome, hasattr(report, "wasxfail")))
def cmd_pytest(args: _Args) -> int:
try:
cases: Final = _load_manifest(Path(args.manifest))
except (OSError, ValueError, json.JSONDecodeError) as exc:
fail(f"manifest invalid: {exc}")
if tuple(cases) != EXPECTED_CASES:
fail(f"manifest case ids must be exactly {list(EXPECTED_CASES)} in order, got {list(cases)}")
node_ids: Final = tuple(cases.values())
if len(set(node_ids)) != len(node_ids):
fail("manifest node ids are not unique")
argv: Final = [
*node_ids,
"-p",
"no:cacheprovider",
"-p",
"no:xdist",
"-p",
"no:rerunfailures",
"-p",
"no:randomly",
"-rA",
"-q",
*(["--rootdir", args.rootdir] if args.rootdir else []),
]
recorder: Final = _Recorder()
code: Final = pytest.main(argv, plugins=[recorder])
name_of: Final = MappingProxyType({node_id: case_id for case_id, node_id in cases.items()})
problems: Final[list[str]] = []
if code != 0:
problems.append(f"pytest exit code {code}")
for failed_id in recorder.collect_failed:
problems.append(f"collection failed: {name_of.get(failed_id, failed_id)}")
expected: Final = Counter(node_ids)
collected: Final = Counter(recorder.collected)
for node_id in expected - collected:
problems.append(f"missing case {name_of[node_id]} ({node_id})")
for node_id in collected - expected:
problems.append(f"unexpected test collected: {node_id}")
for node_id, count in collected.items():
if count > 1:
problems.append(f"duplicated test id: {node_id}")
if len(recorder.collected) != len(EXPECTED_CASES):
problems.append(f"collected {len(recorder.collected)} tests, expected {len(EXPECTED_CASES)}")
rows: Final[list[tuple[str, bool]]] = []
for case_id, node_id in cases.items():
reports = recorder.reports.get(node_id, [])
case_ok = (
bool(reports)
and all(outcome == "passed" and not wasxfail for _, outcome, wasxfail in reports)
and {when for when, _, _ in reports} >= {"setup", "call", "teardown"}
)
rows.append((case_id, case_ok))
if not reports:
problems.append(f"{case_id} ({node_id}) produced no runtest reports")
continue
for when, outcome, wasxfail in reports:
if outcome != "passed":
problems.append(f"{case_id} ({node_id}) {when} outcome={outcome}")
if wasxfail:
problems.append(f"{case_id} ({node_id}) {when} was xfail/xpass")
missing_phases = {"setup", "call", "teardown"} - {when for when, _, _ in reports}
for phase in sorted(missing_phases):
problems.append(f"{case_id} ({node_id}) missing {phase} report")
for case_id, passed in rows:
print(f"{case_id} {'PASS' if passed else 'FAIL'} {cases[case_id]}")
if problems:
for problem in problems:
print(f"merge-smoke: {problem}", file=sys.stderr)
fail("pytest verdict failed")
ok("pytest 11 cases")
return 0
def main() -> int:
parser: Final = argparse.ArgumentParser(description=__doc__)
subs: Final = parser.add_subparsers(dest="command", required=True)
@ -475,16 +341,12 @@ def main() -> int:
p_proxy.add_argument("--ready-deadline", type=float, default=120)
p_proxy.add_argument("--shutdown-deadline", type=float, default=20)
p_proxy.add_argument("--poll-interval", type=float, default=0.5)
p_test: Final = subs.add_parser("pytest")
p_test.add_argument("--manifest", required=True)
p_test.add_argument("--rootdir", default=None)
args: Final = parser.parse_args(namespace=_Args())
handlers: Final = {
"verify-isolation": cmd_verify_isolation,
"interpreter": cmd_interpreter,
"cli": cmd_cli,
"proxy-startup": cmd_proxy_startup,
"pytest": cmd_pytest,
}
return handlers[args.command](args)

View file

@ -1,283 +0,0 @@
name: _Unit Test Base (Reusable)
on:
workflow_call:
inputs:
python-version:
description: "Python version used to install dependencies and run tests"
required: false
type: string
default: "3.12"
test-path:
description: >-
Space-separated pytest paths to run. A path that no longer exists is
dropped with a warning instead of being passed to pytest, because one
missing path makes pytest-xdist collect nothing and report exit 5, which
the step treats as a drained shard. Options are passed through as
written, so use the `--flag=value` form: a bare `--ignore path` would
have its path existence-checked like any other token.
required: true
type: string
workers:
description: "Number of pytest-xdist workers"
required: false
type: number
default: 2
reruns:
description: "Number of reruns for flaky tests"
required: false
type: number
default: 2
timeout-minutes:
description: >-
Timeout for the test step alone. Setup (checkout, dependency install,
Prisma client generation) gets its own allowance on top, so a slow
runner or a cold binary download can never cancel passing tests.
required: false
type: number
default: 20
job-timeout-minutes:
description: >-
Backstop for the whole job. Keep it >= `timeout-minutes` plus 40: 35 for
the per-step ceilings on the setup steps below, and 5 for the runner
overhead the job clock charges but no step owns (job init, step
transitions, post-job cleanup). That headroom is what makes the test
budget a floor rather than a hope, since setup cannot overrun into it
without failing its own step first. GitHub expressions have no
arithmetic, so the sum is passed in rather than computed.
required: false
type: number
default: 60
test-timeout-seconds:
description: >-
Per-test ceiling enforced by pytest-timeout, covering fixture setup and
teardown as well as the test body. A test that hangs fails with a
traceback of where it was stuck instead of idling the shard until
`timeout-minutes` cancels it. Timed-out tests are excluded from reruns
because pytest-timeout arms its timer once per test and
pytest-rerunfailures reruns inside that same window, so a rerun of a
timed-out test would run with no timer at all.
required: false
type: number
default: 120
max-failures:
description: "Stop after this many failures"
required: false
type: number
default: 10
dist:
description: "pytest-xdist distribution mode (loadscope|load|worksteal|loadfile|no)"
required: false
type: string
default: "loadscope"
artifact-name:
description: "Unique name for the coverage artifact (must be unique per run)"
required: true
type: string
rust-bridge-artifact:
description: "Prebuilt editable Rust bridge artifact"
required: false
type: string
default: ""
permissions:
contents: read
env:
UV_PYTHON: ${{ inputs.python-version }}
LITELLM_LOCAL_MODEL_COST_MAP: "True"
jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.job-timeout-minutes }}
permissions:
contents: read
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}
has-coverage: ${{ steps.tests.outputs.has-coverage }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
timeout-minutes: 3
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: ${{ env.UV_PYTHON }}
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
uses: ./.github/actions/cache-uv-downloads
- name: Set up the Rust build
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
uses: ./.github/actions/rust-bridge
with:
artifact: ${{ inputs.rust-bridge-artifact }}
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
env:
RUST_BRIDGE_ARTIFACT: ${{ inputs.rust-bridge-artifact }}
run: |
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
if [ -z "$RUST_BRIDGE_ARTIFACT" ]; then
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils
else
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils --no-install-project
uv pip install --no-deps --python .venv/bin/python rust-bridge-dist/*.whl
cp rust-bridge-dist/litellm/rust_bridge/_native.abi3.so litellm/rust_bridge/_native.abi3.so
uv run --no-sync python -c "import importlib.metadata; import litellm.rust_bridge._native; print(importlib.metadata.version('litellm'))"
fi
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"]'
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
id: tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
TEST_TIMEOUT_SECONDS: ${{ inputs.test-timeout-seconds }}
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
if [ -z "${selection// /}" ]; then
echo "shard selection is empty; nothing to run"
exit 0
fi
pytest_args=()
existing_paths=0
for token in ${selection}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
if [ -e "${token%%::*}" ]; then
pytest_args+=("${token}")
existing_paths=$((existing_paths + 1))
else
echo "::warning::${token} does not exist; drop it from this shard's test-path"
fi
;;
esac
done
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in the selection exists (${selection}); nothing to run"
exit 0
fi
xdist_args=()
if [ "${WORKERS}" != "0" ]; then
xdist_args=(-n "${WORKERS}" --dist="${DIST}")
fi
set +e
uv run --no-sync pytest "${pytest_args[@]}" \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
"${xdist_args[@]}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--timeout="${TEST_TIMEOUT_SECONDS}" \
--rerun-except "from pytest-timeout" \
--durations=20 \
--cov=./litellm --cov=./enterprise/litellm_enterprise \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
status=$?
set -e
if [ -f coverage.xml ]; then
echo "has-coverage=true" >> "$GITHUB_OUTPUT"
fi
if [ "$status" -eq 5 ]; then
echo "pytest collected no tests from ${selection}; passing"
exit 0
fi
exit "$status"
- name: Save coverage report
if: always() && steps.changes.outputs.decision != 'skip'
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage.xml
retention-days: 1
upload-coverage:
name: Upload coverage to Codecov
needs: run
if: always() && needs.run.outputs.decision != 'skip' && needs.run.outputs.has-coverage == 'true'
runs-on: ubuntu-latest
permissions:
contents: read
id-token: write
pull-requests: write
steps:
- name: Checkout code
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Download coverage report
uses: actions/download-artifact@95815c38cf2ff2164869cbab79da8d1f422bc89e # v4.2.1
with:
pattern: coverage-${{ inputs.artifact-name }}-${{ github.run_id }}-${{ github.run_attempt }}
path: coverage-reports
merge-multiple: true
- name: Upload to Codecov
id: codecov-upload
continue-on-error: true
uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5.5.5
with:
use_oidc: true
directory: coverage-reports
root_dir: ${{ github.workspace }}
flags: ${{ inputs.artifact-name }}
fail_ci_if_error: false
- name: Upload to Codecov (retry)
if: steps.codecov-upload.outcome == 'failure'
continue-on-error: true
uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5.5.5
with:
use_oidc: true
directory: coverage-reports
root_dir: ${{ github.workspace }}
flags: ${{ inputs.artifact-name }}
fail_ci_if_error: false

View file

@ -1,134 +0,0 @@
name: Check UI API Types Sync
on:
pull_request:
branches:
- main
- "litellm_**"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
check-sync:
name: Verify schema.d.ts matches the proxy OpenAPI spec
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
fetch-depth: 2
- name: Detect changes that can affect the generated types
id: changes
run: |
set -euo pipefail
if ! base="$(git rev-parse --verify --quiet HEAD^2 >/dev/null && git rev-parse HEAD^1)"; then
echo "Not a pull request merge commit, running the full check."
echo "relevant=true" >> "$GITHUB_OUTPUT"
exit 0
fi
files="$(git diff --name-only "$base" HEAD)"
if grep -Eq '^(litellm/(proxy|types)/|ui/litellm-dashboard/(src/lib/http/schema\.d\.ts|scripts/gen-api-types\.mjs|package(-lock)?\.json)$|\.github/workflows/check-ui-api-types\.yml$)' <<< "$files"; then
echo "relevant=true" >> "$GITHUB_OUTPUT"
else
echo "No proxy, types or generator changes in this pull request, nothing to verify."
echo "relevant=false" >> "$GITHUB_OUTPUT"
fi
- name: Set up Python
if: steps.changes.outputs.relevant == 'true'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.relevant == 'true'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache the Rust build
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-cargo-build
- name: Install backend dependencies
if: steps.changes.outputs.relevant == 'true'
run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.relevant == 'true'
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Regenerate the lazy OpenAPI snapshot
if: steps.changes.outputs.relevant == 'true'
run: uv run --no-sync python -m litellm.proxy._lazy_openapi_snapshot
- name: Fail if the lazy OpenAPI snapshot is stale
if: steps.changes.outputs.relevant == 'true'
run: |
if ! git diff --exit-code -- litellm/proxy/_lazy_openapi_snapshot.json; then
echo "::error file=litellm/proxy/_lazy_openapi_snapshot.json::The lazy OpenAPI snapshot is out of sync with the lazily loaded routes."
echo ""
echo "A lazily loaded route or model changed without regenerating the snapshot that /openapi.json serves for unloaded features."
echo "To fix, run from the repo root:"
echo " uv run python -m litellm.proxy._lazy_openapi_snapshot"
echo "then run npm run gen:api from ui/litellm-dashboard and commit both files."
exit 1
fi
echo "_lazy_openapi_snapshot.json is in sync with the lazily loaded routes."
- name: Set up Node.js
if: steps.changes.outputs.relevant == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dashboard dependencies
if: steps.changes.outputs.relevant == 'true'
working-directory: ui/litellm-dashboard
run: npm ci
- name: Regenerate types from the live spec
if: steps.changes.outputs.relevant == 'true'
working-directory: ui/litellm-dashboard
env:
LITELLM_PYTHON: "uv run --no-sync python"
run: npm run gen:api
- name: Fail if types are stale
if: steps.changes.outputs.relevant == 'true'
run: |
if ! git diff --exit-code -- ui/litellm-dashboard/src/lib/http/schema.d.ts; then
echo "::error file=ui/litellm-dashboard/src/lib/http/schema.d.ts::Generated API types are out of sync with the proxy OpenAPI spec."
echo ""
echo "A backend route or model changed without regenerating the dashboard types."
echo "To fix, run from ui/litellm-dashboard:"
echo " npm run gen:api"
echo "then commit the updated src/lib/http/schema.d.ts."
exit 1
fi
echo "schema.d.ts is in sync with the proxy OpenAPI spec."

View file

@ -1,48 +0,0 @@
name: "CI Coverage"
on:
pull_request:
branches:
- main
- "litellm_**"
push:
branches:
- main
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
assert-ci-coverage:
name: assert-ci-coverage
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
contents: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Assert every test file and Dockerfile is invoked by a job
run: |
python -m pip install "pyyaml==6.0.3"
python .github/scripts/assert_ci_coverage.py
# The census asks whether a job names a file; this asks whether that job's -k
# then throws it back out. A file both globbed and deselected everywhere runs
# nowhere while counting as covered, which is how the caching suite went unrun.
- name: Assert no -k expression deselects a file from every job that globs it
run: python .github/scripts/assert_ci_coverage.py --slices
- name: Assert .github/workflows/ holds only workflows, correctly named
run: python .github/scripts/assert_workflow_dir_hygiene.py

View file

@ -64,7 +64,7 @@ jobs:
uses: ./.github/actions/cache-prisma-binaries
# The three source scanners only need the pinned dev tools (ruff and the
# stdlib checkers), the same versions test-linting.yml's lint job runs.
# stdlib checkers), the same versions test-linting.yml's python job runs.
- name: Install the dev tools
if: matrix.checker != 'basedpyright'
run: |

View file

@ -0,0 +1,397 @@
name: Required checks (legacy)
on:
pull_request:
branches:
- main
- "litellm_**"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: true
jobs:
lint:
name: lint
permissions:
contents: read
pull-requests: read
actions: read
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
ref: ${{ github.event.pull_request.head.sha }}
fetch-depth: 1
clean: true
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Fetch gate base (merge-base with target branch)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
MERGE_BASE=$(retry gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$MERGE_BASE"
retry git fetch --no-tags --depth=1 origin "$MERGE_BASE"
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-lint-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-lint-
- name: Clean Python cache
if: steps.changes.outputs.decision != 'skip'
run: |
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check uv.lock is up to date
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check || (echo "❌ uv.lock is out of sync with pyproject.toml. Run 'uv lock' locally and commit the result." && exit 1)
- name: Cache the Rust build
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Check ruff format
if: steps.changes.outputs.decision != 'skip'
run: |
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- ':(glob)litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
if [ ! -s "$RUNNER_TEMP/ruff_format_files.txt" ]; then
echo "No changed litellm Python files to check with ruff format."
exit 0
fi
xargs uv run --no-sync ruff format --check --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
- name: Debug - Check file state
if: steps.changes.outputs.decision != 'skip'
run: |
echo "Current branch:"
git branch --show-current
echo "Last 3 commits:"
git log --oneline -3
echo "File content around line 43:"
head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10
- name: Check MCP operation boundary
if: steps.changes.outputs.decision != 'skip'
run: uv run --no-sync python scripts/check_mcp_operation_boundary.py
- name: Run Ruff linting
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync ruff check .
cd ..
- name: Run Ruff linting (test tree)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync ruff check --config ruff-tests.toml tests
- name: Check strict ruff rules (delta vs merge-base counts)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/ruff_strict_gate.py --base "$GATE_BASE_SHA"
- name: Check type discipline (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs merge-base counts)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA"
- name: Check test quality (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, conftest snapshot inventory, delta vs merge-base counts)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/test_quality_gate.py --base "$GATE_BASE_SHA"
- name: Print OpenAI version
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- name: Check basedpyright (delta vs merge-base counts)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
- name: Check tests/e2e basedpyright (zero errors)
if: steps.changes.outputs.decision != 'skip'
run: |
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- ':(glob)tests/e2e/**/*.py' ':(glob)tests/e2e_harness/**/*.py' pyrightconfig.json | grep -q .; then
uv run --no-sync basedpyright tests/e2e tests/e2e_harness
else
echo "No changed tests/e2e Python files; skipping."
fi
- name: Run the e2e harness tests
if: steps.changes.outputs.decision != 'skip'
env:
LITELLM_MASTER_KEY: sk-e2e-harness-tests-reach-no-proxy
run: |
if ! git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- tests/e2e tests/e2e_harness ':(exclude)tests/e2e/ui' pyproject.toml uv.lock .github/workflows/test-linting.yml | grep -q .; then
echo "No changed e2e harness files; skipping."
exit 0
fi
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
CLAUDE_VERSION="$(retry uv run --no-sync python tests/e2e/claude_code/pr_gate_version_resolver.py)"
tests/e2e/claude_code/cron_vm/install_claude_code.sh "$CLAUDE_VERSION" "$RUNNER_TEMP/claude-cli"
PATH="$RUNNER_TEMP/claude-cli:$PATH" uv run --no-sync pytest -q tests/e2e_harness
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync python ../tests/documentation_tests/test_circular_imports.py
cd ..
- name: Check import safety
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
frontend-lint:
name: frontend-lint
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 8
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Collect changed files
id: changed
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
[ -f "$f" ] || continue
case "$f" in
*.js | *.jsx | *.ts | *.tsx | *.mjs | *.cjs)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt"
printf '%s\n' "$f" >> "$RUNNER_TEMP/eslint_files.txt" ;;
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$merge_base" "$HEAD_SHA" -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else
echo "has_files=false" >> "$GITHUB_OUTPUT"
echo "No lintable UI files changed in this PR; nothing to check."
fi
- name: Setup Node.js
if: steps.changed.outputs.has_files == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changed.outputs.has_files == 'true'
run: npm ci
- name: Lint changed files (prettier + eslint)
if: steps.changed.outputs.has_files == 'true'
run: |
prettier_files=()
eslint_files=()
while IFS= read -r f; do prettier_files+=("$f"); done < "$RUNNER_TEMP/prettier_files.txt"
while IFS= read -r f; do eslint_files+=("$f"); done < "$RUNNER_TEMP/eslint_files.txt"
status=0
if [ ${#prettier_files[@]} -gt 0 ]; then
echo "::group::Prettier (${#prettier_files[@]} files)"
npx prettier --check "${prettier_files[@]}" || { status=1; echo "::error::Unformatted files. Fix with: npm run format"; }
echo "::endgroup::"
fi
if [ ${#eslint_files[@]} -gt 0 ]; then
echo "::group::ESLint (${#eslint_files[@]} files)"
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_files[@]}" || status=1
echo "::endgroup::"
fi
exit $status
- name: Check lint budgets
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: |
npx eslint . -f json -o "$RUNNER_TEMP/lint-report.json" || true
node scripts/check-lint-budgets.mjs "$RUNNER_TEMP/lint-report.json" eslint-budgets.json
- name: Check for dead code (knip)
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: npm run knip:ci
assert-ci-coverage:
name: assert-ci-coverage
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 5
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Assert every test file and Dockerfile is invoked by a job
run: |
python -m pip install "pyyaml==6.0.3"
python .github/scripts/assert_ci_coverage.py
- name: Assert no -k expression deselects a file from every job that globs it
run: python .github/scripts/assert_ci_coverage.py --slices
- name: Assert .github/workflows/ holds only workflows, correctly named
run: python .github/scripts/assert_workflow_dir_hygiene.py
dashboard-build:
name: Dashboard build
runs-on: ubuntu-24.04
timeout-minutes: 30
steps:
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the dashboard stage
run: docker build --target ui-builder -f Dockerfile .
core-checks:
name: Core checks (Python ${{ matrix.python-version }})
runs-on: ubuntu-24.04
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
env:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: ${{ matrix.python-version }}
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: .github/scripts/uv_sync_with_retries.sh --frozen --extra proxy --extra cli --group dev --group proxy-dev --python ${{ matrix.python-version }}
- name: Create the loopback-only network namespace
run: |
sudo ip netns add smoke
sudo ip netns exec smoke ip link set lo up
cat > "${RUNNER_TEMP}/in-netns" <<'WRAP'
#!/usr/bin/env bash
set -euo pipefail
exec sudo --preserve-env=LITELLM_LOCAL_MODEL_COST_MAP ip netns exec smoke setpriv --reuid "$(id -u)" --regid "$(id -g)" --init-groups -- env HOME="${HOME}" PATH="${PATH}" "$@"
WRAP
chmod +x "${RUNNER_TEMP}/in-netns"
echo "IN_NETNS=${RUNNER_TEMP}/in-netns" >> "${GITHUB_ENV}"
- name: Verify namespace isolation
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py verify-isolation
- name: Verify interpreter version
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py interpreter --expect ${{ matrix.python-version }}
- name: Import and CLI checks
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py cli
- name: Proxy startup check
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py proxy-startup --diagnostics-dir "${RUNNER_TEMP}/smoke-diagnostics"
- name: Upload smoke diagnostics
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: merge-smoke-diagnostics-py${{ matrix.python-version }}
path: ${{ runner.temp }}/smoke-diagnostics
if-no-files-found: ignore
- name: Remove the network namespace
if: always()
run: sudo ip netns delete smoke

View file

@ -1,223 +0,0 @@
name: Code Quality Checks
on:
pull_request:
branches:
- main
- "litellm_**"
push:
branches:
- main
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
code-quality:
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: github.ref == 'refs/heads/main'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache uv dependencies
if: github.ref != 'refs/heads/main'
uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
run: uv sync --frozen --all-groups --all-extras
- name: check_licenses
run: uv run --no-sync python ./tests/code_coverage_tests/check_licenses.py
- name: check_provider_folders_documented
run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py
- name: check_prisma_binary_cache
run: uv run --no-sync python ./tests/code_coverage_tests/check_prisma_binary_cache.py
- name: check_workflow_startup_safety
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py
- name: check_workflow_job_name_collisions
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_job_name_collisions.py
- name: test_workflow_job_name_collisions
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_workflow_job_name_collisions.py
- name: test_unit_passed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_unit_passed_gate.py
- name: test_e2e_changed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py
- name: test_e2e_metadata
env:
PYTHONPATH: tests/e2e
run: |
export LITELLM_MASTER_KEY="sk-$(openssl rand -hex 16)"
uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_metadata.py tests/code_coverage_tests/test_e2e_junit_report.py
- name: Check merge smoke harness
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_merge_smoke.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
- name: test_chat_completion_imports
run: uv run --no-sync python ./tests/code_coverage_tests/test_chat_completion_imports.py
- name: info_log_check
run: uv run --no-sync python ./tests/code_coverage_tests/info_log_check.py
- name: check_guardrail_apply_decorator
run: uv run --no-sync python ./tests/code_coverage_tests/check_guardrail_apply_decorator.py
- name: test_ban_set_verbose
run: uv run --no-sync python ./tests/code_coverage_tests/test_ban_set_verbose.py
- name: code_qa_check_tests
run: uv run --no-sync python ./tests/code_coverage_tests/code_qa_check_tests.py
- name: check_get_model_cost_key_performance
run: uv run --no-sync python ./tests/code_coverage_tests/check_get_model_cost_key_performance.py
- name: test_proxy_types_import
run: uv run --no-sync python ./tests/code_coverage_tests/test_proxy_types_import.py
- name: callback_manager_test
run: uv run --no-sync python ./tests/code_coverage_tests/callback_manager_test.py
- name: recursive_detector
run: uv run --no-sync python ./tests/code_coverage_tests/recursive_detector.py
- name: test_router_strategy_async
run: uv run --no-sync python ./tests/code_coverage_tests/test_router_strategy_async.py
- name: litellm_logging_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/litellm_logging_code_coverage.py
- name: ensure_async_clients_test
run: uv run --no-sync python ./tests/code_coverage_tests/ensure_async_clients_test.py
- name: enforce_llms_folder_style
run: uv run --no-sync python ./tests/code_coverage_tests/enforce_llms_folder_style.py
- name: prevent_key_leaks_in_exceptions
run: uv run --no-sync python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py
- name: check_unsafe_enterprise_import
run: uv run --no-sync python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py
- name: ban_copy_deepcopy_kwargs
run: uv run --no-sync python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py
- name: check_fastuuid_usage
run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py
- name: check_py310_typing_imports
run: uv run --no-sync python ./tests/code_coverage_tests/check_py310_typing_imports.py
- name: check_e2e_no_raw_requests
run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py
- name: check_migrations_no_data_rewrites
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
- name: check_no_publicly_known_master_key
run: uv run --no-sync python ./tests/code_coverage_tests/check_no_publicly_known_master_key.py
- name: test_check_no_publicly_known_master_key
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_check_no_publicly_known_master_key.py
- name: check_unbounded_in_lists (fails on findings not in the baseline)
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
- name: documentation_test_env_keys
run: uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
- name: documentation_test_router_settings
run: uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
- name: documentation_test_api_docs
run: uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
python-310-import-smoke:
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.10"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: uv sync --frozen --extra proxy --extra cli --python 3.10
- run: uv run --no-sync python --version
- name: Import litellm
run: uv run --no-sync python -c "import litellm"
- name: Check litellm CLI
run: uv run --no-sync litellm --version
- name: Check lite CLI
run: uv run --no-sync lite version

View file

@ -1,4 +1,4 @@
name: LiteLLM Linting
name: lint
on:
pull_request:
@ -14,22 +14,16 @@ concurrency:
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
lint:
runs-on: ubuntu-latest
timeout-minutes: 15
# actions: read lets the four lint gates download the base-counts artifacts
# published by publish-lint-base-counts.yml instead of re-scanning the
# merge-base tree in a throwaway worktree.
python:
name: python
permissions:
contents: read
pull-requests: read
actions: read
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
# Check out the PR head, not the default refs/pull/N/merge: the merge ref
# folds in newer base commits, which the diff-based gates (ruff delta,
# Any-discipline) would otherwise blame on this branch.
with:
ref: ${{ github.event.pull_request.head.sha }}
fetch-depth: 1
@ -100,9 +94,6 @@ jobs:
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
# DB wrappers typed against the generated client would degrade to Unknown.
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
run: |
@ -211,13 +202,11 @@ jobs:
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
secret-scan:
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 5
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
@ -249,3 +238,418 @@ jobs:
else
echo "GITGUARDIAN_API_KEY not set, skipping ggshield scan"
fi
ui:
name: ui
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 8
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Collect changed files
id: changed
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
# base.sha is the base branch tip from when the PR was opened, while
# actions/checkout leaves HEAD on a merge of the PR into the *current*
# base tip. "$BASE_SHA"...HEAD therefore spans every base-branch commit
# landed since, so a PR that touches no UI file still gets linted
# against hundreds of other people's files. Diff the PR head against its
# own merge base instead, which is exactly what this PR changed.
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
[ -f "$f" ] || continue
case "$f" in
*.js | *.jsx | *.ts | *.tsx | *.mjs | *.cjs)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt"
printf '%s\n' "$f" >> "$RUNNER_TEMP/eslint_files.txt" ;;
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$merge_base" "$HEAD_SHA" -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else
echo "has_files=false" >> "$GITHUB_OUTPUT"
echo "No lintable UI files changed in this PR; nothing to check."
fi
- name: Setup Node.js
if: steps.changed.outputs.has_files == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changed.outputs.has_files == 'true'
run: npm ci
- name: Lint changed files (prettier + eslint)
if: steps.changed.outputs.has_files == 'true'
run: |
prettier_files=()
eslint_files=()
while IFS= read -r f; do prettier_files+=("$f"); done < "$RUNNER_TEMP/prettier_files.txt"
while IFS= read -r f; do eslint_files+=("$f"); done < "$RUNNER_TEMP/eslint_files.txt"
status=0
if [ ${#prettier_files[@]} -gt 0 ]; then
echo "::group::Prettier (${#prettier_files[@]} files)"
npx prettier --check "${prettier_files[@]}" || { status=1; echo "::error::Unformatted files. Fix with: npm run format"; }
echo "::endgroup::"
fi
if [ ${#eslint_files[@]} -gt 0 ]; then
echo "::group::ESLint (${#eslint_files[@]} files)"
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_files[@]}" || status=1
echo "::endgroup::"
fi
exit $status
- name: Check lint budgets
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: |
npx eslint . -f json -o "$RUNNER_TEMP/lint-report.json" || true
node scripts/check-lint-budgets.mjs "$RUNNER_TEMP/lint-report.json" eslint-budgets.json
- name: Check for dead code (knip)
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: npm run knip:ci
ui-api-types:
name: ui-api-types
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
fetch-depth: 2
- name: Detect changes that can affect the generated types
id: changes
run: |
set -euo pipefail
if ! base="$(git rev-parse --verify --quiet HEAD^2 >/dev/null && git rev-parse HEAD^1)"; then
echo "Not a pull request merge commit, running the full check."
echo "relevant=true" >> "$GITHUB_OUTPUT"
exit 0
fi
files="$(git diff --name-only "$base" HEAD)"
if grep -Eq '^(litellm/(proxy|types)/|ui/litellm-dashboard/(src/lib/http/schema\.d\.ts|scripts/gen-api-types\.mjs|package(-lock)?\.json)$|\.github/workflows/test-linting\.yml$)' <<< "$files"; then
echo "relevant=true" >> "$GITHUB_OUTPUT"
else
echo "No proxy, types or generator changes in this pull request, nothing to verify."
echo "relevant=false" >> "$GITHUB_OUTPUT"
fi
- name: Set up Python
if: steps.changes.outputs.relevant == 'true'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.relevant == 'true'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache the Rust build
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-cargo-build
- name: Install backend dependencies
if: steps.changes.outputs.relevant == 'true'
run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.relevant == 'true'
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Regenerate the lazy OpenAPI snapshot
if: steps.changes.outputs.relevant == 'true'
run: uv run --no-sync python -m litellm.proxy._lazy_openapi_snapshot
- name: Fail if the lazy OpenAPI snapshot is stale
if: steps.changes.outputs.relevant == 'true'
run: |
if ! git diff --exit-code -- litellm/proxy/_lazy_openapi_snapshot.json; then
echo "::error file=litellm/proxy/_lazy_openapi_snapshot.json::The lazy OpenAPI snapshot is out of sync with the lazily loaded routes."
echo ""
echo "A lazily loaded route or model changed without regenerating the snapshot that /openapi.json serves for unloaded features."
echo "To fix, run from the repo root:"
echo " uv run python -m litellm.proxy._lazy_openapi_snapshot"
echo "then run npm run gen:api from ui/litellm-dashboard and commit both files."
exit 1
fi
echo "_lazy_openapi_snapshot.json is in sync with the lazily loaded routes."
- name: Set up Node.js
if: steps.changes.outputs.relevant == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dashboard dependencies
if: steps.changes.outputs.relevant == 'true'
working-directory: ui/litellm-dashboard
run: npm ci
- name: Regenerate types from the live spec
if: steps.changes.outputs.relevant == 'true'
working-directory: ui/litellm-dashboard
env:
LITELLM_PYTHON: "uv run --no-sync python"
run: npm run gen:api
- name: Fail if types are stale
if: steps.changes.outputs.relevant == 'true'
run: |
if ! git diff --exit-code -- ui/litellm-dashboard/src/lib/http/schema.d.ts; then
echo "::error file=ui/litellm-dashboard/src/lib/http/schema.d.ts::Generated API types are out of sync with the proxy OpenAPI spec."
echo ""
echo "A backend route or model changed without regenerating the dashboard types."
echo "To fix, run from ui/litellm-dashboard:"
echo " npm run gen:api"
echo "then commit the updated src/lib/http/schema.d.ts."
exit 1
fi
echo "schema.d.ts is in sync with the proxy OpenAPI spec."
code-quality:
name: code-quality
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
run: uv sync --frozen --all-groups --all-extras
- name: check_licenses
run: uv run --no-sync python ./tests/code_coverage_tests/check_licenses.py
- name: check_provider_folders_documented
run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py
- name: check_prisma_binary_cache
run: uv run --no-sync python ./tests/code_coverage_tests/check_prisma_binary_cache.py
- name: check_workflow_startup_safety
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py
- name: check_workflow_job_name_collisions
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_job_name_collisions.py
- name: test_workflow_job_name_collisions
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_workflow_job_name_collisions.py
- name: test_unit_passed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_unit_passed_gate.py
- name: test_e2e_changed_gate
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py
- name: test_e2e_metadata
env:
PYTHONPATH: tests/e2e
run: |
export LITELLM_MASTER_KEY="sk-$(openssl rand -hex 16)"
uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_metadata.py tests/code_coverage_tests/test_e2e_junit_report.py
- name: Check merge smoke harness
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_merge_smoke.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
- name: test_chat_completion_imports
run: uv run --no-sync python ./tests/code_coverage_tests/test_chat_completion_imports.py
- name: info_log_check
run: uv run --no-sync python ./tests/code_coverage_tests/info_log_check.py
- name: check_guardrail_apply_decorator
run: uv run --no-sync python ./tests/code_coverage_tests/check_guardrail_apply_decorator.py
- name: test_ban_set_verbose
run: uv run --no-sync python ./tests/code_coverage_tests/test_ban_set_verbose.py
- name: code_qa_check_tests
run: uv run --no-sync python ./tests/code_coverage_tests/code_qa_check_tests.py
- name: check_get_model_cost_key_performance
run: uv run --no-sync python ./tests/code_coverage_tests/check_get_model_cost_key_performance.py
- name: test_proxy_types_import
run: uv run --no-sync python ./tests/code_coverage_tests/test_proxy_types_import.py
- name: callback_manager_test
run: uv run --no-sync python ./tests/code_coverage_tests/callback_manager_test.py
- name: recursive_detector
run: uv run --no-sync python ./tests/code_coverage_tests/recursive_detector.py
- name: test_router_strategy_async
run: uv run --no-sync python ./tests/code_coverage_tests/test_router_strategy_async.py
- name: litellm_logging_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/litellm_logging_code_coverage.py
- name: ensure_async_clients_test
run: uv run --no-sync python ./tests/code_coverage_tests/ensure_async_clients_test.py
- name: enforce_llms_folder_style
run: uv run --no-sync python ./tests/code_coverage_tests/enforce_llms_folder_style.py
- name: prevent_key_leaks_in_exceptions
run: uv run --no-sync python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py
- name: check_unsafe_enterprise_import
run: uv run --no-sync python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py
- name: ban_copy_deepcopy_kwargs
run: uv run --no-sync python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py
- name: check_fastuuid_usage
run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py
- name: check_py310_typing_imports
run: uv run --no-sync python ./tests/code_coverage_tests/check_py310_typing_imports.py
- name: check_e2e_no_raw_requests
run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py
- name: check_migrations_no_data_rewrites
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
- name: check_no_publicly_known_master_key
run: uv run --no-sync python ./tests/code_coverage_tests/check_no_publicly_known_master_key.py
- name: test_check_no_publicly_known_master_key
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_check_no_publicly_known_master_key.py
- name: check_unbounded_in_lists (fails on findings not in the baseline)
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
- name: documentation_test_env_keys
run: uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
- name: documentation_test_router_settings
run: uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
- name: documentation_test_api_docs
run: uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
ci-coverage:
name: ci-coverage
permissions:
contents: read
runs-on: ubuntu-latest
timeout-minutes: 5
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Assert every test file and Dockerfile is invoked by a job
run: |
python -m pip install "pyyaml==6.0.3"
python .github/scripts/assert_ci_coverage.py
- name: Assert no -k expression deselects a file from every job that globs it
run: python .github/scripts/assert_ci_coverage.py --slices
- name: Assert .github/workflows/ holds only workflows, correctly named
run: python .github/scripts/assert_workflow_dir_hygiene.py
lint-passed:
name: lint passed
needs: [python, secret-scan, ui, ui-api-types, code-quality, ci-coverage]
if: always()
runs-on: ubuntu-latest
timeout-minutes: 2
permissions: {}
steps:
- name: Require every lint job to succeed
env:
NEEDS: ${{ toJSON(needs) }}
run: |
jq -r 'to_entries[] | "\(.key): \(.value.result)"' <<< "$NEEDS"
jq -e 'all(.[]; .result == "success")' <<< "$NEEDS" > /dev/null

View file

@ -1,105 +0,0 @@
name: UI Lint
permissions:
contents: read
on:
pull_request:
branches:
- main
- "litellm_**"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
frontend-lint:
runs-on: ubuntu-latest
timeout-minutes: 8
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Collect changed files
id: changed
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
# base.sha is the base branch tip from when the PR was opened, while
# actions/checkout leaves HEAD on a merge of the PR into the *current*
# base tip. "$BASE_SHA"...HEAD therefore spans every base-branch commit
# landed since, so a PR that touches no UI file still gets linted
# against hundreds of other people's files. Diff the PR head against its
# own merge base instead, which is exactly what this PR changed.
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
[ -f "$f" ] || continue
case "$f" in
*.js | *.jsx | *.ts | *.tsx | *.mjs | *.cjs)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt"
printf '%s\n' "$f" >> "$RUNNER_TEMP/eslint_files.txt" ;;
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$merge_base" "$HEAD_SHA" -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else
echo "has_files=false" >> "$GITHUB_OUTPUT"
echo "No lintable UI files changed in this PR; nothing to check."
fi
- name: Setup Node.js
if: steps.changed.outputs.has_files == 'true'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changed.outputs.has_files == 'true'
run: npm ci
- name: Lint changed files (prettier + eslint)
if: steps.changed.outputs.has_files == 'true'
run: |
prettier_files=()
eslint_files=()
while IFS= read -r f; do prettier_files+=("$f"); done < "$RUNNER_TEMP/prettier_files.txt"
while IFS= read -r f; do eslint_files+=("$f"); done < "$RUNNER_TEMP/eslint_files.txt"
status=0
if [ ${#prettier_files[@]} -gt 0 ]; then
echo "::group::Prettier (${#prettier_files[@]} files)"
npx prettier --check "${prettier_files[@]}" || { status=1; echo "::error::Unformatted files. Fix with: npm run format"; }
echo "::endgroup::"
fi
if [ ${#eslint_files[@]} -gt 0 ]; then
echo "::group::ESLint (${#eslint_files[@]} files)"
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_files[@]}" || status=1
echo "::endgroup::"
fi
exit $status
- name: Check lint budgets
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: |
npx eslint . -f json -o "$RUNNER_TEMP/lint-report.json" || true
node scripts/check-lint-budgets.mjs "$RUNNER_TEMP/lint-report.json" eslint-budgets.json
- name: Check for dead code (knip)
if: ${{ !cancelled() && steps.changed.outputs.has_files == 'true' }}
run: npm run knip:ci

View file

@ -1,99 +0,0 @@
name: UI Unit Tests
permissions:
contents: read
pull-requests: read
on:
pull_request:
branches:
- main
- "litellm_**"
push:
branches:
- main
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
ui-unit-tests:
runs-on: ubuntu-latest-16-cores
timeout-minutes: 20
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Check UI production source types
if: steps.changes.outputs.decision != 'skip'
run: npm run typecheck
- name: Run UI type tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
run: npm run test:types
- name: Run UI unit tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
full_suite() { npm run test -- --run --pool forks --maxWorkers=14; }
if [ -z "$BASE_SHA" ]; then
echo "Push to $GITHUB_REF_NAME: running the full suite"
full_suite
exit 0
fi
merge_base=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$merge_base"
git fetch --no-tags --depth=1 origin "$merge_base" "$HEAD_SHA"
changed_files=()
while IFS= read -r f; do
changed_files+=("$f")
done < <(git diff --name-only --relative "$merge_base" "$HEAD_SHA" -- .)
if [ ${#changed_files[@]} -eq 0 ]; then
echo "No UI files changed in this PR; skipping unit tests."
exit 0
fi
scope=$(printf '%s\n' "${changed_files[@]}" | bash "$GITHUB_WORKSPACE/.github/scripts/select_ui_test_scope.sh")
if [ "$scope" != related ]; then
echo "Pull request: ${#changed_files[@]} changed UI files reach outside src/, so related would miss their dependents; running the full suite"
full_suite
exit 0
fi
echo "Pull request: running tests related to ${#changed_files[@]} changed UI files"
npm run test -- related "${changed_files[@]}" --run --passWithNoTests \
--pool forks --maxWorkers=14

View file

@ -1,4 +1,4 @@
name: Merge smoke checks
name: smoke
on:
pull_request:
@ -14,7 +14,7 @@ concurrency:
jobs:
dashboard-build:
name: Dashboard build
name: dashboard-build
runs-on: ubuntu-24.04
timeout-minutes: 30
steps:
@ -27,7 +27,7 @@ jobs:
run: docker build --target ui-builder -f Dockerfile .
core-checks:
name: Core checks (Python ${{ matrix.python-version }})
name: py${{ matrix.python-version }}
runs-on: ubuntu-24.04
timeout-minutes: 30
strategy:
@ -79,9 +79,6 @@ jobs:
- name: Proxy startup check
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py proxy-startup --diagnostics-dir "${RUNNER_TEMP}/smoke-diagnostics"
- name: Run curated smoke cases
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py pytest --manifest .github/merge-smoke-tests.json
- name: Upload smoke diagnostics
if: always()
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
@ -93,3 +90,18 @@ jobs:
- name: Remove the network namespace
if: always()
run: sudo ip netns delete smoke
smoke-passed:
name: smoke passed
needs: [dashboard-build, core-checks]
if: always()
runs-on: ubuntu-latest
timeout-minutes: 2
permissions: {}
steps:
- name: Require every smoke job to succeed
env:
NEEDS: ${{ toJSON(needs) }}
run: |
jq -r 'to_entries[] | "\(.key): \(.value.result)"' <<< "$NEEDS"
jq -e 'all(.[]; .result == "success")' <<< "$NEEDS" > /dev/null

View file

@ -1,103 +0,0 @@
name: "Unit Tests: Documentation Validation"
on:
pull_request:
branches:
- main
- "litellm_**"
push:
branches:
- main
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
documentation:
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
pull-requests: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
if: steps.changes.outputs.decision != 'skip'
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip' && github.ref == 'refs/heads/main'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip' && github.ref != 'refs/heads/main'
uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Cache the Rust build
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
# Run the same documentation tests that CircleCI ran (as direct Python scripts)
- name: Run documentation validation tests
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py

File diff suppressed because it is too large Load diff

View file

@ -137,7 +137,7 @@ format-check: install-dev
lint-fetch-base:
@$(RESOLVE_BASE)
# Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated
# Mirror test-linting.yml's python job environment: the proxy-dev group plus a generated
# Prisma client, so `basedpyright tests/e2e` resolves the same modules CI does. The
# basedpyright gate itself no longer measures here (scripts/type_check_gate.py provisions its
# own .venv-typecheck). --inexact tops up the venv instead of pruning the proxy extras
@ -236,7 +236,7 @@ check-circular-imports: $(LINT_DEP_INSTALL)
check-import-safety: $(LINT_DEP_INSTALL)
@$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
# Combined linting, isomorphic to test-linting.yml's lint job so a local pass means a
# Combined linting, isomorphic to test-linting.yml's python job so a local pass means a
# green CI lint: it installs the same env (proxy-dev + generated Prisma client) and then
# runs the diff-scoped ruff format check, whole-tree ruff check, the strict-rule /
# type-discipline / basedpyright gates as a delta vs the base, then the circular-import
@ -260,8 +260,7 @@ lint-dev: lint-format-changed check-circular-imports check-import-safety
# is staged (warning about changed files left unstaged); with nothing staged it falls
# back to the working tree's diff against the merge base with the base branch, so a
# fresh merge commit or an unstaged working tree still gets checked. Mirrors
# test-linting.yml (Python), test-litellm-ui-build.yml's frontend-lint (dashboard), and
# check-ui-api-types.yml (API-type drift), skipping any whose files aren't in scope.
# test-linting.yml (Python, UI, and API types), skipping any whose files aren't in scope.
# Not auto-installed as a git hook so it never slows an unrelated human commit.
check:
@$(GATE_SLOT_LOCK) $(MAKE) check-inner

View file

@ -6,7 +6,7 @@ codecov:
ignore:
- "litellm-rust/**"
# Uploads are flagged per workflow/shard (GHA) or "circleci". carryforward makes
# Uploads are flagged per workflow tier (GHA) or "circleci". carryforward makes
# a re-upload of a flag replace its prior session instead of accumulating a
# conflicting one, and lets a commit reuse a flag from its parent when that flag
# was not re-uploaded. Required because the same commit can receive the
@ -27,6 +27,76 @@ flag_management:
carryforward: false
- name: circleci
carryforward: false
- name: core-utils
carryforward: false
- name: enterprise-routing
carryforward: false
- name: integrations
carryforward: false
- name: llm-vertex-ai
carryforward: false
- name: llm-other-providers
carryforward: false
- name: llm-openai-meta
carryforward: false
- name: misc
carryforward: false
- name: misc-dirs
carryforward: false
- name: proxy-auth
carryforward: false
- name: proxy-hooks-client
carryforward: false
- name: proxy-endpoints
carryforward: false
- name: proxy-feature-endpoints
carryforward: false
- name: proxy-server
carryforward: false
- name: mcp-elicitation
carryforward: false
- name: proxy-infra
carryforward: false
- name: proxy-infra-root
carryforward: false
- name: caching-local
carryforward: false
- name: proxy-extras
carryforward: false
- name: enterprise-package
carryforward: false
- name: enterprise-managed-files
carryforward: false
- name: responses-caching-types
carryforward: false
- name: lens-python-310
carryforward: false
- name: proxy-db-key-generation
carryforward: false
- name: proxy-db-auth-checks
carryforward: false
- name: proxy-db-jwt-and-keys
carryforward: false
- name: proxy-db-proxy-utils
carryforward: false
- name: proxy-db-proxy-server-core
carryforward: false
- name: proxy-db-proxy-runtime
carryforward: false
- name: proxy-db-mcp-oauth
carryforward: false
- name: proxy-db-custom-logging
carryforward: false
- name: proxy-db-logging-misc
carryforward: false
- name: proxy-db-db-and-spend
carryforward: false
- name: proxy-db-guardrails-hooks
carryforward: false
- name: proxy-db-budgets
carryforward: false
- name: proxy-db-endpoints-and-responses
carryforward: false
component_management:
individual_components:

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_CredentialsTable" ADD COLUMN IF NOT EXISTS "display_name" TEXT;

View file

@ -42,6 +42,7 @@ model LiteLLM_BudgetTable {
model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
display_name String?
credential_values Json
credential_info Json?
created_at DateTime @default(now()) @map("created_at")

View file

@ -680,6 +680,10 @@ EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE: Final = float(
ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01")
ANTHROPIC_SKILLS_API_BETA_VERSION: Final = "skills-2025-10-02"
ANTHROPIC_BATCHES_ROUTE: Final = "/v1/messages/batches"
ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX: Final = 1568
ANTHROPIC_IMAGE_MAX_PIXELS: Final = 1_150_000
ANTHROPIC_IMAGE_PIXELS_PER_TOKEN: Final = 750
PDF_DATA_URL_PREFIX: Final = "data:application/pdf;base64,"
VERTEX_BATCH_PREDICTION_JOBS_ROUTE: Final = "batchPredictionJobs"
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES: Final = {
"low": 1,

View file

@ -20,6 +20,7 @@ import httpx2
from httpx2._client import UseClientDefault
from httpx2._types import AuthTypes
from mcp import ClientSession, MCPError, ReadResourceResult, Resource, StdioServerParameters
from mcp.client._input_required import run_input_required_driver
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamable_http_client
@ -49,6 +50,7 @@ from mcp.types import (
InitializeRequestParams,
InitializeResult,
InputRequiredResult,
InputResponses,
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
@ -100,6 +102,7 @@ from litellm.types.mcp import (
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult
from litellm.proxy._experimental.mcp_server.legacy_callbacks import ElicitationCallback
def to_basic_auth(auth_value: str) -> str:
@ -411,7 +414,7 @@ class MCPClient:
aws_auth: httpx2.Auth | None = None,
resolved_auth: httpx2.Auth | None = None,
sampling_callback: Callable | None = None,
elicitation_callback: Callable | None = None,
elicitation_callback: "ElicitationCallback | None" = None,
logging_callback: Callable | None = None,
protocol_version: MCPUpstreamProtocol = "auto",
):
@ -438,7 +441,7 @@ class MCPClient:
self._resolved_auth: httpx2.Auth | None = resolved_auth
self._last_initialize_instructions: str | None = None
self._sampling_callback: Callable | None = sampling_callback
self._elicitation_callback: Callable | None = elicitation_callback
self._elicitation_callback: ElicitationCallback | None = elicitation_callback
self._logging_callback: Callable | None = logging_callback
# handle the basic auth value if provided
if auth_value:
@ -631,15 +634,6 @@ class MCPClient:
# The SDK closes pending requests when its message handler raises.
raise RuntimeError("MCP response stream failed")
session_kwargs: Final = {
name: callback
for name, callback in (
("sampling_callback", self._sampling_callback),
("elicitation_callback", self._elicitation_callback),
("logging_callback", self._logging_callback),
)
if callback is not None
}
# The SDK drops a response stream that ends without a JSON-RPC reply, so nothing else
# ever fails the request.
session_ctx: Final = ClientSession(
@ -647,7 +641,9 @@ class MCPClient:
write_stream,
read_timeout_seconds=self.timeout,
message_handler=receive_message,
**session_kwargs,
sampling_callback=self._sampling_callback,
elicitation_callback=self._elicitation_callback,
logging_callback=self._logging_callback,
)
session: Final = await session_ctx.__aenter__()
try:
@ -948,6 +944,28 @@ class MCPClient:
"""The error result ``call_tool`` returns when it swallows a failure (no re-execution)."""
return error_text_result(exc)
async def _request_with_interaction(
self,
session: ClientSession,
request: Callable[[InputResponses | None, str | None], Awaitable[TSessionResult | InputRequiredResult]],
input_responses: InputResponses | None,
request_state: str | None,
allow_input_required: bool,
) -> TSessionResult | InputRequiredResult:
from litellm.proxy._experimental.mcp_server.contracts import ClientInteraction
from litellm.proxy._experimental.mcp_server.interactions import LegacyClientInteraction, ModernClientInteraction
with anyio.fail_after(self.timeout):
first: Final = await request(input_responses, request_state)
if allow_input_required:
return await ModernClientInteraction(
session, allow_elicitation=self._elicitation_callback is not None
).complete(first, request)
if not isinstance(first, InputRequiredResult):
return first
interaction: Final[ClientInteraction] = LegacyClientInteraction(session)
return await run_input_required_driver(first, dispatch=interaction.request, retry=request)
async def call_tool(
self,
call_tool_request_params: MCPCallToolRequestParams,
@ -990,11 +1008,25 @@ class MCPClient:
)
if not any(tool.name == call_tool_request_params.name for tool in tools):
raise MCPError(code=-32603, message="Tool schema is unavailable from the bounded upstream catalog")
return await session.call_tool(
name=call_tool_request_params.name,
arguments=call_tool_request_params.arguments,
progress_callback=on_progress,
allow_input_required=allow_input_required,
async def request(
responses: InputResponses | None, state: str | None
) -> MCPCallToolResult | InputRequiredResult:
return await session.call_tool(
name=call_tool_request_params.name,
arguments=call_tool_request_params.arguments,
input_responses=responses,
request_state=state,
progress_callback=on_progress,
allow_input_required=True,
)
return await self._request_with_interaction(
session,
request,
call_tool_request_params.input_responses,
call_tool_request_params.request_state,
allow_input_required,
)
try:
@ -1129,15 +1161,32 @@ class MCPClient:
# Return empty list instead of raising to allow graceful degradation
return ListPromptsResult(prompts=[])
async def get_prompt(self, get_prompt_request_params: GetPromptRequestParams) -> GetPromptResult:
async def get_prompt(
self, get_prompt_request_params: GetPromptRequestParams, *, allow_input_required: bool = False
) -> GetPromptResult | InputRequiredResult:
"""Fetch a prompt definition from the MCP server."""
verbose_logger.info("MCP client fetching prompt '%s'", get_prompt_request_params.name)
async def _get_prompt_operation(session: ClientSession):
verbose_logger.debug("MCP client sending get_prompt request to session")
return await session.get_prompt(
name=get_prompt_request_params.name,
arguments=get_prompt_request_params.arguments,
async def request(
responses: InputResponses | None, state: str | None
) -> GetPromptResult | InputRequiredResult:
return await session.get_prompt(
name=get_prompt_request_params.name,
arguments=get_prompt_request_params.arguments,
input_responses=responses,
request_state=state,
allow_input_required=True,
)
return await self._request_with_interaction(
session,
request,
get_prompt_request_params.input_responses,
get_prompt_request_params.request_state,
allow_input_required,
)
try:
@ -1285,13 +1334,37 @@ class MCPClient:
# Return empty list instead of raising to allow graceful degradation
return ListResourceTemplatesResult(resource_templates=[])
async def read_resource(self, url: AnyUrl) -> ReadResourceResult:
async def read_resource(
self,
url: AnyUrl,
*,
input_responses: InputResponses | None = None,
request_state: str | None = None,
allow_input_required: bool = False,
) -> ReadResourceResult | InputRequiredResult:
"""Fetch resource contents from the MCP server."""
verbose_logger.info("MCP client fetching resource '%s'", url)
async def _read_resource_operation(session: ClientSession):
verbose_logger.debug("MCP client sending read_resource request to session")
return await session.read_resource(str(url))
async def request(
responses: InputResponses | None, state: str | None
) -> ReadResourceResult | InputRequiredResult:
return await session.read_resource(
str(url),
input_responses=responses,
request_state=state,
allow_input_required=True,
)
return await self._request_with_interaction(
session,
request,
input_responses,
request_state,
allow_input_required,
)
try:
read_resource_result: Final = await self.run_with_session(_read_resource_operation)

View file

@ -20,6 +20,7 @@ overwrite each other within the same day, producing incomplete data.
from __future__ import annotations
import os
from collections.abc import Callable
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final
@ -28,6 +29,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import MAVVRIK_FOCUS_EXPORT_JOB_NAME
from litellm.integrations.focus.destinations.base import FocusTimeWindow
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.utils import get_utc_datetime
if TYPE_CHECKING:
from apscheduler.schedulers.asyncio import AsyncIOScheduler
@ -83,7 +85,8 @@ def _is_empty_metrics_marker(marker: object | None) -> bool:
class MavvrikFocusLogger(FocusLogger):
"""FOCUS-based export logger that routes to the Mavvrik destination."""
def __init__(self, **kwargs: Any) -> None:
def __init__(self, *, clock: Callable[[], datetime] = get_utc_datetime, **kwargs: Any) -> None:
self._clock: Final = clock
frequency: Final = os.getenv("MAVVRIK_FOCUS_FREQUENCY", "daily").lower()
if frequency != "daily":
raise ValueError(
@ -174,7 +177,7 @@ class MavvrikFocusLogger(FocusLogger):
# metricsMarker may be a Unix timestamp (int/float) or an ISO date string.
marker: Final = await destination.get_metrics_marker()
now: Final = datetime.now(timezone.utc)
now: Final = self._clock()
yesterday: Final = now.replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(days=1)
last_ingested: Final = _parse_metrics_marker(marker)

View file

@ -31,7 +31,7 @@ Anthropic wire shape is built later by ``anthropic_messages_pt``.
from collections.abc import Iterator, Mapping, Sequence
from itertools import chain, groupby
from typing import Final, Literal, TypeAlias
from typing import Final, Literal, TypeAlias, TypeVar
from litellm.types.llms.anthropic import AnthropicMessagesSystemMessageParam, AnthropicSystemMessageContent
from litellm.types.llms.openai import (
@ -55,6 +55,7 @@ _RENDERED_ASSISTANT_PART_TYPES: Final = frozenset({"text", "server_tool_use"})
_THINKING_BLOCK_TYPES: Final = frozenset({"thinking", "redacted_thinking"})
_MessageKind: TypeAlias = Literal["system", "tool", "user", "other"]
_Message: Final = TypeVar("_Message")
_TextPart: TypeAlias = tuple[str, ChatCompletionCachedContent | None]
@ -96,8 +97,8 @@ def _kind(message: object) -> _MessageKind:
def split_leading_system_run(
messages: Sequence[AllMessageValues],
) -> tuple[tuple[AllMessageValues, ...], tuple[AllMessageValues, ...]]:
messages: Sequence[_Message],
) -> tuple[tuple[_Message, ...], tuple[_Message, ...]]:
"""Split ``messages`` into the leading run of system messages and everything after it."""
leading_count: Final = next(
(index for index, message in enumerate(messages) if not is_system_message(message)),
@ -160,9 +161,16 @@ def _anthropic_text_block(part: _TextPart) -> AnthropicSystemMessageContent:
return cached
def anthropic_system_blocks(run: Sequence[object]) -> tuple[AnthropicSystemMessageContent, ...]:
"""The top-level ``system`` blocks for a run of system messages: every non-empty text part in order,
each keeping its ``cache_control``, which is the shape the chat path sends for the leading run."""
parts: Final = chain.from_iterable(_text_parts(message) for message in run)
return tuple(_anthropic_text_block(part) for part in parts)
def anthropic_system_messages(message: object) -> tuple[AnthropicMessagesSystemMessageParam, ...]:
"""The Anthropic wire message for a system message, or nothing when it carries no text."""
blocks: Final = tuple(_anthropic_text_block(part) for part in _text_parts(message))
blocks: Final = anthropic_system_blocks((message,))
if not blocks:
return ()
wire: Final[AnthropicMessagesSystemMessageParam] = {

View file

@ -2,6 +2,7 @@
## Helper utilities for token counting
import base64
import io
import math
import struct
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from itertools import accumulate
@ -17,6 +18,9 @@ import litellm
from litellm import verbose_logger
from litellm._lazy_imports import get_default_encoding
from litellm.constants import (
ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX,
ANTHROPIC_IMAGE_MAX_PIXELS,
ANTHROPIC_IMAGE_PIXELS_PER_TOKEN,
DEFAULT_IMAGE_HEIGHT,
DEFAULT_IMAGE_TOKEN_COUNT,
DEFAULT_IMAGE_WIDTH,
@ -25,6 +29,7 @@ from litellm.constants import (
MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES,
MAX_TILE_HEIGHT,
MAX_TILE_WIDTH,
PDF_DATA_URL_PREFIX,
TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS,
TOKEN_COUNTER_MAX_CONCURRENT_COUNTS,
TOKEN_COUNTER_MAX_EXACT_CHARS,
@ -805,6 +810,49 @@ def _anthropic_image_source_data(
return ""
def _anthropic_rendered_page_image_tokens(width: float, height: float) -> int:
"""A PDF page is rasterized within Anthropic's image limits, then billed by its pixel area."""
if width <= 0 or height <= 0:
return 0
area_at_max_edge: Final = ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX**2 * min(width, height) / max(width, height)
return math.ceil(min(float(ANTHROPIC_IMAGE_MAX_PIXELS), area_at_max_edge) / ANTHROPIC_IMAGE_PIXELS_PER_TOKEN)
def _count_inline_pdf_tokens(data_url: str, count_function: TokenCounterFunction) -> int | None:
if not data_url.startswith(PDF_DATA_URL_PREFIX):
return None
try:
from pypdf import PdfReader
except ImportError:
verbose_logger.debug("pypdf is not installed, so the PDF document is priced like one image")
return None
try:
reader: Final = PdfReader(io.BytesIO(base64.b64decode(data_url[len(PDF_DATA_URL_PREFIX) :])))
return sum(
count_function(page.extract_text() or "")
+ _anthropic_rendered_page_image_tokens(float(page.mediabox.width), float(page.mediabox.height))
for page in reader.pages
)
except Exception as e:
verbose_logger.debug("Could not read the PDF document's pages (%s), so it is priced like one image", e)
return None
def _count_opaque_document_tokens(
data_url: str,
count_function: TokenCounterFunction,
use_default_image_token_count: bool,
) -> int:
pdf_tokens: Final = _count_inline_pdf_tokens(data_url, count_function)
if pdf_tokens is not None:
return pdf_tokens
return calculate_img_tokens(
data=data_url,
mode="auto",
use_default_image_token_count=use_default_image_token_count,
)
def _count_document_tokens(
document: ChatCompletionDocumentObject | AnthropicMessagesDocumentParam,
count_function: TokenCounterFunction,
@ -824,10 +872,8 @@ def _count_document_tokens(
return metadata_tokens + _count_content_list(
count_function, content, use_default_image_token_count, default_token_count
)
return metadata_tokens + calculate_img_tokens(
data=_anthropic_image_source_data(source),
mode="auto",
use_default_image_token_count=use_default_image_token_count,
return metadata_tokens + _count_opaque_document_tokens(
_anthropic_image_source_data(source), count_function, use_default_image_token_count
)
@ -844,11 +890,7 @@ def _count_file_tokens(
name_tokens: Final = count_function(filename) if isinstance(filename, str) and filename else 0
if not isinstance(file_data, str) or not file_data:
return name_tokens
return name_tokens + calculate_img_tokens(
data=file_data,
mode="auto",
use_default_image_token_count=use_default_image_token_count,
)
return name_tokens + _count_opaque_document_tokens(file_data, count_function, use_default_image_token_count)
def _count_anthropic_content(

View file

@ -11,11 +11,16 @@ from typing import Final
from pydantic import JsonValue, TypeAdapter
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
anthropic_system_blocks,
split_leading_system_run,
)
from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers
from litellm.llms.anthropic.wif import resolve_anthropic_base
from litellm.types.llms.openai import ChatCompletionImageObject
_COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue])
_SYSTEM_BLOCKS: Final = TypeAdapter(list[JsonValue])
_IMAGE_BLOCK: Final = TypeAdapter(ChatCompletionImageObject)
COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config")
@ -48,6 +53,27 @@ def _count_content(content: JsonValue) -> JsonValue:
return [_count_block(block) for block in content] if isinstance(content, list) else content
def _lift_leading_system(
messages: Sequence[Mapping[str, JsonValue]], system: JsonValue
) -> tuple[tuple[Mapping[str, JsonValue], ...], JsonValue]:
"""Move the leading run of system-role messages into the top-level ``system`` parameter.
count_tokens only takes the initial system prompt there and answers 400 on ``role: "system"``
at the head of ``messages``; the chat path sends the same run as ``system``. A caller's own
``system`` keeps its place ahead of the lifted blocks, and a ``system`` that is neither text
nor a block list is left as sent, messages included, for the provider to judge.
"""
leading, conversation = split_leading_system_run(messages)
if not leading or not (system is None or isinstance(system, (str, list))):
return tuple(messages), system
lifted: Final = _SYSTEM_BLOCKS.validate_python(list(anthropic_system_blocks(leading)))
if isinstance(system, list):
return conversation, [*system, *lifted]
if isinstance(system, str) and system:
return conversation, [{"type": "text", "text": system}, *lifted]
return conversation, lifted or system
class AnthropicCountTokensConfig:
"""
Configuration and transformation logic for Anthropic CountTokens API.
@ -85,16 +111,24 @@ class AnthropicCountTokensConfig:
"""
Transform request to Anthropic CountTokens format.
Includes optional system and tools fields for accurate token counting.
Includes optional system and tools fields for accurate token counting; a leading run of
system-role messages is counted through ``system``, the only place count_tokens accepts it.
"""
options: Final[Mapping[str, JsonValue]] = optional_params or MappingProxyType({})
counted_messages, counted_system = _lift_leading_system(messages, system)
return _COUNT_REQUEST.validate_python(
MappingProxyType(
{
"model": model,
"messages": [{**message, "content": _count_content(message["content"])} for message in messages],
"messages": [
{**message, "content": _count_content(message["content"])} for message in counted_messages
],
**MappingProxyType(
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
{
key: value
for key, value in (("system", counted_system), ("tools", tools))
if value is not None
}
),
**MappingProxyType(
{key: value for key, value in options.items() if key in COUNT_TOKEN_OPTION_NAMES}

View file

@ -1319,6 +1319,7 @@ def _register_custom_pricing_for_request(
},
persist_across_reloads=False,
warning_display_name=shared_key,
custom_llm_provider=custom_llm_provider,
)

View file

@ -2999,6 +2999,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -3295,6 +3296,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -23585,6 +23587,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -32372,6 +32375,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -37547,6 +37551,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -47801,6 +47806,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -80595,17 +80601,25 @@
"global.openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3e-05,
"cache_creation_input_token_cost_ultrafast": 1.5e-05,
"cache_read_input_token_cost": 1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-06,
"cache_read_input_token_cost_ultrafast": 6e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.4e-05,
"input_cost_per_token_ultrafast": 1.2e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9e-05,
"output_cost_per_token_ultrafast": 6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
@ -80633,17 +80647,25 @@
"bedrock_mantle/openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.75e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3.3e-05,
"cache_creation_input_token_cost_ultrafast": 1.65e-05,
"cache_read_input_token_cost": 1.1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.32e-06,
"cache_read_input_token_cost_ultrafast": 6.6e-07,
"input_cost_per_token": 2.2e-06,
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.64e-05,
"input_cost_per_token_ultrafast": 1.32e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "responses",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9.9e-05,
"output_cost_per_token_ultrafast": 6.6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
@ -80672,17 +80694,25 @@
"us.openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.75e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3.3e-05,
"cache_creation_input_token_cost_ultrafast": 1.65e-05,
"cache_read_input_token_cost": 1.1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.32e-06,
"cache_read_input_token_cost_ultrafast": 6.6e-07,
"input_cost_per_token": 2.2e-06,
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.64e-05,
"input_cost_per_token_ultrafast": 1.32e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9.9e-05,
"output_cost_per_token_ultrafast": 6.6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",

View file

@ -6,14 +6,18 @@ layer; ``litellm.types.utils`` re-exports them for backwards compatibility.
"""
from collections.abc import Mapping
from typing import Literal, TypeAlias
from pydantic import Field, model_validator
from pydantic import ConfigDict, Field, model_validator
from litellm.types.llms.base import LiteLLMBaseModel
CredentialSource: TypeAlias = Literal["db", "config"]
class CredentialBase(LiteLLMBaseModel):
credential_name: str
display_name: str | None = None
credential_info: dict
@ -23,6 +27,14 @@ class CredentialItem(CredentialBase):
# edit rather than the credential, so it stays out of dumps: those feed config loading, the DB
# write, and the in-memory list, none of which have a place for it.
credential_values_to_delete: tuple[str, ...] | None = Field(default=None, exclude=True)
source: CredentialSource = Field(default="db", exclude=True)
class CredentialView(CredentialBase):
model_config = ConfigDict(frozen=True)
credential_values: Mapping[str, object]
source: CredentialSource
class CreateCredentialItem(CredentialBase):
@ -38,7 +50,8 @@ class CreateCredentialItem(CredentialBase):
class UpdateCredentialItem(LiteLLMBaseModel):
credential_name: str
credential_name: str | None = None
display_name: str | None = None
credential_info: Mapping[str, object]
credential_values: Mapping[str, object] | None = None
model_id: str | None = None

View file

@ -11,7 +11,13 @@ from mcp_types.methods import CLIENT_REQUESTS
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
from pydantic import TypeAdapter
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport
from litellm.types.mcp import (
MCP_LEGACY_VERSIONS,
MCPAdvertisedVersion,
MCPAdvertisedVersions,
MCPSpecVersion,
MCPTransport,
)
GATEWAY_OPERATIONS: Final = frozenset(
{
@ -46,14 +52,14 @@ REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
else frozenset({"complete", "input_required"}),
extensions=frozenset(),
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS or version.value == "2026-07-28",
)
for version in MCPSpecVersion
}
)
_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed)
TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2))
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPAdvertisedVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
def configured_versions() -> tuple[str, ...]:

View file

@ -7,6 +7,8 @@ from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias
from mcp.types import ErrorData, InputRequest, InputResponse, InputResponses
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -121,6 +123,10 @@ class ProgressCallback(Protocol):
async def __call__(self, progress: float, total: float | None, /) -> None: ...
class ClientInteraction(Protocol):
async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData: ...
@dataclass(frozen=True, slots=True)
class AuthorizedToolCall:
name: str
@ -130,3 +136,5 @@ class AuthorizedToolCall:
host_progress_callback: ProgressCallback | None
guardrail_context: Mapping[str, object] | None
logging_data: Mapping[str, object]
input_responses: InputResponses | None = None
request_state: str | None = None

View file

@ -0,0 +1,262 @@
import asyncio
import hashlib
import json
import secrets
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Final, Literal, TypeAlias, TypeVar
from uuid import uuid4
from mcp import MCPError
from mcp.client._input_required import run_input_required_driver
from mcp.client.session import ClientRequestContext, ClientSession
from mcp.types import (
CallToolRequest,
ElicitRequest,
ElicitRequestURLParams,
ErrorData,
GetPromptRequest,
InputRequest,
InputRequests,
InputRequiredResult,
InputResponse,
InputResponses,
ReadResourceRequest,
)
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.state_tokens import StateTokenError, open_state, seal_state
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@dataclass(frozen=True, slots=True)
class LegacyClientInteraction:
session: ClientSession
async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData:
context: Final = ClientRequestContext(
session=self.session, request_id=key, meta=request.params.meta if request.params else None
)
legacy_request: Final = (
request.model_copy(update={"params": request.params.model_copy(update={"elicitation_id": str(uuid4())})})
if isinstance(request, ElicitRequest)
and isinstance(request.params, ElicitRequestURLParams)
and request.params.elicitation_id is None
else request
)
return await self.session.dispatch_input_request(context, legacy_request)
InteractionOperation: TypeAlias = CallToolRequest | GetPromptRequest | ReadResourceRequest
_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
_PURPOSE: Final = "mcp:interaction:repeatable:v1"
class BoundInputRequiredResult(InputRequiredResult):
target_id: str | None = Field(default=None, exclude=True)
target_digest: str | None = Field(default=None, exclude=True)
gateway_responses: InputResponses | None = Field(default=None, exclude=True)
class ContinuationState(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
principal: str
operation: str
target_id: str
target_digest: str
upstream_state: str | None
gateway_responses: InputResponses | None = None
policy: Literal["repeatable"] = "repeatable"
expires_at: int
nonce: str
def _digest(value: JsonValue) -> str:
return hashlib.sha256(
json.dumps(value, sort_keys=True, separators=(",", ":"), allow_nan=False).encode()
).hexdigest()
def target_digest(server: MCPServer) -> str:
return _digest(
_JSON.validate_python(
{
"id": server.server_id,
"url": server.url,
"transport": server.transport,
"protocol": server.protocol_version,
"command": server.command,
"args": server.args,
}
)
)
def bind_target(result: InputRequiredResult, server: MCPServer) -> BoundInputRequiredResult:
bound: Final = (
result
if isinstance(result, BoundInputRequiredResult)
else BoundInputRequiredResult.model_validate(result.model_dump())
)
return bound.model_copy(update={"target_id": server.server_id, "target_digest": target_digest(server)})
def _principal(context: OperationContext) -> str:
caller: Final = context.user_api_key_auth
if caller is None or not caller.user_id:
raise MCPError(code=-32602, message="MCP continuations require an authenticated caller identity")
return _digest(
_JSON.validate_python(
{
"user": caller.user_id,
"team": caller.team_id,
"org": caller.org_id,
"end_user": caller.end_user_id,
"servers": sorted(context.mcp_servers) if context.mcp_servers is not None else None,
}
)
)
def _operation(operation: InteractionOperation) -> str:
return _digest(
_JSON.validate_python(
{
"method": operation.method,
"params": operation.params.model_dump(
mode="json", by_alias=True, exclude={"meta", "input_responses", "request_state"}
),
}
)
)
def _state_error(error: StateTokenError) -> MCPError:
return MCPError(
code=-32602,
message=(
"Set the same LITELLM_SALT_KEY on every replica to enable MCP continuations"
if error is StateTokenError.MISSING_KEY
else "Invalid or expired MCP continuation; start a fresh request"
),
)
def open_continuation(
operation: InteractionOperation, context: OperationContext, *, now: int
) -> ContinuationState | None:
token: Final = operation.params.request_state
if token is None:
if operation.params.input_responses:
raise MCPError(code=-32602, message="Input responses require a gateway continuation")
return None
opened: Final = open_state(token, purpose=_PURPOSE, now=now)
if isinstance(opened, Error):
raise _state_error(opened.error)
try:
state: Final = ContinuationState.model_validate(opened.ok)
except ValidationError as error:
raise _state_error(StateTokenError.INVALID) from error
if state.principal != _principal(context) or state.operation != _operation(operation) or state.expires_at <= now:
raise _state_error(StateTokenError.INVALID)
return state
def seal_continuation(
result: BoundInputRequiredResult,
operation: InteractionOperation,
context: OperationContext,
*,
now: int,
previous: ContinuationState | None = None,
) -> InputRequiredResult:
if result.target_id is None or result.target_digest is None:
raise MCPError(code=-32602, message="MCP continuation target is unavailable")
state: Final = ContinuationState(
principal=_principal(context),
operation=_operation(operation),
target_id=result.target_id,
target_digest=result.target_digest,
upstream_state=result.request_state,
gateway_responses=result.gateway_responses,
expires_at=previous.expires_at if previous is not None else now + 600,
nonce=previous.nonce if previous is not None else secrets.token_urlsafe(24),
)
sealed: Final = seal_state(
_JSON.validate_json(state.model_dump_json()), purpose=_PURPOSE, expires_at=state.expires_at, now=now
)
if isinstance(sealed, Error):
raise _state_error(sealed.error)
return InputRequiredResult(input_requests=result.input_requests, request_state=sealed.ok, _meta=result.meta)
_Terminal: Final = TypeVar("_Terminal")
@dataclass(frozen=True, slots=True)
class _DeferredInteraction:
result: BoundInputRequiredResult
@dataclass(frozen=True, slots=True)
class ModernClientInteraction:
session: ClientSession
allow_elicitation: bool
async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData:
if isinstance(request, ElicitRequest):
return ErrorData(code=-32602, message="Modern elicitation requires a continuation")
return await LegacyClientInteraction(self.session).request(key, request)
async def prepare(
self, result: _Terminal | InputRequiredResult
) -> _Terminal | InputRequiredResult | _DeferredInteraction:
if not isinstance(result, InputRequiredResult):
return result
pending: Final[InputRequests] = {
key: request for key, request in (result.input_requests or {}).items() if isinstance(request, ElicitRequest)
}
if pending and not self.allow_elicitation:
raise MCPError(code=-32602, message="Elicitation is disabled for this MCP server")
if result.input_requests and not pending:
return result
local: Final = tuple(
(key, request) for key, request in (result.input_requests or {}).items() if key not in pending
)
responses: Final = await asyncio.gather(*(self.request(key, request) for key, request in local))
for response in responses:
if isinstance(response, ErrorData):
raise MCPError(code=response.code, message=response.message)
return _DeferredInteraction(
BoundInputRequiredResult(
input_requests=pending or None,
request_state=result.request_state,
_meta=result.meta,
gateway_responses={
key: response for (key, _), response in zip(local, responses) if not isinstance(response, ErrorData)
}
or None,
)
)
async def complete(
self,
first: _Terminal | InputRequiredResult,
retry: Callable[[InputResponses | None, str | None], Awaitable[_Terminal | InputRequiredResult]],
) -> _Terminal | InputRequiredResult:
prepared: Final = await self.prepare(first)
if isinstance(prepared, _DeferredInteraction):
return prepared.result
if not isinstance(prepared, InputRequiredResult):
return prepared
async def resume(
responses: InputResponses | None, state: str | None
) -> _Terminal | InputRequiredResult | _DeferredInteraction:
return await self.prepare(await retry(responses, state))
completed: Final = await run_input_required_driver(prepared, dispatch=self.request, retry=resume)
return completed.result if isinstance(completed, _DeferredInteraction) else completed

View file

@ -24,7 +24,7 @@ class SamplingCallback(Protocol):
class ElicitationCallback(Protocol):
async def __call__(self, context: object, params: ElicitRequestParams, /) -> ElicitResult | ErrorData: ...
async def __call__(self, context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData: ...
def create_sampling_callback(
@ -79,6 +79,11 @@ def create_elicitation_callback(timeout: float | None = None) -> ElicitationCall
relay_timeout: Final = timeout if timeout is not None else MCP_CLIENT_TIMEOUT
async def callback(context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData:
if request is not None and request.protocol_version == "2026-07-28":
return ErrorData(
code=-32602,
message="A legacy upstream cannot resume input for a modern client; the operation may have partially completed",
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import handle_elicitation_request
return await handle_elicitation_request(

View file

@ -44,6 +44,7 @@ from mcp.types import (
GetPromptRequestParams,
GetPromptResult,
InputRequiredResult,
InputResponses,
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
@ -96,6 +97,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.interactions import bind_target
from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
@ -4888,7 +4890,10 @@ class MCPServerManager:
extra_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> ReadResourceResult:
input_responses: InputResponses | None = None,
request_state: str | None = None,
allow_input_required: bool = False,
) -> ReadResourceResult | InputRequiredResult:
"""Read resource contents from a specific MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
@ -4913,7 +4918,10 @@ class MCPServerManager:
user_api_key_auth=user_api_key_auth,
)
return await client.read_resource(url)
result: Final = await client.read_resource(
url, input_responses=input_responses, request_state=request_state, allow_input_required=allow_input_required
)
return bind_target(result, server) if isinstance(result, InputRequiredResult) else result
async def get_prompt_from_server(
self,
@ -4925,7 +4933,10 @@ class MCPServerManager:
extra_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> GetPromptResult:
input_responses: InputResponses | None = None,
request_state: str | None = None,
allow_input_required: bool = False,
) -> GetPromptResult | InputRequiredResult:
"""Fetch a specific prompt definition from a single MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
@ -4953,8 +4964,11 @@ class MCPServerManager:
get_prompt_request_params: Final = GetPromptRequestParams(
name=prompt_name,
arguments=arguments,
input_responses=input_responses,
request_state=request_state,
)
return await client.get_prompt(get_prompt_request_params)
result: Final = await client.get_prompt(get_prompt_request_params, allow_input_required=allow_input_required)
return bind_target(result, server) if isinstance(result, InputRequiredResult) else result
@staticmethod
def _is_same_authority_metadata_url(url: str, server_url: str) -> bool:
@ -6178,6 +6192,8 @@ class MCPServerManager:
user_api_key_auth: UserAPIKeyAuth | None = None,
client_ip: str | None = None,
allow_input_required: bool = False,
input_responses: InputResponses | None = None,
request_state: str | None = None,
) -> CallToolResult | InputRequiredResult:
"""
Call a regular MCP tool using the MCP client.
@ -6319,6 +6335,8 @@ class MCPServerManager:
call_tool_params: Final = MCPCallToolRequestParams(
name=original_tool_name,
arguments=arguments,
input_responses=input_responses,
request_state=request_state,
)
if _obo_retry_applies(mcp_server, subject_token):
@ -6425,7 +6443,11 @@ class MCPServerManager:
result: Final = mcp_responses[result_index]
self._remember_upstream_initialize_instructions(mcp_server, client)
return cast("CallToolResult | InputRequiredResult", result)
return (
bind_target(result, mcp_server)
if isinstance(result, InputRequiredResult)
else cast("CallToolResult", result)
)
def _resolve_mcp_server_for_tool_call(
self,
@ -6637,6 +6659,8 @@ class MCPServerManager:
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
input_responses: InputResponses | None = None,
request_state: str | None = None,
*,
catalog_auth_header: str | None | EllipsisType = ...,
listed_tool: MCPTool | None | EllipsisType = ...,
@ -6779,6 +6803,8 @@ class MCPServerManager:
hook_extra_headers=hook_result.get("extra_headers"),
user_api_key_auth=user_api_key_auth,
allow_input_required=wire_compat is WireCompat.MODERN,
input_responses=input_responses,
request_state=request_state,
)
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)

View file

@ -1,18 +1,19 @@
"""Shared MCP operation policy and dispatch."""
import asyncio
import time
import traceback
import types
import uuid
from collections.abc import Mapping, Sequence
from contextvars import ContextVar
from dataclasses import dataclass
from dataclasses import dataclass, replace
from datetime import datetime
from functools import partial
from typing import Any, Final, NoReturn, TypeAlias, overload
from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp import MCPError, ReadResourceResult, Resource
from mcp.types import (
CallToolRequest,
CallToolRequestParams,
@ -23,6 +24,7 @@ from mcp.types import (
GetPromptRequestParams,
GetPromptResult,
InputRequiredResult,
InputResponses,
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
@ -83,6 +85,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
classify_list_exception,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.interactions import (
BoundInputRequiredResult,
ContinuationState,
InteractionOperation,
open_continuation,
seal_continuation,
target_digest,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports
MCPServerManager,
_caller_authorization_fans_out, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
@ -1862,6 +1872,8 @@ async def execute_mcp_tool(
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
input_responses: InputResponses | None = None,
request_state: str | None = None,
**kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract
) -> CallToolResult | InputRequiredResult:
context: Final = prepare_context(
@ -1881,6 +1893,8 @@ async def execute_mcp_tool(
host_progress_callback=host_progress_callback,
guardrail_context=guardrail_context,
logging_data=types.MappingProxyType(kwargs),
input_responses=input_responses,
request_state=request_state,
)
return await GatewayOperations().execute(operation, context)
@ -1899,6 +1913,8 @@ async def _execute_mcp_tool(
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
input_responses: InputResponses | None = None,
request_state: str | None = None,
**kwargs: Any,
) -> CallToolResult | InputRequiredResult:
"""
@ -2194,6 +2210,8 @@ async def _execute_mcp_tool(
guardrail_context=guardrail_context,
host_progress_callback=host_progress_callback,
wire_compat=wire_compat,
input_responses=input_responses,
request_state=request_state,
)
# Fall back to local tool registry with original name (legacy support)
@ -2435,6 +2453,8 @@ async def call_mcp_tool(
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
input_responses: InputResponses | None = None,
request_state: str | None = None,
**kwargs: Any,
) -> CallToolResult | InputRequiredResult:
"""
@ -2493,6 +2513,8 @@ async def call_mcp_tool(
raw_headers=raw_headers,
client_ip=client_ip,
wire_compat=wire_compat,
input_responses=input_responses,
request_state=request_state,
**kwargs,
)
except Exception as e:
@ -2525,7 +2547,10 @@ async def mcp_get_prompt(
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> GetPromptResult:
input_responses: InputResponses | None = None,
request_state: str | None = None,
allow_input_required: bool = False,
) -> GetPromptResult | InputRequiredResult:
"""
Fetch a specific MCP prompt, handling both prefixed and unprefixed names.
"""
@ -2570,6 +2595,9 @@ async def mcp_get_prompt(
extra_headers=extra_headers,
raw_headers=raw_headers,
client_ip=client_ip,
input_responses=input_responses,
request_state=request_state,
allow_input_required=allow_input_required,
)
@ -2582,7 +2610,10 @@ async def mcp_read_resource(
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> ReadResourceResult:
input_responses: InputResponses | None = None,
request_state: str | None = None,
allow_input_required: bool = False,
) -> ReadResourceResult | InputRequiredResult:
"""Read resource contents from upstream MCP servers."""
allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
@ -2623,6 +2654,9 @@ async def mcp_read_resource(
extra_headers=extra_headers,
raw_headers=raw_headers,
client_ip=client_ip,
input_responses=input_responses,
request_state=request_state,
allow_input_required=allow_input_required,
)
@ -2671,6 +2705,8 @@ async def _handle_managed_mcp_tool(
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
input_responses: InputResponses | None = None,
request_state: str | None = None,
*,
catalog_auth_header: str | None,
) -> CallToolResult | InputRequiredResult:
@ -2696,6 +2732,8 @@ async def _handle_managed_mcp_tool(
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
wire_compat=wire_compat,
input_responses=input_responses,
request_state=request_state,
)
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
return call_tool_result
@ -2917,6 +2955,8 @@ async def _execute_mcp_server_tool_call(
client_ip=_client_ip,
host_progress_callback=host_progress_callback,
wire_compat=context.wire_compat,
input_responses=params.input_responses,
request_state=params.request_state,
**data, # for logging
)
except MCPMissingUserEnvVarsError as e:
@ -3008,7 +3048,7 @@ async def _execute_list_prompts(
async def _execute_get_prompt(
context: OperationContext, params: GetPromptRequestParams, host_progress_callback: ProgressCallback | None = None
) -> GetPromptResult:
) -> GetPromptResult | InputRequiredResult:
if context.mcp_proxy_mode:
_reject_mcp_proxy_operation()
(
@ -3032,6 +3072,9 @@ async def _execute_get_prompt(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
input_responses=params.input_responses,
request_state=params.request_state,
allow_input_required=context.wire_compat is WireCompat.MODERN,
)
@ -3093,7 +3136,7 @@ async def _execute_list_resource_templates(
async def _execute_read_resource(
context: OperationContext, params: ReadResourceRequestParams, host_progress_callback: ProgressCallback | None = None
) -> ReadResourceResult:
) -> ReadResourceResult | InputRequiredResult:
if context.mcp_proxy_mode:
_reject_mcp_proxy_operation()
(
@ -3115,6 +3158,9 @@ async def _execute_read_resource(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
input_responses=params.input_responses,
request_state=params.request_state,
allow_input_required=context.wire_compat is WireCompat.MODERN,
)
return read_resource_result
@ -3177,6 +3223,17 @@ GatewayResult: TypeAlias = (
)
def validate_continuation(operation: InteractionOperation, context: OperationContext) -> ContinuationState | None:
state: Final = open_continuation(operation, context, now=int(time.time()))
if state is not None:
target: Final = global_mcp_server_manager.get_mcp_server_by_id(state.target_id)
if target is None or target_digest(target) != state.target_digest:
raise MCPError(code=-32602, message="MCP continuation target changed; start a fresh request")
if set(operation.params.input_responses or {}) & set(state.gateway_responses or {}):
raise MCPError(code=-32602, message="Cannot replace gateway input responses")
return state
class GatewayOperations:
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
self._host_progress_callback = host_progress_callback
@ -3201,7 +3258,9 @@ class GatewayOperations:
async def execute(self, operation: ListPromptsRequest, context: OperationContext) -> ListPromptsResult: ...
@overload
async def execute(self, operation: GetPromptRequest, context: OperationContext) -> GetPromptResult: ...
async def execute(
self, operation: GetPromptRequest, context: OperationContext
) -> GetPromptResult | InputRequiredResult: ...
@overload
async def execute(self, operation: ListResourcesRequest, context: OperationContext) -> ListResourcesResult: ...
@ -3212,10 +3271,43 @@ class GatewayOperations:
) -> ListResourceTemplatesResult: ...
@overload
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
async def execute(
self, operation: ReadResourceRequest, context: OperationContext
) -> ReadResourceResult | InputRequiredResult: ...
@catalog_operation(lambda: global_mcp_server_manager)
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
if not isinstance(operation, (CallToolRequest, GetPromptRequest, ReadResourceRequest)):
return await self._execute(operation, context)
if context.wire_compat is not WireCompat.MODERN:
if operation.params.request_state is not None or operation.params.input_responses:
raise MCPError(code=-32602, message="Continuations require the modern MCP protocol")
return await self._execute(operation, context)
state: Final = validate_continuation(operation, context)
upstream: Final = operation.model_copy(
update={
"params": operation.params.model_copy(
update={
"request_state": state.upstream_state if state is not None else None,
"input_responses": {
**(state.gateway_responses or {}),
**(operation.params.input_responses or {}),
}
if state is not None
else None,
}
)
}
)
dispatch_context: Final = replace(context, mcp_servers=(state.target_id,)) if state is not None else context
result: Final = await self._execute(upstream, dispatch_context)
if isinstance(result, BoundInputRequiredResult):
return seal_continuation(result, operation, context, now=int(time.time()), previous=state)
if isinstance(result, InputRequiredResult):
raise MCPError(code=-32602, message="MCP continuation target is unavailable")
return result
async def _execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case DiscoverRequest():
listings: Final = (
@ -3282,6 +3374,8 @@ class GatewayOperations:
host_progress_callback=operation.host_progress_callback,
guardrail_context=operation.guardrail_context,
wire_compat=context.wire_compat,
input_responses=operation.input_responses,
request_state=operation.request_state,
**operation.logging_data,
)
case ListToolsRequest(params=params):

View file

@ -131,12 +131,7 @@ def reject_disallowed_mcp_origin(request: StarletteRequest) -> None:
def unsupported_protocol_version(scope: Scope) -> str | None:
"""Return the unsupported ``MCP-Protocol-Version`` header value, if any.
SDK 2's ``StreamableHTTPSessionManager`` routes any version outside
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
bypasses litellm's session/auth model, so the ASGI entry rejects it.
"""
"""Admit configured HTTP revisions while keeping SSE on the legacy protocol."""
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
@ -144,6 +139,8 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
)
for value in values:
if value == "2026-07-28" and scope.get("path", "").rstrip("/").endswith("/sse"):
return value
if value and value not in configured_versions():
return value
return None
@ -507,9 +504,11 @@ if MCP_AVAILABLE:
from mcp.server.lowlevel.server import NotificationOptions
from mcp.server.models import InitializationOptions
from mcp.shared.exceptions import MCPError
from mcp.shared.inbound import InboundLadderRejection, classify_inbound_request, find_duplicated_routing_header
from mcp.types import (
CallToolRequest,
GetPromptRequest,
JSONRPCRequest,
ListPromptsRequest,
ListResourcesRequest,
ListResourceTemplatesRequest,
@ -519,6 +518,7 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.interactions import InteractionOperation
from litellm.proxy._experimental.mcp_server.operations import (
_invalidate_byok_cred_cache,
_mcp_session_id_from_headers,
@ -926,7 +926,9 @@ if MCP_AVAILABLE:
verbose_logger.exception("Error in list_prompts endpoint: %s", exc)
return ListPromptsResult(prompts=[])
async def get_prompt(ctx: ServerRequestContext, params: GetPromptRequestParams) -> GetPromptResult:
async def get_prompt(
ctx: ServerRequestContext, params: GetPromptRequestParams
) -> GetPromptResult | InputRequiredResult:
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
async with _legacy_operation_context(ctx, trace=False) as context:
@ -964,7 +966,9 @@ if MCP_AVAILABLE:
verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc)
return ListResourceTemplatesResult(resource_templates=[])
async def read_resource(ctx: ServerRequestContext, params: ReadResourceRequestParams) -> ReadResourceResult:
async def read_resource(
ctx: ServerRequestContext, params: ReadResourceRequestParams
) -> ReadResourceResult | InputRequiredResult:
if _mcp_proxy_mode.get():
_reject_mcp_proxy_operation()
async with _legacy_operation_context(ctx, trace=False) as context:
@ -1389,6 +1393,8 @@ if MCP_AVAILABLE:
async def _read_request_body_for_routing(
receive: Receive,
*,
full_body: bool = False,
) -> tuple[list[Message], bytes]:
"""
Read just enough of the request body to decide whether this is a
@ -1401,7 +1407,8 @@ if MCP_AVAILABLE:
The remainder of an oversized body is streamed lazily through
``wrapped_receive`` in the caller — so an authenticated client cannot
force the proxy to buffer an arbitrarily large payload just to make a
routing decision.
routing decision. Modern interaction preflight requests the full body,
which the SDK's single-exchange transport also requires.
"""
consumed_messages: Final[list[Message]] = []
body_chunks: Final[list[bytes]] = []
@ -1409,6 +1416,12 @@ if MCP_AVAILABLE:
while True:
message = await receive()
if (
full_body
and peeked_bytes + len(message.get("body", b"") or b"")
> session_manager_stateless.max_request_body_size
):
raise HTTPException(status_code=413, detail="Request body too large")
consumed_messages.append(message)
if message.get("type") != "http.request":
@ -1422,7 +1435,7 @@ if MCP_AVAILABLE:
# handler via ``consumed_messages``, but ``body_chunks`` is
# purely for the JSON-RPC method check — there is no reason
# to copy a large body frame into a second buffer.
remaining = _MCP_ROUTING_PEEK_MAX_BYTES - peeked_bytes
remaining = len(body) if full_body else _MCP_ROUTING_PEEK_MAX_BYTES - peeked_bytes
if remaining > 0:
body_chunks.append(body[:remaining])
peeked_bytes += min(len(body), remaining)
@ -1430,7 +1443,7 @@ if MCP_AVAILABLE:
if not message.get("more_body", False):
break
if peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES:
if not full_body and peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES:
# Stop draining; downstream replay will pull remaining chunks
# directly from the original `receive` via wrapped_receive.
break
@ -1997,6 +2010,59 @@ if MCP_AVAILABLE:
detail="Forbidden",
)
_INTERACTION_REQUEST: Final[TypeAdapter[InteractionOperation]] = TypeAdapter(InteractionOperation)
@catalog_operation(lambda: operations.global_mcp_server_manager)
async def _preflight_modern_interaction(
scope: Scope, body: bytes, context: OperationContext
) -> JSONResponse | None:
try:
envelope: Final = JSONRPCRequest.model_validate_json(body)
operation: Final = _INTERACTION_REQUEST.validate_json(body)
except ValidationError:
return None
headers: Final = StarletteRequest(scope).headers
if find_duplicated_routing_header(headers.items()) is not None or isinstance(
classify_inbound_request(envelope.model_dump(by_alias=True), headers=dict(headers)), InboundLadderRejection
):
return None
try:
state: Final = operations.validate_continuation(operation, context)
except MCPError as error:
return JSONResponse(
status_code=400,
content={"jsonrpc": "2.0", "id": envelope.id, "error": error.error.model_dump(exclude_none=True)},
)
if state is None and context.mcp_servers is None:
return None
targets: Final = (
[state.target_id]
if state is not None
else list(context.mcp_servers)
if context.mcp_servers is not None
else None
)
allowed: Final = await operations._get_allowed_mcp_servers(
user_api_key_auth=context.user_api_key_auth, mcp_servers=targets, client_ip=context.client_ip
)
if state is not None and not any(target.server_id == state.target_id for target in allowed):
raise HTTPException(status_code=403, detail="MCP continuation target is no longer authorized")
authorized_names: Final = [target.alias or target.name for target in allowed]
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=list(context.mcp_servers) if context.mcp_servers is not None else authorized_names,
oauth2_headers=dict(context.oauth2_headers) if context.oauth2_headers is not None else None,
mcp_server_auth_headers={key: dict(value) for key, value in context.mcp_server_auth_headers.items()}
if context.mcp_server_auth_headers is not None
else None,
user_api_key_auth=context.user_api_key_auth,
client_ip=context.client_ip,
allowed_server_ids={target.server_id for target in allowed},
raw_headers=context.raw_headers,
)
await _check_passthrough_upstream_auth(scope, context.user_api_key_auth, authorized_names, context.client_ip)
return None
async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
"""Handle MCP requests through StreamableHTTP."""
try:
@ -2027,6 +2093,10 @@ if MCP_AVAILABLE:
) = await extract_mcp_auth_context(scope, path)
reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
request_headers: Final = StarletteRequest(scope).headers
defer_upstream_probes: Final = request_headers.get(
"mcp-protocol-version"
) == "2026-07-28" and request_headers.get("mcp-method") in {"tools/call", "prompts/get", "resources/read"}
# Extract client IP for MCP access control
_client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope))
@ -2055,21 +2125,22 @@ if MCP_AVAILABLE:
# from the fully-authorized server set: a passthrough server that
# the active toolset excludes should not trigger an OAuth flow
# for a server the caller will be 403'd on after authentication.
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
if not defer_upstream_probes:
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
# Pre-flight auth check for pass-through servers. Must run after
# toolset scoping so the probe list is derived from the fully-authorized
# server set, not the raw user-supplied names.
await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip)
# Pre-flight auth check for pass-through servers. Must run after
# toolset scoping so the probe list is derived from the fully-authorized
# server set, not the raw user-supplied names.
await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip)
# Inject masked debug headers when client sends x-litellm-mcp-debug: true
_debug_headers: Final = MCPDebug.maybe_build_debug_headers(
@ -2138,7 +2209,24 @@ if MCP_AVAILABLE:
body = b""
if scope.get("method") == "POST":
consumed_messages, body = await _read_request_body_for_routing(receive)
consumed_messages, body = await _read_request_body_for_routing(receive, full_body=defer_upstream_probes)
if defer_upstream_probes:
rejection: Final = await _preflight_modern_interaction(
scope,
body,
OperationContext(
_caller=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=tuple(mcp_servers) if mcp_servers is not None else None,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
),
)
if rejection is not None:
await rejection(scope, receive, send)
return
is_initialize = _is_initialize_request(body)
use_stateful: Final = bool(session_id or is_initialize)
@ -2279,12 +2367,16 @@ if MCP_AVAILABLE:
client_info=_extract_initialize_client_info(body),
)
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=is_initialize,
async with (
contextlib.nullcontext()
if defer_upstream_probes
else _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=is_initialize,
)
):
await target_manager.handle_request(scope, receive, local_send)
if use_stateful and session_id and scope.get("method") == "DELETE":

View file

@ -3,7 +3,7 @@ Per-feature OpenAPI snapshot for lazy-loaded routers.
The committed JSON is generated by `python -m litellm.proxy._lazy_openapi_snapshot`
and consumed at runtime so /openapi.json can show full route info for unloaded
features without importing them. check-ui-api-types.yml (mirrored locally by
features without importing them. test-linting.yml (mirrored locally by
`make check`) regenerates this file and fails when the committed copy differs,
then rebuilds schema.d.ts from app.openapi() with the snapshot injected. After
changing any lazily loaded route or this generator, rerun the module and commit

View file

@ -3206,7 +3206,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
mcp_advertised_versions: MCPAdvertisedVersions | None = Field(
None,
description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. "
"Modern protocol serving and Apps/Tasks remain disabled.",
"Modern protocol serving requires explicit opt-in. Apps/Tasks remain disabled.",
)
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
None,

View file

@ -18,6 +18,8 @@ class _CredentialInfo(TypedDict):
class _CredentialItem(TypedDict):
credential_name: ReadOnly[NotRequired[str]]
display_name: ReadOnly[NotRequired[str | None]]
source: ReadOnly[NotRequired[str]]
credential_info: ReadOnly[NotRequired[_CredentialInfo]]
@ -33,6 +35,10 @@ class _JsonBodyView(TypedDict):
body: ReadOnly[object]
def _print_json(data: object) -> None:
rich.print_json(data=data)
@click.group()
def credentials():
"""Manage credentials for the LiteLLM proxy server"""
@ -55,12 +61,14 @@ def list(ctx: click.Context, output_format: Literal["table", "json"]):
assert isinstance(response, dict)
if output_format == "json":
rich.print_json(data=response)
_print_json(response)
else: # table format
table: Final = Table(title="Credentials")
# Add columns
table.add_column("Credential Name", style="cyan")
table.add_column("Display Name")
table.add_column("Source")
table.add_column("Custom LLM Provider", style="green")
# Add rows
@ -69,6 +77,8 @@ def list(ctx: click.Context, output_format: Literal["table", "json"]):
info = cred.get("credential_info", {})
table.add_row(
str(cred.get("credential_name", "")),
cred.get("display_name") or "",
str(cred.get("source", "")),
str(info.get("custom_llm_provider", "")),
)
@ -89,8 +99,9 @@ def list(ctx: click.Context, output_format: Literal["table", "json"]):
help="JSON string containing credential values",
required=True,
)
@click.option("--display-name", type=str, default=None, help="Optional label shown in the UI")
@click.pass_context
def create(ctx: click.Context, credential_name: str, info: str, values: str):
def create(ctx: click.Context, credential_name: str, info: str, values: str, display_name: str | None) -> None:
"""Create a new credential"""
context: Final = cli_context_values(ctx)
client: Final = CredentialsManagementClient(context["base_url"], context["api_key"])
@ -101,13 +112,39 @@ def create(ctx: click.Context, credential_name: str, info: str, values: str):
raise click.BadParameter(f"Invalid JSON: {e}")
try:
response: Final = client.create(credential_name, credential_info["value"], credential_values["value"])
rich.print_json(data=response)
response: Final = client.create(
credential_name, credential_info["value"], credential_values["value"], display_name=display_name
)
_print_json(response)
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body: Final[_JsonBodyView] = {"body": e.response.json()}
rich.print_json(data=error_body["body"])
_print_json(error_body["body"])
except json.JSONDecodeError:
click.echo(e.response.text, err=True)
raise click.Abort()
@credentials.command()
@click.argument("credential_name")
@click.option("--display-name", type=str, default=None, help="New label shown in the UI")
@click.option("--clear-display-name", is_flag=True, help="Remove the label so the UI shows the credential name")
@click.pass_context
def update(ctx: click.Context, credential_name: str, display_name: str | None, clear_display_name: bool) -> None:
"""Change a credential's display name. The credential name itself cannot change"""
if (display_name is None) == (not clear_display_name):
raise click.UsageError("Pass exactly one of --display-name or --clear-display-name")
context: Final = cli_context_values(ctx)
client: Final = CredentialsManagementClient(context["base_url"], context["api_key"])
try:
response: Final = client.update_display_name(credential_name, display_name)
_print_json(response)
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body: Final[_JsonBodyView] = {"body": e.response.json()}
_print_json(error_body["body"])
except json.JSONDecodeError:
click.echo(e.response.text, err=True)
raise click.Abort()
@ -122,12 +159,12 @@ def delete(ctx: click.Context, credential_name: str):
client: Final = CredentialsManagementClient(context["base_url"], context["api_key"])
try:
response: Final = client.delete(credential_name)
rich.print_json(data=response)
_print_json(response)
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
try:
error_body: Final[_JsonBodyView] = {"body": e.response.json()}
rich.print_json(data=error_body["body"])
_print_json(error_body["body"])
except json.JSONDecodeError:
click.echo(e.response.text, err=True)
raise click.Abort()
@ -141,4 +178,4 @@ def get(ctx: click.Context, credential_name: str):
context: Final = cli_context_values(ctx)
client: Final = CredentialsManagementClient(context["base_url"], context["api_key"])
response: Final = client.get(credential_name)
rich.print_json(data=response)
_print_json(response)

View file

@ -1,5 +1,6 @@
from collections.abc import Mapping
from typing import Any, Final
from urllib.parse import quote
import requests
@ -73,6 +74,7 @@ class CredentialsManagementClient:
credential_info: Mapping[str, object],
credential_values: Mapping[str, object],
return_request: bool = False,
display_name: str | None = None,
) -> dict[str, Any] | requests.Request:
"""
Create a new credential.
@ -97,6 +99,7 @@ class CredentialsManagementClient:
"credential_name": credential_name,
"credential_info": credential_info,
"credential_values": credential_values,
**({} if display_name is None else {"display_name": display_name}),
}
request: Final = requests.Request("POST", url, headers=self._get_headers(), json=data)
@ -151,6 +154,30 @@ class CredentialsManagementClient:
raise UnauthorizedError(e)
raise
def update_display_name(
self,
credential_name: str,
display_name: str | None,
return_request: bool = False,
) -> Mapping[str, object] | requests.Request:
url: Final = f"{self._base_url}/credentials/{quote(credential_name, safe='')}"
data: Final[Mapping[str, object]] = {"display_name": display_name, "credential_info": {}}
request: Final = requests.Request("PATCH", url, headers=self._get_headers(), json=data)
if return_request:
return request
session: Final = requests.Session()
try:
response: Final = session.send(request.prepare(), timeout=self._timeout)
response.raise_for_status()
return response.json()
except requests.exceptions.HTTPError as e:
if e.response.status_code == 401:
raise UnauthorizedError(e)
raise
def get(
self,
credential_name: str,

View file

@ -22,7 +22,7 @@ from litellm.llms.anthropic.wif import (
UnbuildableIdentitySource,
anthropic_internal_issuer_jwks,
)
from litellm.models.credentials import UpdateCredentialItem
from litellm.models.credentials import CredentialView, UpdateCredentialItem
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
@ -32,7 +32,6 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.credential_hydration import (
hydrate_named_credential,
hydrate_named_credential_authoritative,
named_credential_wif_fields,
stored_credential_provider,
@ -46,6 +45,7 @@ from litellm.types.utils import CreateCredentialItem, CredentialItem
router: Final = APIRouter()
_CREDENTIAL_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
_DISPLAY_NAME_MAX_LENGTH: Final = 255
def _reject_non_admin_wif_fields(
@ -105,7 +105,7 @@ def _without_null_values(credential_values: Mapping[str, object]) -> dict[str, o
return {key: value for key, value in credential_values.items() if value is not None}
def _sync_in_memory_credential(credential: CredentialItem, credential_name: str, new_name: str) -> None:
def _sync_in_memory_credential(credential: CredentialItem, credential_name: str) -> None:
"""Mirror a DB credential update into the in-memory ``credential_list`` used by request-time
resolution; a no-op if the credential isn't loaded in memory (e.g. proxy restarted since boot).
"""
@ -127,13 +127,11 @@ def _sync_in_memory_credential(credential: CredentialItem, credential_name: str,
if credential.credential_info:
in_memory_info.update(credential.credential_info)
updated_in_memory: Final = CredentialItem(
credential_name=new_name,
credential_name=credential_name,
display_name=credential.display_name,
credential_values=in_memory_values,
credential_info=in_memory_info,
)
# Remove old entry if renamed, then use upsert_credentials to handle duplicates
if new_name != credential_name:
litellm.credential_list = [c for c in litellm.credential_list if c.credential_name != credential_name]
CredentialAccessor.upsert_credentials([updated_in_memory])
@ -149,11 +147,56 @@ class CredentialHelperUtils:
# is kept in memory and should remain unencrypted.
return CredentialItem(
credential_name=credential.credential_name,
display_name=credential.display_name,
credential_values=encrypted_credential_values,
credential_info=credential.credential_info or {},
)
def _normalized_display_name(display_name: str | None) -> str | None:
if display_name is None:
return None
trimmed: Final = display_name.strip()
if not trimmed:
raise ProxyException(
message="display_name cannot be blank. Send null to clear it or omit the field to leave it unchanged.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="display_name",
)
if len(trimmed) > _DISPLAY_NAME_MAX_LENGTH:
raise ProxyException(
message=f"display_name cannot be longer than {_DISPLAY_NAME_MAX_LENGTH} characters.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="display_name",
)
return trimmed
def _not_found_unless_config_defined(credential_name: str, not_found_detail: str) -> HTTPException | ProxyException:
in_memory: Final = CredentialAccessor.find_credential(credential_name)
if in_memory is None or in_memory.source != "config":
return HTTPException(status_code=404, detail=not_found_detail)
return ProxyException(
message=f"Credential '{credential_name}' is defined in config and cannot be edited from the API or UI.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_405_METHOD_NOT_ALLOWED,
param="credential_name",
headers={"Allow": "GET"},
)
def _credential_view(credential: CredentialItem, credential_values: Mapping[str, object]) -> CredentialView:
return CredentialView(
credential_name=credential.credential_name,
display_name=credential.display_name,
credential_values=credential_values,
credential_info=credential.credential_info,
source=credential.source,
)
def _credential_exists_detail(credential_name: str) -> str:
return (
f"Credential '{credential_name}' already exists. "
@ -222,6 +265,7 @@ async def create_credential(
)
processed_credential: Final = CredentialItem(
credential_name=credential.credential_name,
display_name=_normalized_display_name(credential.display_name),
credential_values=_without_null_values(_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values)),
credential_info=credential.credential_info,
)
@ -267,11 +311,7 @@ async def get_credentials(
"""
try:
masked_credentials: Final = [
{
"credential_name": credential.credential_name,
"credential_values": get_masked_values(credential.credential_values),
"credential_info": credential.credential_info,
}
_credential_view(credential, get_masked_values(credential.credential_values))
for credential in litellm.credential_list
]
return {"success": True, "credentials": masked_credentials}
@ -283,7 +323,7 @@ async def get_credentials(
"/credentials/by_name/{credential_name:path}",
dependencies=[Depends(user_api_key_auth)],
tags=["credential management"],
response_model=CredentialItem,
response_model=CredentialView,
)
async def get_credential_by_name(
request: Request,
@ -297,16 +337,10 @@ async def get_credential_by_name(
try:
for credential in litellm.credential_list:
if credential.credential_name == credential_name:
masked_credential = CredentialItem(
credential_name=credential.credential_name,
credential_values=get_masked_values(
credential.credential_values,
unmasked_length=4,
number_of_asterisks=4,
),
credential_info=credential.credential_info,
return _credential_view(
credential,
get_masked_values(credential.credential_values, unmasked_length=4, number_of_asterisks=4),
)
return masked_credential
raise HTTPException(
status_code=404,
detail="Credential not found. Got credential name: " + credential_name,
@ -444,9 +478,8 @@ async def delete_credential(
)
deleted: Final = await CredentialsRepository(prisma_client).delete_by_name(credential_name)
if deleted is None:
raise HTTPException(
status_code=404,
detail="Credential not found. Got credential name: " + credential_name,
raise _not_found_unless_config_defined(
credential_name, "Credential not found. Got credential name: " + credential_name
)
## DELETE FROM LITELLM ##
@ -466,6 +499,7 @@ def update_db_credential(
"""
merged_credential: Final = CredentialItem(
credential_name=db_credential.credential_name,
display_name=updated_patch.display_name,
credential_info=db_credential.credential_info,
credential_values=db_credential.credential_values,
)
@ -474,10 +508,6 @@ def update_db_credential(
updated_patch,
new_encryption_key,
)
# update model name
if encrypted_credential.credential_name:
merged_credential.credential_name = encrypted_credential.credential_name
# update litellm params
if encrypted_credential.credential_values:
# Encrypt any sensitive values
@ -515,6 +545,14 @@ async def update_credential(
from litellm.proxy.proxy_server import prisma_client
try:
if credential.credential_name and credential.credential_name != credential_name:
raise ProxyException(
message="credential_name is immutable. Set display_name to change how the credential is labeled.",
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="credential_name",
)
requested_display_name: Final = _normalized_display_name(credential.display_name)
_reject_overlapping_credential_values(credential)
incoming_values: Final = _CREDENTIAL_DICT_ADAPTER.validate_python(
_resolve_deployment_credentials(llm_router, credential.model_id)
@ -530,14 +568,13 @@ async def update_credential(
credentials_repository: Final = CredentialsRepository(prisma_client)
db_credential: Final = await credentials_repository.find_by_name(credential_name)
if db_credential is None:
raise HTTPException(status_code=404, detail="Credential not found in DB.")
raise _not_found_unless_config_defined(credential_name, "Credential not found in DB.")
_reject_non_admin_wif_fields(_stored_wif_fields(db_credential), user_api_key_dict)
if credential.credential_name != credential_name:
shadowed_credential: Final = await hydrate_named_credential(credential.credential_name, prisma_client)
if shadowed_credential is not None:
_reject_non_admin_wif_fields(_stored_wif_fields(shadowed_credential), user_api_key_dict)
patch: Final = CredentialItem(
credential_name=credential.credential_name,
credential_name=credential_name,
display_name=(
requested_display_name if "display_name" in credential.model_fields_set else db_credential.display_name
),
credential_info=_CREDENTIAL_DICT_ADAPTER.validate_python(credential.credential_info),
credential_values=incoming_values,
credential_values_to_delete=credential.credential_values_to_delete,
@ -550,12 +587,13 @@ async def update_credential(
credential_name,
data={
**credential_object_jsonified,
"display_name": merged_credential.display_name,
"updated_by": user_api_key_dict.user_id,
},
)
# Sync in-memory credential_list (skip if not in memory - e.g., proxy restarted)
_sync_in_memory_credential(patch, credential_name, merged_credential.credential_name)
_sync_in_memory_credential(patch, credential_name)
return {"success": True, "message": "Credential updated successfully"}
except Exception as e:

View file

@ -20,8 +20,8 @@ the deployment as a whole costs the primary one statement per interval.
"""
import json
from collections.abc import AsyncIterator, Iterable
from datetime import datetime, timezone
from collections.abc import AsyncIterator, Callable, Iterable
from datetime import datetime
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
@ -40,6 +40,7 @@ from litellm.types.proxy.gateway_requests import (
GatewayRequestKey,
GatewayRequestSnapshot,
)
from litellm.utils import get_utc_datetime
_GATEWAY_REQUEST_QUEUE_TARGET: Final = "gateway_request_queue"
@ -58,18 +59,15 @@ _BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...])
_NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({})
def _utc_date() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%d")
class GatewayRequestAccumulator:
"""Sink for the request-metrics middleware. ``record`` is sync and never awaits."""
def __init__(self) -> None:
def __init__(self, *, clock: Callable[[], datetime] = get_utc_datetime) -> None:
self._clock: Final = clock
self._counts: dict[GatewayRequestKey, GatewayRequestCounts] = {}
def record(self, *, category: BillableCategory, route: str, status_code: int) -> None:
key: Final = GatewayRequestKey(date=_utc_date(), category=category.value, route=route)
key: Final = GatewayRequestKey(date=self._clock().strftime("%Y-%m-%d"), category=category.value, route=route)
self._counts[key] = self._counts.get(key, _EMPTY).plus(succeeded=200 <= status_code < 300)
def drain(self) -> GatewayRequestSnapshot:

View file

@ -2078,6 +2078,7 @@ async def test_model_connection(
"responses",
"anthropic_messages",
"ocr",
"evaluation",
]
| None = fastapi.Body(
None,

View file

@ -1,5 +1,6 @@
import asyncio
import sys
from collections.abc import Callable
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
@ -36,6 +37,10 @@ else:
InternalUsageCache = Any
def _precise_minute(now: datetime) -> str:
return now.strftime("%Y-%m-%d-%H-%M")
def _response_total_tokens(response_obj: object) -> int:
if not isinstance(response_obj, (ModelResponse, EmbeddingResponse, TextCompletionResponse)):
return 0
@ -54,8 +59,9 @@ class CacheObject(TypedDict):
class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# Class variables or attributes
def __init__(self, internal_usage_cache: InternalUsageCache):
def __init__(self, internal_usage_cache: InternalUsageCache, *, clock: Callable[[], datetime] = datetime.now):
self.internal_usage_cache = internal_usage_cache
self._clock: Final = clock
def print_verbose(self, print_statement) -> None:
try:
@ -149,7 +155,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
def time_to_next_minute(self) -> float:
# Get the current time
now: Final = datetime.now()
now: Final = self._clock()
# Calculate the next minute
next_minute: Final = (now + timedelta(minutes=1)).replace(second=0, microsecond=0)
@ -306,10 +312,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
)
_model = data.get("model", None)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
current_hour: Final = datetime.now().strftime("%H")
current_minute: Final = datetime.now().strftime("%M")
precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
precise_minute: Final = _precise_minute(self._clock())
cache_objects: Final[CacheObject] = await self.get_all_cache_objects(
current_global_requests=(
@ -538,10 +541,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
litellm_parent_otel_span=litellm_parent_otel_span,
)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
current_hour: Final = datetime.now().strftime("%H")
current_minute: Final = datetime.now().strftime("%M")
precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
precise_minute: Final = _precise_minute(self._clock())
total_tokens: int = _response_total_tokens(response_obj)
@ -737,10 +737,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
litellm_parent_otel_span=litellm_parent_otel_span,
)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
current_hour: Final = datetime.now().strftime("%H")
current_minute: Final = datetime.now().strftime("%M")
precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
precise_minute: Final = _precise_minute(self._clock())
request_count_api_key: Final = f"{user_api_key}::{precise_minute}::request_count"
@ -813,10 +810,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
Retrieve the key's remaining rate limits.
"""
api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
current_date: Final = datetime.now().strftime("%Y-%m-%d")
current_hour: Final = datetime.now().strftime("%H")
current_minute: Final = datetime.now().strftime("%M")
precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
precise_minute: Final = _precise_minute(self._clock())
request_count_api_key: Final = f"{api_key}::{precise_minute}::request_count"
current: Final[CurrentItemRateLimit | None] = await self.internal_usage_cache.async_get_cache(
key=request_count_api_key,

View file

@ -6313,7 +6313,10 @@ class ProxyConfig:
credential_list_dict: Final = config.get("credential_list")
credential_list = []
if credential_list_dict:
credential_list = [CredentialItem(**cred) for cred in credential_list_dict]
credential_list = [
CredentialItem.model_validate({**cred, "display_name": None, "source": "config"})
for cred in credential_list_dict
]
return credential_list
def parse_search_tools(self, config: dict) -> list[SearchToolTypedDict] | None:

View file

@ -42,6 +42,7 @@ model LiteLLM_BudgetTable {
model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
display_name String?
credential_values Json
credential_info Json?
created_at DateTime @default(now()) @map("created_at")

View file

@ -36,6 +36,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attributio
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import PrismaTableRepository
from litellm.utils import get_utc_datetime
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -607,6 +608,8 @@ async def run_scheduled_ptu_rollup(
target_date: date | None = None,
alert: Callable[[str], Awaitable[None]] | None = None,
router: object | None = None,
*,
clock: Callable[[], datetime] = get_utc_datetime,
) -> RollupResult | None:
"""Run the daily rollup under a cross-pod lock so only one proxy reconciles a day.
@ -630,8 +633,12 @@ async def run_scheduled_ptu_rollup(
if not is_ptu_cost_attribution_enabled():
return None
today: Final = clock().date()
if pod_lock_manager is None or pod_lock_manager.redis_cache is None:
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
return await _run_and_alert(
prisma_client, target_date=target_date, today=today, alert=alert, may_prune=False, router=router
)
if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS):
if await _lock_is_held(pod_lock_manager):
@ -645,10 +652,14 @@ async def run_scheduled_ptu_rollup(
"PTU rollup: could not take the rollup lock and no other pod holds it, "
"running unguarded rather than skipping the day"
)
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
return await _run_and_alert(
prisma_client, target_date=target_date, today=today, alert=alert, may_prune=False, router=router
)
try:
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True, router=router)
return await _run_and_alert(
prisma_client, target_date=target_date, today=today, alert=alert, may_prune=True, router=router
)
finally:
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
@ -672,6 +683,7 @@ async def _run_and_alert(
prisma_client: "PrismaClient",
*,
target_date: date | None,
today: date,
alert: "Callable[[str], Awaitable[None]] | None",
may_prune: bool = True,
router: object | None = None,
@ -687,8 +699,9 @@ async def _run_and_alert(
explicit date means reconcile exactly that day, so it stays a single-day operation.
Its failure is contained: the day's own result is returned either way.
"""
rollup_date: Final = target_date or today - timedelta(days=1)
result: Final = await run_ptu_flat_cost_rollup(
prisma_client, target_date=target_date, may_prune=may_prune, router=router
prisma_client, target_date=rollup_date, may_prune=may_prune, router=router
)
if result.rows_failed:
await _deliver_alert(
@ -706,13 +719,14 @@ async def _run_and_alert(
"by the provider with nothing attributing it here. Extend the window, or retire the deployment.",
)
if target_date is None:
await _backfill_and_alert(prisma_client, alert=alert, router=router)
await _backfill_and_alert(prisma_client, today=today, alert=alert, router=router)
return result
async def _backfill_and_alert(
prisma_client: "PrismaClient",
*,
today: date,
alert: "Callable[[str], Awaitable[None]] | None",
router: object | None = None,
) -> None:
@ -722,7 +736,7 @@ async def _backfill_and_alert(
caller whatever the catch-up pass does.
"""
try:
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, router=router)
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, today=today, router=router)
except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup
verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc)
return

View file

@ -58,6 +58,7 @@ class CredentialsRepository:
return CredentialItem.model_validate(
{
"credential_name": data["credential_name"],
"display_name": data.get("display_name"),
"credential_values": data.get("credential_values") or {},
"credential_info": data.get("credential_info") or {},
}

View file

@ -10451,6 +10451,7 @@ class Router:
model_cost={model_id: model_info},
persist_across_reloads=False,
warning_display_name=model,
custom_llm_provider=custom_llm_provider,
)
## OLD MODEL REGISTRATION ## Kept to prevent breaking changes

View file

@ -4,7 +4,7 @@ import enum
import re
from collections.abc import Awaitable, Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias
from urllib.parse import urlsplit
import httpx
@ -71,7 +71,8 @@ def validate_mcp_protocol_transport(protocol_version: MCPUpstreamProtocol, trans
raise ValueError("Modern MCP requires HTTP or stdio transport")
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
MCPAdvertisedVersion: TypeAlias = MCPLegacyVersion | Literal["2026-07-28"]
MCPAdvertisedVersions: TypeAlias = Annotated[tuple[MCPAdvertisedVersion, ...], Field(min_length=1)]
MCPSpecVersionType = Literal[
MCPSpecVersion.nov_2024,
MCPSpecVersion.mar_2025,

View file

@ -3260,7 +3260,7 @@ def is_generalized_model_info(model_info: ModelInfo) -> bool:
return key not in litellm.model_cost and match_capability_generalizations(key) is not None
def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None:
def _get_builtin_model_info_for_registration(model: str, custom_llm_provider: str | None) -> ModelInfo | None:
"""Resolve ``model`` to its built-in cost-map entry for registration merging.
Returns ``None`` when the lookup raises or when it resolved via a
@ -3269,7 +3269,7 @@ def _get_builtin_model_info_for_registration(model: str) -> ModelInfo | None:
inheritance for prefix-mangled keys.
"""
try:
info: Final = get_model_info(model=model)
info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return None
return None if is_generalized_model_info(info) else info
@ -3344,6 +3344,7 @@ def register_model(
*,
persist_across_reloads: bool = True,
warning_display_name: str | None = None,
custom_llm_provider: str | None = None,
):
"""
Register new / Override existing models (and their pricing) to specific providers.
@ -3368,6 +3369,10 @@ def register_model(
``warning_display_name`` names the model in the missing-cache-pricing
warning instead of the registered key, for callers that register under an
opaque key (e.g. the router's hashed deployment ids).
``custom_llm_provider`` scopes the built-in cost-map match to entries for
that provider, so a deployment id that happens to equal another provider's
catalog key stays its own provider-less entry instead of merging into it.
"""
loaded_model_cost = {}
@ -3394,7 +3399,9 @@ def register_model(
existing_model = litellm.model_cost.get(key, {})
model_cost_key = key
else:
builtin_model_info = _get_builtin_model_info_for_registration(model=_key_str)
builtin_model_info = _get_builtin_model_info_for_registration(
model=_key_str, custom_llm_provider=custom_llm_provider
)
if builtin_model_info is not None:
existing_model = cast(dict, builtin_model_info)
model_cost_key = existing_model["key"]

View file

@ -2999,6 +2999,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -3295,6 +3296,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -23585,6 +23587,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -32372,6 +32375,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -37547,6 +37551,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -47801,6 +47806,7 @@
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"deprecation_date": "2027-04-08",
"litellm_provider": "bedrock_converse",
"supports_tool_search": true,
"max_input_tokens": 200000,
@ -80595,17 +80601,25 @@
"global.openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3e-05,
"cache_creation_input_token_cost_ultrafast": 1.5e-05,
"cache_read_input_token_cost": 1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-06,
"cache_read_input_token_cost_ultrafast": 6e-07,
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.4e-05,
"input_cost_per_token_ultrafast": 1.2e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9e-05,
"output_cost_per_token_ultrafast": 6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
@ -80633,17 +80647,25 @@
"bedrock_mantle/openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.75e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3.3e-05,
"cache_creation_input_token_cost_ultrafast": 1.65e-05,
"cache_read_input_token_cost": 1.1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.32e-06,
"cache_read_input_token_cost_ultrafast": 6.6e-07,
"input_cost_per_token": 2.2e-06,
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.64e-05,
"input_cost_per_token_ultrafast": 1.32e-05,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "responses",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9.9e-05,
"output_cost_per_token_ultrafast": 6.6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",
@ -80672,17 +80694,25 @@
"us.openai.gpt-6.1-sol": {
"cache_creation_input_token_cost": 2.75e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5.5e-06,
"cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3.3e-05,
"cache_creation_input_token_cost_ultrafast": 1.65e-05,
"cache_read_input_token_cost": 1.1e-07,
"cache_read_input_token_cost_above_272k_tokens": 2.2e-07,
"cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.32e-06,
"cache_read_input_token_cost_ultrafast": 6.6e-07,
"input_cost_per_token": 2.2e-06,
"input_cost_per_token_above_272k_tokens": 4.4e-06,
"input_cost_per_token_above_272k_tokens_ultrafast": 2.64e-05,
"input_cost_per_token_ultrafast": 1.32e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 1.1e-05,
"output_cost_per_token_above_272k_tokens": 1.65e-05,
"output_cost_per_token_above_272k_tokens_ultrafast": 9.9e-05,
"output_cost_per_token_ultrafast": 6.6e-05,
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html",
"supported_endpoints": [
"/v1/chat/completions",

View file

@ -42,6 +42,7 @@ model LiteLLM_BudgetTable {
model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
display_name String?
credential_values Json
credential_info Json?
created_at DateTime @default(now()) @map("created_at")

View file

@ -9,16 +9,16 @@
# - nothing staged -> scope is the working tree's diff against the merge base
# with origin's current default branch, untracked files included
# The per-area checks:
# - litellm/ Python -> `make lint` (test-linting.yml's lint job)
# - litellm/ Python -> `make lint` (test-linting.yml's python job)
# - tests/e2e and tests/e2e_harness Python
# -> `make lint-e2e-basedpyright` (test-linting.yml's e2e type-check step)
# + raw HTTP client ban (test-code-quality.yml's check_e2e_no_raw_requests)
# + raw HTTP client ban (test-linting.yml's check_e2e_no_raw_requests)
# - tests/ Python, ruff-tests.toml, scripts/check_test_quality.py,
# scripts/test_quality_gate.py
# -> ruff over ruff-tests.toml + `make lint-test-quality` (test-linting.yml's
# test-tree ruff and test-quality gate steps)
# - dashboard -> prettier + eslint + lint budgets (test-litellm-ui-build.yml's frontend-lint)
# - proxy/types -> regenerate the lazy OpenAPI snapshot and dashboard API types, fail on drift (check-ui-api-types.yml)
# - dashboard -> prettier + eslint + lint budgets (test-linting.yml's ui job)
# - proxy/types -> regenerate the lazy OpenAPI snapshot and dashboard API types, fail on drift (test-linting.yml's ui-api-types job)
#
# Each block is skipped when no matching files are in scope, so unrelated commits
# stay fast. This is intentionally not auto-installed as a git hook (see
@ -110,11 +110,11 @@ e2e_py_files=$(scope_match "$e2e_py_pattern")
test_tree_files=$(scope_match "$test_tree_pattern")
# ruff format (and CI's format step) skip enterprise; the rest of make lint covers it.
fmt_files=$(printf '%s\n' "$litellm_py_files" | grep -v '^litellm/enterprise/' | existing_files)
# check-ui-api-types.yml triggers on any file under litellm/proxy or litellm/types
# test-linting.yml triggers on any file under litellm/proxy or litellm/types
# (Prisma schema and configs included, not just Python) plus the generator and its
# lockfiles, so match that whole trigger set rather than a Python subset.
spec_files=$(scope_match "$spec_pattern")
# CI's frontend-lint runs prettier over a wider extension set than eslint; keep that
# CI's ui job runs prettier over a wider extension set than eslint; keep that
# split so this flags exactly what the job would.
ui_prettier_changed=$(scope_match "$ui_prettier_pattern")
ui_eslint_changed=$(scope_match "$ui_eslint_pattern")
@ -176,7 +176,7 @@ EOF
if [ ${#eslint_rel[@]} -gt 0 ]; then
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_rel[@]}" || rc=1
fi
# Whole-folder lint budgets, exactly as the frontend-lint job runs them: the
# Whole-folder lint budgets, exactly as the ui job runs them: the
# counts are not diff-scoped, so a local pass here means the budget step will
# pass in CI too.
report=$(mktemp)
@ -259,7 +259,7 @@ genapi_checks() {
local status=0
echo "check: checking the lazy OpenAPI snapshot and dashboard API types are in sync (npm run gen:api)"
# gen-api-types.mjs imports litellm.proxy.proxy_server, which needs the proxy deps
# and an up-to-date Prisma client; check-ui-api-types.yml installs those and runs
# and an up-to-date Prisma client; test-linting.yml installs those and runs
# prisma generate before gen:api, so mirror that here or a stale client can mask
# drift that CI will still flag.
if [ ! -d ui/litellm-dashboard/node_modules ]; then

View file

@ -12,7 +12,7 @@ because CI cannot enforce them on itself.
flagged: ``-`` appears in hyphenated input names like ``inputs.timeout-minutes``
and ``/`` inside ref strings, so neither can be told apart from arithmetic by
inspection alone.
2. Callers of the reusable unit-test workflow keep the job timeout at or above
2. Unit matrix rows keep the job timeout at or above
the test budget plus the setup ceilings plus the runner overhead below.
Otherwise the job deadline preempts pytest inside its own advertised budget,
which is the failure the split timeouts exist to prevent, and it shows up as
@ -24,7 +24,6 @@ because CI cannot enforce them on itself.
import re
import sys
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
@ -33,8 +32,7 @@ from pydantic import BaseModel, Field, ValidationError
REPO_ROOT: Final = Path(__file__).resolve().parent.parent.parent
WORKFLOWS_DIR: Final = REPO_ROOT / ".github" / "workflows"
BASE_WORKFLOW: Final = "./.github/workflows/_test-unit-base.yml"
BASE_WORKFLOW_PATH: Final = WORKFLOWS_DIR / "_test-unit-base.yml"
UNIT_WORKFLOW_PATH: Final = WORKFLOWS_DIR / "test-unit.yml"
# Runner time the job clock charges but no step owns: job init, the gaps between
# steps, and post-job cleanup. Without it a job capped at exactly test + setup
@ -44,24 +42,26 @@ JOB_OVERHEAD_MINUTES: Final = 5
EXPRESSION: Final = re.compile(r"\$\{\{(?P<body>.*?)\}\}", re.DOTALL)
QUOTED: Final = re.compile(r"'[^']*'")
ARITHMETIC: Final = re.compile(r"[+*]")
MATRIX_REF: Final = re.compile(r"^\$\{\{\s*matrix\.(?P<key>[\w-]+)\s*\}\}$")
JOB_DEADLINE: Final = "${{ matrix.job-timeout-minutes }}"
TEST_DEADLINE: Final = "${{ matrix.timeout-minutes }}"
class WorkflowStartupError(Exception):
pass
class ReusableCall(BaseModel):
uses: str | None = None
with_: Mapping[str, object] = Field(default_factory=dict, alias="with")
class WorkflowJob(BaseModel):
strategy: Mapping[str, object] = Field(default_factory=dict)
steps: tuple[Mapping[str, object], ...] = ()
timeout_minutes: object = Field(default=None, alias="timeout-minutes")
model_config = {"populate_by_name": True}
model_config = {"populate_by_name": True, "frozen": True}
class WorkflowFile(BaseModel):
jobs: Mapping[str, ReusableCall] = Field(default_factory=dict)
jobs: Mapping[str, WorkflowJob] = Field(default_factory=dict)
model_config = {"frozen": True}
def parse_workflow(text: str) -> WorkflowFile | str:
@ -79,10 +79,10 @@ def arithmetic_expressions(text: str) -> Iterator[str]:
yield body.strip()
def setup_ceiling_minutes(base_text: str) -> int:
"""Sum the per-step timeouts on everything the base workflow runs before pytest."""
base: Final = yaml.safe_load(base_text)
steps: Final = base["jobs"]["run"]["steps"]
def setup_ceiling_minutes(unit_text: str) -> int:
"""Sum the per-step timeouts on everything the unit job runs before pytest."""
workflow: Final = yaml.safe_load(unit_text)
steps: Final = workflow["jobs"]["unit"]["steps"]
return sum(
s["timeout-minutes"]
for s in steps
@ -90,100 +90,38 @@ def setup_ceiling_minutes(base_text: str) -> int:
)
def base_default(base_text: str, name: str) -> int:
base: Final = yaml.safe_load(base_text)
return base[True]["workflow_call"]["inputs"][name]["default"]
@dataclass(frozen=True, slots=True)
class Column:
"""A budget the caller reads from one column of its own matrix."""
name: str
def budget_source(job: ReusableCall, key: str, fallback: int) -> int | Column | str:
"""A caller passes a literal, or `${{ matrix.x }}` naming a column of its matrix.
Anything else comes back as the reason it could not be read, since a budget
nothing can resolve has to be reported rather than passed over.
"""
value: Final = job.with_.get(key)
if value is None:
return fallback
if isinstance(value, int):
return value
matrix_ref: Final = MATRIX_REF.match(str(value))
if not matrix_ref:
return f"passes `{key}: {value}`, which is neither a number nor a `matrix` reference."
return Column(matrix_ref.group("key"))
def matrix_rows(job: ReusableCall) -> Sequence[Mapping[str, object]]:
matrix: Final = job.strategy.get("matrix", {})
entries: Final = matrix.get("include", ()) if isinstance(matrix, dict) else ()
return tuple(e for e in entries if isinstance(e, dict))
def budget_pairs(job: ReusableCall, test_source: int | Column, job_source: int | Column) -> Iterator[tuple[int, int]]:
"""Pair each shard's test budget with the job budget of that same shard.
Matrix-sourced budgets resolve per `include` row, so two matrix columns are
read off the same row rather than cross-producted across rows.
"""
if isinstance(test_source, int) and isinstance(job_source, int):
yield test_source, job_source
def timeout_contract_errors(rel: Path, workflow: WorkflowFile, ceiling: int) -> Iterator[str]:
job: Final = workflow.jobs.get("unit")
if job is None:
return
for row in matrix_rows(job):
test_budget = row.get(test_source.name) if isinstance(test_source, Column) else test_source
job_budget = row.get(job_source.name) if isinstance(job_source, Column) else job_source
if isinstance(test_budget, int) and isinstance(job_budget, int):
yield test_budget, job_budget
def unresolved_message(where: str, job: ReusableCall, sources: Sequence[int | Column]) -> str:
"""Why no shard yielded a pair of budgets to compare.
Naming only the columns that resolve nowhere keeps the message honest: a
column every row supplies is not what left the pair unchecked.
"""
rows: Final = matrix_rows(job)
missing: Final = tuple(
f"`matrix.{s.name}`"
for s in sources
if isinstance(s, Column) and not any(isinstance(row.get(s.name), int) for row in rows)
)
if missing:
return (
f"{where} reads a budget from {', '.join(missing)}, which no `include` row supplies "
"as a number, so the pair would go unchecked."
)
return (
f"{where} reads both budgets from its matrix, but no single `include` row supplies both "
"as numbers, so the pair would go unchecked."
)
def job_errors(rel: Path, job_name: str, job: ReusableCall, ceiling: int, base_text: str) -> Iterator[str]:
where: Final = f"{rel}: job `{job_name}`"
test_source: Final = budget_source(job, "timeout-minutes", base_default(base_text, "timeout-minutes"))
job_source: Final = budget_source(job, "job-timeout-minutes", base_default(base_text, "job-timeout-minutes"))
sources: Final = (test_source, job_source)
unreadable: Final = tuple(f"{where} {reason}" for reason in sources if isinstance(reason, str))
if unreadable:
yield from unreadable
if job.timeout_minutes != JOB_DEADLINE:
yield f"{rel}: job `unit` sets timeout-minutes to `{job.timeout_minutes}`, not `{JOB_DEADLINE}`"
test_deadlines: Final = tuple(s.get("timeout-minutes") for s in job.steps if s.get("name") == "Run tests")
if test_deadlines != (TEST_DEADLINE,):
yield f"{rel}: job `unit` step `Run tests` must set timeout-minutes to `{TEST_DEADLINE}`, found {test_deadlines}"
matrix_value: Final = job.strategy.get("matrix", {})
if not isinstance(matrix_value, Mapping):
yield f"{rel}: job `unit` has no readable matrix"
return
pairs: Final = tuple(budget_pairs(job, test_source, job_source))
if not pairs:
yield unresolved_message(where, job, sources)
entries_value: Final = matrix_value.get("include", ())
if not isinstance(entries_value, Sequence) or isinstance(entries_value, str) or not entries_value:
yield f"{rel}: job `unit` has no readable matrix include rows"
return
for test_budget, job_budget in pairs:
required = test_budget + ceiling + JOB_OVERHEAD_MINUTES
for entry in entries_value:
if not isinstance(entry, Mapping):
yield f"{rel}: job `unit` has an unreadable matrix row"
continue
shard: Final = entry.get("shard", "<unnamed>")
test_budget: Final = entry.get("timeout-minutes")
job_budget: Final = entry.get("job-timeout-minutes")
where: Final = f"{rel}: job `unit`, shard `{shard}`"
if not isinstance(test_budget, int) or isinstance(test_budget, bool):
yield f"{where} has no integer timeout-minutes"
continue
if not isinstance(job_budget, int) or isinstance(job_budget, bool):
yield f"{where} has no integer job-timeout-minutes"
continue
required: Final = test_budget + ceiling + JOB_OVERHEAD_MINUTES
if job_budget < required:
yield (
f"{where} gives pytest {test_budget}m but caps the job at "
@ -193,13 +131,7 @@ def job_errors(rel: Path, job_name: str, job: ReusableCall, ceiling: int, base_t
)
def timeout_contract_errors(rel: Path, workflow: WorkflowFile, ceiling: int, base_text: str) -> Iterator[str]:
for job_name, job in workflow.jobs.items():
if job.uses == BASE_WORKFLOW:
yield from job_errors(rel, job_name, job, ceiling, base_text)
def workflow_errors(rel: Path, text: str, ceiling: int, base_text: str) -> Iterator[str]:
def workflow_errors(rel: Path, text: str, ceiling: int) -> Iterator[str]:
for expression in arithmetic_expressions(text):
yield (
f"{rel}: `${{{{ {expression} }}}}` uses arithmetic, which GitHub expressions do not "
@ -211,22 +143,20 @@ def workflow_errors(rel: Path, text: str, ceiling: int, base_text: str) -> Itera
yield f"{rel}: {workflow}"
return
yield from timeout_contract_errors(rel, workflow, ceiling, base_text)
yield from timeout_contract_errors(rel, workflow, ceiling)
def main() -> None:
base_text: Final = BASE_WORKFLOW_PATH.read_text()
ceiling: Final = setup_ceiling_minutes(base_text)
unit_text: Final = UNIT_WORKFLOW_PATH.read_text()
ceiling: Final = setup_ceiling_minutes(unit_text)
errors: Final = tuple(
error
for path in sorted(WORKFLOWS_DIR.glob("*.y*ml"))
for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text(), ceiling, base_text)
for error in workflow_errors(path.relative_to(REPO_ROOT), path.read_text(), ceiling)
)
if errors:
raise WorkflowStartupError(
"Workflow startup invariants violated:\n - " + "\n - ".join(errors)
)
raise WorkflowStartupError("Workflow startup invariants violated:\n - " + "\n - ".join(errors))
print(f"Workflow startup invariants hold (setup ceiling {ceiling}m)")

View file

@ -2,7 +2,7 @@
Harness logic, so it lives here rather than under tests/e2e, which holds only
tests that drive a live proxy. The harness modules are imported off
``PYTHONPATH=tests/e2e``, the way the Code Quality workflow's
``PYTHONPATH=tests/e2e``, the way the lint workflow's code-quality job's
test_e2e_metadata step runs this file. Call order, the failing test's last step,
the per-test reset and the JUnit attach are pinned end to end in
test_e2e_junit_report.py.

View file

@ -11,37 +11,6 @@ import pytest
HARNESS: Final = Path(__file__).parents[2] / ".github" / "scripts" / "run_merge_smoke.py"
CASE_IDS: Final = (
"CHAT-JSON",
"CHAT-TEXT-STREAM",
"CHAT-TOOL-STREAM",
"MODEL-ALLOW",
"MODEL-DENY",
"COST-EXPLICIT",
"COST-ZERO",
"LOG-CONTENT-ON",
"LOG-CONTENT-OFF",
"CALLBACK-SUCCESS",
"CALLBACK-FAILURE",
)
def _write_fake_tests(root: Path, body: str) -> Path:
package: Final = root / "fake_tests"
package.mkdir()
(package / "test_cases.py").write_text(body)
return package
def _manifest(root: Path, **overrides: str) -> Path:
cases: Final[dict[str, str]] = {
case_id: f"fake_tests/test_cases.py::test_{case_id.lower().replace('-', '_')}" for case_id in CASE_IDS
}
cases.update(overrides)
path: Final = root / "manifest.json"
path.write_text(json.dumps({"cases": cases}))
return path
def _run(root: Path, *argv: str) -> subprocess.CompletedProcess[str]:
return subprocess.run(
@ -53,132 +22,6 @@ def _run(root: Path, *argv: str) -> subprocess.CompletedProcess[str]:
)
def _passing_tests() -> str:
return "\n".join(f"def test_{case_id.lower().replace('-', '_')}():\n assert True" for case_id in CASE_IDS)
def test_all_eleven_cases_pass(tmp_path: Path) -> None:
_write_fake_tests(tmp_path, _passing_tests())
manifest: Final = _manifest(tmp_path)
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode == 0, proc.stderr
assert proc.stdout.count("PASS") >= 11
for case_id in CASE_IDS:
assert f"{case_id} PASS" in proc.stdout
def test_missing_test_node_id_fails(tmp_path: Path) -> None:
_write_fake_tests(tmp_path, _passing_tests())
manifest: Final = _manifest(tmp_path, **{"COST-ZERO": "fake_tests/test_cases.py::test_does_not_exist"})
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode != 0
assert "COST-ZERO" in proc.stderr or "test_does_not_exist" in proc.stderr
def test_skipped_case_fails(tmp_path: Path) -> None:
_write_fake_tests(
tmp_path,
_passing_tests().replace(
"def test_cost_zero():\n assert True",
"def test_cost_zero():\n import pytest\n pytest.skip('nope')",
),
)
manifest: Final = _manifest(tmp_path)
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode != 0
assert "COST-ZERO" in proc.stderr
def test_xfail_case_fails(tmp_path: Path) -> None:
_write_fake_tests(
tmp_path,
"import pytest\n"
+ _passing_tests().replace(
"def test_cost_zero():\n assert True",
"@pytest.mark.xfail\ndef test_cost_zero():\n assert False",
),
)
manifest: Final = _manifest(tmp_path)
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode != 0
assert "COST-ZERO" in proc.stderr
def test_xpass_case_fails(tmp_path: Path) -> None:
_write_fake_tests(
tmp_path,
"import pytest\n"
+ _passing_tests().replace(
"def test_cost_zero():\n assert True",
"@pytest.mark.xfail\ndef test_cost_zero():\n assert True",
),
)
manifest: Final = _manifest(tmp_path)
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode != 0
assert "COST-ZERO" in proc.stderr
def test_duplicate_manifest_key_fails(tmp_path: Path) -> None:
manifest: Final = tmp_path / "manifest.json"
manifest.write_text('{"cases": {"CHAT-JSON": "a::b", "CHAT-JSON": "a::c"}}')
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest))
assert proc.returncode != 0
assert "CHAT-JSON" in proc.stderr
def test_missing_case_id_fails(tmp_path: Path) -> None:
manifest: Final = tmp_path / "manifest.json"
cases: Final = {c: f"t::{c}" for c in CASE_IDS[:-1]}
manifest.write_text(json.dumps({"cases": cases}))
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest))
assert proc.returncode != 0
assert "case ids" in proc.stderr
def test_extra_case_id_fails(tmp_path: Path) -> None:
manifest: Final = tmp_path / "manifest.json"
cases: Final = {c: f"t::{c}" for c in CASE_IDS}
cases["EXTRA"] = "t::x"
manifest.write_text(json.dumps({"cases": cases}))
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest))
assert proc.returncode != 0
assert "case ids" in proc.stderr
def test_teardown_error_fails(tmp_path: Path) -> None:
body: Final = (
"import pytest\n\n@pytest.fixture\ndef boom():\n yield\n raise RuntimeError('teardown-boom')\n\n"
+ _passing_tests().replace(
"def test_cost_zero():\n assert True",
"def test_cost_zero(boom):\n assert True",
)
)
_write_fake_tests(tmp_path, body)
manifest: Final = _manifest(tmp_path)
proc: Final = _run(tmp_path, "pytest", "--manifest", str(manifest), "--rootdir", str(tmp_path))
assert proc.returncode != 0
assert "COST-ZERO" in proc.stderr
def _fake_litellm(tmp_path: Path, script: str) -> Path:
path: Final = tmp_path / "fake-litellm"
path.write_text(f"#!{sys.executable}\n" + textwrap.dedent(script))

View file

@ -7,24 +7,30 @@ from typing import Final
import pytest
import yaml
_UNIT_WORKFLOW: Final = Path(__file__).resolve().parents[2] / ".github" / "workflows" / "test-unit.yml"
_GATE_JOB: Final = "unit-passed"
_WORKFLOWS_DIR: Final = Path(__file__).resolve().parents[2] / ".github" / "workflows"
_TIERS: Final = (
("test-linting.yml", "lint-passed", frozenset()),
("test-unit.yml", "unit-passed", frozenset({"coverage"})),
("test-merge-smoke.yml", "smoke-passed", frozenset()),
)
_Tier = tuple[str, str, frozenset[str]]
def _jobs() -> dict[str, dict[str, object]]:
return yaml.safe_load(_UNIT_WORKFLOW.read_text())["jobs"]
def _jobs(tier: _Tier) -> dict[str, dict[str, object]]:
workflow: Final = yaml.safe_load((_WORKFLOWS_DIR / tier[0]).read_text())
return workflow["jobs"]
def _gate_script() -> str:
steps: Final = _jobs()[_GATE_JOB]["steps"]
def _gate_script(tier: _Tier) -> str:
steps: Final = _jobs(tier)[tier[1]]["steps"]
assert isinstance(steps, list) and len(steps) == 1
return steps[0]["run"]
def _run_gate(results: dict[str, str]) -> subprocess.CompletedProcess[str]:
def _run_gate(tier: _Tier, results: dict[str, str]) -> subprocess.CompletedProcess[str]:
needs: Final = {job: {"result": result, "outputs": {}} for job, result in results.items()}
return subprocess.run(
("bash", "--noprofile", "--norc", "-eo", "pipefail", "-c", _gate_script()),
("bash", "--noprofile", "--norc", "-eo", "pipefail", "-c", _gate_script(tier)),
env={**os.environ, "NEEDS": json.dumps(needs)},
capture_output=True,
text=True,
@ -33,36 +39,40 @@ def _run_gate(results: dict[str, str]) -> subprocess.CompletedProcess[str]:
)
def _needed_jobs() -> tuple[str, ...]:
needs: Final = _jobs()[_GATE_JOB]["needs"]
def _needed_jobs(tier: _Tier) -> tuple[str, ...]:
needs: Final = _jobs(tier)[tier[1]]["needs"]
assert isinstance(needs, list)
return tuple(needs)
def test_the_gate_passes_when_every_needed_job_succeeded() -> None:
jobs: Final = _needed_jobs()
@pytest.mark.parametrize("tier", _TIERS, ids=("lint", "unit", "smoke"))
def test_the_gate_passes_when_every_needed_job_succeeded(tier: _Tier) -> None:
jobs: Final = _needed_jobs(tier)
assert jobs
result: Final = _run_gate(dict.fromkeys(jobs, "success"))
result: Final = _run_gate(tier, {job: "success" for job in jobs})
assert result.returncode == 0, result.stdout + result.stderr
assert all(f"{job}: success" in result.stdout for job in jobs), result.stdout
@pytest.mark.parametrize("outcome", ("failure", "cancelled", "skipped"))
def test_the_gate_fails_when_any_needed_job_did_not_succeed(outcome: str) -> None:
jobs: Final = _needed_jobs()
@pytest.mark.parametrize("tier", _TIERS, ids=("lint", "unit", "smoke"))
def test_the_gate_fails_when_any_needed_job_did_not_succeed(tier: _Tier, outcome: str) -> None:
jobs: Final = _needed_jobs(tier)
assert len(jobs) > 1
result: Final = _run_gate({**dict.fromkeys(jobs, "success"), jobs[-1]: outcome})
result: Final = _run_gate(tier, {**{job: "success" for job in jobs}, jobs[-1]: outcome})
assert result.returncode != 0
assert f"{jobs[-1]}: {outcome}" in result.stdout, result.stdout
def test_the_gate_waits_for_every_other_job_and_reports_even_when_they_fail() -> None:
jobs: Final = _jobs()
gate: Final = jobs[_GATE_JOB]
@pytest.mark.parametrize("tier", _TIERS, ids=("lint", "unit", "smoke"))
def test_the_gate_waits_for_every_other_job_and_reports_even_when_they_fail(tier: _Tier) -> None:
jobs: Final = _jobs(tier)
gate: Final = jobs[tier[1]]
assert set(_needed_jobs()) == set(jobs) - {_GATE_JOB}
assert set(_needed_jobs(tier)) == set(jobs) - {tier[1]} - set(tier[2])
assert gate["if"] == "always()"
assert gate["permissions"] == {}

View file

@ -20,9 +20,13 @@ import {
* traffic, so each assertion is scoped to a key this test minted and to the requests it sent.
*/
const escapeRegExp = (value: string): string => value.replace(/[.*+?^${}()|[\]\\]/g, "\\$&");
/** Each breakdown renders one expandable card per entity, named "<entity> $x.xx N requests". */
const entityCard = (page: PlaywrightPage, tab: string, name: string): Locator =>
page.getByRole("tabpanel", { name: tab }).getByRole("button", { name: new RegExp(`^${name}\\s`) });
page
.getByRole("tabpanel", { name: tab })
.getByRole("button", { name: new RegExp(`(?:^|\\s)${escapeRegExp(name)}\\s`) });
async function openUsageTab(page: PlaywrightPage, tab: string): Promise<Locator> {
await navigateToPage(page, Page.NewUsage);

View file

@ -14,4 +14,4 @@ LITELLM_MASTER_KEY=sk-harness uv run pytest tests/e2e_harness
Rules: no `e2e` marker and no `@meta`, since nothing here drives the proxy; `@pytest.mark.covers` only where the test proves the collector or the JUnit properties read it; inputs via arguments or env vars (setting an env var through pytest's `monkeypatch` fixture is fine, patching a function, class or module is not); and the same typing bar as the suite, `make lint-e2e-basedpyright` covers this folder and allows zero errors. The raw HTTP client ban (`tests/code_coverage_tests/check_e2e_no_raw_requests.py`) applies here too
CI: the `lint` job in `.github/workflows/test-linting.yml` runs this folder whenever anything under `tests/e2e/` (except `ui/`) or `tests/e2e_harness/` changes, with the `claude` CLI installed. The CircleCI `provider_replay_harness` job also runs the provider-edge and fixture tests at the root of this folder next to `tests/code_coverage_tests/test_provider_replay_harness.py`, which imports helpers from `test_provider_edge.py`
CI: the `python` job in `.github/workflows/test-linting.yml` runs this folder whenever anything under `tests/e2e/` (except `ui/`) or `tests/e2e_harness/` changes, with the `claude` CLI installed. The CircleCI `provider_replay_harness` job also runs the provider-edge and fixture tests at the root of this folder next to `tests/code_coverage_tests/test_provider_replay_harness.py`, which imports helpers from `test_provider_edge.py`

View file

@ -0,0 +1,149 @@
import base64
import json
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from itertools import accumulate
from types import MappingProxyType
from typing import Final
from integration._support.wire import Reply
from pydantic import JsonValue
PDF_MEDIA_TYPE: Final = "application/pdf"
COUNT_TOKENS_TARGET: Final = "/v1/messages/count_tokens"
COUNT_REFUSED: Final = Reply(
status=400,
body=json.dumps(
{"type": "error", "error": {"type": "invalid_request_error", "message": "count_tokens is not supported"}}
).encode(),
)
@dataclass(frozen=True, slots=True)
class Page:
width: int = 612
height: int = 792
text: str | None = None
LETTER: Final = Page()
NARROW: Final = Page(width=100, height=1000)
_RENDERED_TOKENS: Final = MappingProxyType({(612, 792): 1534, (100, 1000): 328})
def rendered_tokens(pages: Sequence[Page]) -> int:
return sum(_RENDERED_TOKENS[(page.width, page.height)] for page in pages)
def _escaped(text: str) -> str:
return text.replace("\\", "\\\\").replace("(", "\\(").replace(")", "\\)")
def _content_stream(page: Page) -> bytes:
stream: Final = f"BT /F1 12 Tf 72 {page.height - 72} Td ({_escaped(page.text or '')}) Tj ET".encode()
return b"<< /Length " + str(len(stream)).encode() + b" >>\nstream\n" + stream + b"\nendstream"
def _page_object(page: Page, contents_id: int | None) -> bytes:
contents: Final = f" /Contents {contents_id} 0 R" if contents_id is not None else ""
return (
f"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 {page.width} {page.height}]"
f" /Resources << /Font << /F1 1 0 R >> >>{contents} >>"
).encode()
def _page_bodies(pages: Sequence[Page], starts: Sequence[int]) -> Iterator[bytes]:
for page, start in zip(pages, starts):
if page.text is None:
yield _page_object(page, None)
else:
yield _content_stream(page)
yield _page_object(page, start)
def pdf_bytes(pages: Sequence[Page]) -> bytes:
starts: Final = tuple(accumulate((1 if page.text is None else 2 for page in pages), initial=3))
page_ids: Final = tuple(start + (0 if page.text is None else 1) for start, page in zip(starts, pages))
kids: Final = " ".join(f"{identity} 0 R" for identity in page_ids)
bodies: Final = (
b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>",
f"<< /Type /Pages /Kids [{kids}] /Count {len(pages)} >>".encode(),
*_page_bodies(pages, starts),
b"<< /Type /Catalog /Pages 2 0 R >>",
)
header: Final = b"%PDF-1.4\n"
objects: Final = tuple(
f"{number} 0 obj\n".encode() + body + b"\nendobj\n" for number, body in enumerate(bodies, start=1)
)
offsets: Final = tuple(accumulate((len(chunk) for chunk in objects), initial=len(header)))
entries: Final = b"".join(f"{offset:010d} 00000 n \n".encode() for offset in offsets[:-1])
trailer: Final = (
f"xref\n0 {len(bodies) + 1}\n".encode()
+ b"0000000000 65535 f \n"
+ entries
+ f"trailer\n<< /Size {len(bodies) + 1} /Root {starts[-1]} 0 R >>\nstartxref\n{offsets[-1]}\n%%EOF\n".encode()
)
return header + b"".join(objects) + trailer
def encoded(raw: bytes) -> str:
return base64.b64encode(raw).decode()
def data_url(media_type: str, raw: bytes) -> str:
return f"data:{media_type};base64,{encoded(raw)}"
def pdf_data_url(pages: Sequence[Page]) -> str:
return data_url(PDF_MEDIA_TYPE, pdf_bytes(pages))
def base64_source(data: JsonValue, media_type: JsonValue = PDF_MEDIA_TYPE) -> dict[str, JsonValue]:
return {"type": "base64", "media_type": media_type, "data": data}
def document(source: Mapping[str, JsonValue], **fields: JsonValue) -> dict[str, JsonValue]:
return {"type": "document", "source": dict(source), **fields}
def pdf_document(pages: Sequence[Page], **fields: JsonValue) -> dict[str, JsonValue]:
return document(base64_source(encoded(pdf_bytes(pages))), **fields)
def text_document(text: str) -> dict[str, JsonValue]:
return document({"type": "text", "media_type": "text/plain", "data": text})
def chat_file(file_data: JsonValue, filename: JsonValue = "document.pdf") -> dict[str, JsonValue]:
return {"type": "file", "file": {"filename": filename, "file_data": file_data}}
def responses_input_file(file_data: JsonValue, filename: JsonValue = "document.pdf") -> dict[str, JsonValue]:
return {"type": "input_file", "filename": filename, "file_data": file_data}
def messages_body(model: str, blocks: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": [*blocks, {"type": "text", "text": text}]}],
**fields,
}
def chat_body(model: str, parts: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": [*parts, {"type": "text", "text": text}]}],
**fields,
}
def responses_body(model: str, items: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]:
return {
"model": model,
"max_output_tokens": 64,
"input": [{"role": "user", "content": [*items, {"type": "input_text", "text": text}]}],
**fields,
}

View file

@ -17,6 +17,8 @@ from integration._support.process import LEGACY_MIGRATE_DEPLOY, MIGRATE_DEPLOY,
from psycopg import sql
from psycopg.rows import class_row
pytestmark: Final = pytest.mark.timeout(240)
REPO_ROOT: Final = Path(__file__).resolve().parents[3]
PRISMA_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras"
PARTITION_SCRIPT: Final = REPO_ROOT / "db_scripts" / "partition_spend_logs.sql"

View file

@ -0,0 +1,129 @@
import uuid
from typing import Final
from pydantic import JsonValue
from tests.integration._support.client import Gateway, Scenario
from tests.integration._support.database import read_rows
def _credential_with_display_name(gateway: Gateway, scenario: Scenario, display_name: str) -> str:
name: Final = f"credential-{uuid.uuid4().hex}"
gateway.post(
"/credentials",
{
"credential_name": name,
"display_name": display_name,
"credential_values": {"api_key": "synthetic-credential"},
"credential_info": {"custom_llm_provider": "openai"},
},
)
scenario.cleanups.callback(_delete_credential_if_present, gateway, name)
return name
def _delete_credential_if_present(gateway: Gateway, name: str) -> None:
response: Final = gateway.request("DELETE", f"/credentials/{name}")
assert response.status_code in (200, 404), response.text
def _stored_display_name(name: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT credential_name, display_name FROM "LiteLLM_CredentialsTable" WHERE credential_name = %s',
(name,),
)
def _served_credential(gateway: Gateway, name: str) -> dict[str, JsonValue]:
by_name: Final = gateway.get(f"/credentials/by_name/{name}")
listed: Final = [entry for entry in gateway.get("/credentials")["credentials"] if isinstance(entry, dict)]
from_list: Final = [entry for entry in listed if entry["credential_name"] == name]
assert len(from_list) == 1, listed
assert from_list[0]["display_name"] == by_name["display_name"], (from_list[0], by_name)
return by_name
def _model_using(gateway: Gateway, scenario: Scenario, credential: str) -> str:
return scenario.model(
model="openai/gpt-4o-mini",
api_base=f"{gateway.upstream_url}/v1",
litellm_credential_name=credential,
)
def test_display_name_round_trips_and_patch_keeps_clears_and_rejects_blank(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
name: Final = _credential_with_display_name(gateway, scenario, "Prod OpenAI")
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Prod OpenAI"}]
assert _served_credential(gateway, name)["display_name"] == "Prod OpenAI"
kept: Final = gateway.request("PATCH", f"/credentials/{name}", {"credential_info": {"description": "d"}})
assert kept.status_code == 200, kept.text
assert _served_credential(gateway, name)["display_name"] == "Prod OpenAI"
relabeled: Final = gateway.request(
"PATCH", f"/credentials/{name}", {"display_name": "Staging OpenAI", "credential_info": {}}
)
assert relabeled.status_code == 200, relabeled.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Staging OpenAI"}]
assert _served_credential(gateway, name)["display_name"] == "Staging OpenAI"
blank: Final = gateway.request("PATCH", f"/credentials/{name}", {"display_name": " ", "credential_info": {}})
assert blank.status_code == 400, blank.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Staging OpenAI"}]
too_long: Final = gateway.request(
"PATCH", f"/credentials/{name}", {"display_name": "x" * 256, "credential_info": {}}
)
assert too_long.status_code == 400, too_long.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Staging OpenAI"}]
trimmed: Final = gateway.request(
"PATCH", f"/credentials/{name}", {"display_name": " Staging OpenAI v2 ", "credential_info": {}}
)
assert trimmed.status_code == 200, trimmed.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Staging OpenAI v2"}]
assert _served_credential(gateway, name)["source"] == "db"
cleared: Final = gateway.request("PATCH", f"/credentials/{name}", {"display_name": None, "credential_info": {}})
assert cleared.status_code == 200, cleared.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": None}]
assert _served_credential(gateway, name)["display_name"] is None
def test_blank_display_name_on_create_is_rejected(gateway: Gateway) -> None:
name: Final = f"credential-{uuid.uuid4().hex}"
rejected: Final = gateway.request(
"POST",
"/credentials",
{"credential_name": name, "display_name": "", "credential_values": {"api_key": "k"}, "credential_info": {}},
)
assert rejected.status_code == 400, rejected.text
assert _stored_display_name(name) == []
def test_patch_rename_is_rejected_and_model_keeps_working(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
name: Final = _credential_with_display_name(gateway, scenario, "Prod OpenAI")
model: Final = _model_using(gateway, scenario, name)
assert gateway.chat(model)["object"] == "chat.completion"
other: Final = f"credential-{uuid.uuid4().hex}"
renamed: Final = gateway.request(
"PATCH", f"/credentials/{name}", {"credential_name": other, "credential_info": {"description": "x"}}
)
assert renamed.status_code == 400, renamed.text
assert "display_name" in renamed.text, renamed.text
assert _stored_display_name(name) == [{"credential_name": name, "display_name": "Prod OpenAI"}]
assert _stored_display_name(other) == []
assert gateway.request("GET", f"/credentials/by_name/{other}").status_code == 404
same_name: Final = gateway.request(
"PATCH", f"/credentials/{name}", {"credential_name": name, "credential_info": {"description": "y"}}
)
assert same_name.status_code == 200, same_name.text
assert _served_credential(gateway, name)["credential_info"] == {
"custom_llm_provider": "openai",
"description": "y",
}
assert gateway.chat(model)["object"] == "chat.completion"

View file

@ -0,0 +1,395 @@
import json
import os
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from pydantic import JsonValue
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
_KEY: Final = "sk-interaction-test"
_PROXY_PYTHONPATH: Final = os.pathsep.join(
(str(Path(__file__).resolve().parents[3]), str(Path(__file__).resolve().parents[2]))
)
_META: Final = {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientInfo": {"name": "continuation-test", "version": "1"},
"io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}, "url": {}}},
}
def interaction_peer(request: Request) -> Reply:
if request.method != "POST":
return Reply(status=405)
body: Final = json.loads(request.body)
method: Final = body["method"]
params: Final = body.get("params", {})
if method == "server/discover":
result = {
"supportedVersions": ["2026-07-28"],
"capabilities": {"tools": {}, "prompts": {}, "resources": {}},
"cacheScope": "private",
"ttlMs": 0,
}
elif method == "tools/list":
result = {
"tools": [{"name": "confirm", "inputSchema": {"type": "object"}}],
"cacheScope": "private",
"ttlMs": 0,
}
elif method == "prompts/list":
result = {"prompts": [{"name": "confirm"}], "cacheScope": "private", "ttlMs": 0}
elif method == "resources/list":
result = {"resources": [{"name": "confirm", "uri": "test://confirm"}], "cacheScope": "private", "ttlMs": 0}
elif method == "resources/templates/list":
result = {"resourceTemplates": [], "cacheScope": "private", "ttlMs": 0}
elif not params.get("requestState"):
return Reply(
body=json.dumps(
{
"jsonrpc": "2.0",
"id": body["id"],
"result": {
"resultType": "input_required",
"requestState": "opaque:" + method,
"inputRequests": {
"consent": {
"method": "elicitation/create",
"params": {
"mode": "form",
"message": "Confirm",
"requestedSchema": {"type": "object", "properties": {}},
},
}
},
},
}
).encode()
)
else:
assert params["requestState"] == "opaque:" + method
assert params["inputResponses"] == {"consent": {"action": "accept"}}
if method == "tools/call":
result = {"content": [{"type": "text", "text": "confirmed"}], "isError": False}
elif method == "prompts/get":
result = {"messages": [{"role": "user", "content": {"type": "text", "text": "confirmed"}}]}
else:
assert method == "resources/read"
result = {"contents": [{"uri": "test://confirm", "text": "confirmed"}]}
return Reply(
body=json.dumps({"jsonrpc": "2.0", "id": body["id"], "result": {"resultType": "complete", **result}}).encode()
)
def rpc(gateway: Gateway, method: str, params: dict[str, JsonValue], *, status: int = 200) -> dict:
response: Final = gateway.client.post(
"/mcp/",
headers={
"Authorization": "Bearer " + gateway.key,
"MCP-Protocol-Version": "2026-07-28",
"Mcp-Method": method,
"Mcp-Name": str(params.get("name", params.get("uri", ""))),
"Accept": "application/json, text/event-stream",
},
json={"jsonrpc": "2.0", "id": 1, "method": method, "params": {**params, "_meta": _META}},
)
assert response.status_code == status, response.text
return response.json()
@pytest.mark.parametrize("changed_target", [False, True])
def test_continuations_resume_on_another_replica_and_reject_changed_operations(
tmp_path: Path, changed_target: bool
) -> None:
with wire_server(interaction_peer) as peer, httpx.Client() as client:
config: Final = tmp_path / "proxy.yaml"
config.write_text(
yaml.safe_dump(
{
"model_list": [],
"mcp_servers": {
"mrtr": {
"url": peer.url,
"transport": "http",
"protocol_version": "2026-07-28",
"allow_elicitation": True,
}
},
"general_settings": {
"master_key": _KEY,
"store_model_in_db": False,
"mcp_advertised_versions": ["2025-11-25", "2026-07-28"],
},
}
)
)
other_config: Final = tmp_path / "other-proxy.yaml"
other_config.write_text(
config.read_text().replace(peer.url, peer.url + "/changed-target") if changed_target else config.read_text()
)
seed: Final = Gateway(client, _KEY, peer.url)
environment: Final = {
"PYTHONPATH": _PROXY_PYTHONPATH,
"STORE_MODEL_IN_DB": "False",
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "shared-interaction-test",
}
options: Final = {
"database_setup": (),
"remove_environment": (
"DATABASE_URL",
"DATABASE_URL_READ_REPLICA",
"LITELLM_LICENSE",
"LITELLM_LICENSE_PATH",
"REDIS_URL",
"REDIS_HOST",
),
}
with (
owned_proxy(seed, tmp_path / "a", environment, config=config, **options) as first,
owned_proxy(seed, tmp_path / "b", environment, config=other_config, **options) as second,
):
for method, params, terminal_field in (
("tools/call", {"name": "mrtr-confirm", "arguments": {}}, "content"),
("prompts/get", {"name": "mrtr-confirm", "arguments": {}}, "messages"),
("resources/read", {"uri": "test://confirm"}, "contents"),
):
initial: Final = rpc(first, method, params)
assert initial["result"]["resultType"] == "input_required", initial
state: Final = initial["result"]["requestState"]
assert state.startswith("mcp_state_v1."), initial
retry: Final = {**params, "requestState": state, "inputResponses": {"consent": {"action": "accept"}}}
if changed_target:
peer.drain()
refused: Final = rpc(second, method, retry, status=400)
assert refused["error"]["code"] == -32602, refused
assert peer.drain() == (), "Changed target must reject before upstream dispatch"
continue
completed: Final = rpc(second, method, retry)
assert "confirmed" in json.dumps(completed["result"][terminal_field]), completed
assert rpc(first, method, retry) == completed
peer.drain()
changed: Final = {
**retry,
**({"uri": "test://other"} if method == "resources/read" else {"name": "mrtr-other"}),
}
rejected: Final = rpc(second, method, changed, status=400)
assert rejected["error"]["code"] == -32602, rejected
assert peer.drain() == (), "Rejected continuation must not contact the upstream"
def test_continuation_reauthenticates_caller_and_rechecks_revoked_permissions(tmp_path: Path) -> None:
auth_module: Final = tmp_path / "interaction_auth.py"
auth_module.write_text(
"from pathlib import Path\n"
"from fastapi import HTTPException, Request\n"
"from litellm.proxy._types import UserAPIKeyAuth\n"
"async def authenticate(request: Request, api_key: str) -> UserAPIKeyAuth:\n"
" if api_key != 'sk-interaction-test':\n"
" raise HTTPException(status_code=401, detail='Unknown test caller')\n"
" return UserAPIKeyAuth.model_validate_json(Path(__file__).with_suffix('.json').read_text())\n"
)
identity: Final = {
"user_id": "alice",
"team_id": "team-a",
"user_role": "internal_user",
"object_permission": {"object_permission_id": "test-permission", "mcp_servers": ["interaction-server"]},
}
auth_state: Final = auth_module.with_suffix(".json")
auth_state.write_text(json.dumps(identity))
with wire_server(interaction_peer) as peer, wire_server(interaction_peer) as other_peer, httpx.Client() as client:
config: Final = tmp_path / "proxy.yaml"
config.write_text(
yaml.safe_dump(
{
"model_list": [],
"mcp_servers": {
"mrtr": {
"server_id": "interaction-server",
"url": peer.url,
"transport": "http",
"protocol_version": "2026-07-28",
"allow_elicitation": True,
},
"other": {
"server_id": "other-server",
"url": other_peer.url,
"transport": "http",
"protocol_version": "2026-07-28",
"allow_elicitation": True,
},
},
"general_settings": {
"master_key": _KEY,
"custom_auth": "interaction_auth.authenticate",
"store_model_in_db": False,
"mcp_advertised_versions": ["2025-11-25", "2026-07-28"],
},
}
)
)
with owned_proxy(
Gateway(client, _KEY, peer.url),
tmp_path / "proxy",
{
"PYTHONPATH": _PROXY_PYTHONPATH,
"STORE_MODEL_IN_DB": "False",
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "caller-test-salt",
},
config=config,
database_setup=(),
remove_environment=("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "REDIS_URL", "REDIS_HOST"),
) as gateway:
params: Final = {"name": "mrtr-confirm", "arguments": {}}
initial: Final = rpc(gateway, "tools/call", params)
assert initial["result"]["resultType"] == "input_required", initial
retry: Final = {
**params,
"requestState": initial["result"]["requestState"],
"inputResponses": {"consent": {"action": "accept"}},
}
for changed in ({"user_id": "bob"}, {"team_id": "team-b"}):
auth_state.write_text(json.dumps({**identity, **changed}))
peer.drain()
rejected: Final = rpc(gateway, "tools/call", retry, status=400)
assert rejected["error"]["code"] == -32602, rejected
assert peer.drain() == (), "Caller-bound state must reject before contacting upstream"
auth_state.write_text(json.dumps(identity))
resumed: Final = rpc(gateway, "tools/call", retry)
assert resumed["result"]["content"] == [{"type": "text", "text": "confirmed"}], resumed
auth_state.write_text(
json.dumps(
{
**identity,
"object_permission": {
"object_permission_id": "test-permission",
"mcp_servers": ["denied-server"],
},
}
)
)
peer.drain()
revoked: Final = rpc(gateway, "tools/call", retry, status=403)
assert revoked["detail"] == "MCP continuation target is no longer authorized", revoked
assert peer.drain() == (), "Revoked access must reject a valid continuation before upstream dispatch"
auth_state.write_text(json.dumps(identity))
resource: Final = rpc(gateway, "resources/read", {"uri": "test://confirm"})
assert resource["result"]["resultType"] == "input_required", resource
auth_state.write_text(
json.dumps(
{
**identity,
"object_permission": {
"object_permission_id": "test-permission",
"mcp_servers": ["other-server"],
},
}
)
)
peer.drain()
other_peer.drain()
response: Final = gateway.client.post(
"/mcp/",
headers={
"Authorization": "Bearer " + gateway.key,
"MCP-Protocol-Version": "2026-07-28",
"Mcp-Method": "resources/read",
"Mcp-Name": "test://confirm",
"Accept": "application/json, text/event-stream",
},
json={
"jsonrpc": "2.0",
"id": 1,
"method": "resources/read",
"params": {
"uri": "test://confirm",
"_meta": _META,
"requestState": resource["result"]["requestState"],
"inputResponses": {"consent": {"action": "accept"}},
},
},
)
other_requests: Final = tuple(
(json.loads(request.body)["method"], json.loads(request.body).get("params", {}).get("requestState"))
for request in other_peer.drain()
)
assert other_requests == (), "Continuation must never send upstream state to another authorized server"
assert peer.drain() == (), "Revoked original target must not receive a retry"
assert response.status_code >= 400 or "error" in response.json(), response.text
auth_state.write_text(
json.dumps(
{
**identity,
"object_permission": {
"object_permission_id": "test-permission",
"mcp_servers": ["interaction-server", "other-server"],
},
}
)
)
completed: Final = rpc(
gateway,
"resources/read",
{
"uri": "test://confirm",
"requestState": resource["result"]["requestState"],
"inputResponses": {"consent": {"action": "accept"}},
},
)
assert completed["result"]["contents"] == [{"uri": "test://confirm", "text": "confirmed"}], completed
assert other_peer.drain() == (), "Expanded access must keep the continuation on its original target"
def test_missing_continuation_key_reports_configuration_for_each_carrier(tmp_path: Path) -> None:
with wire_server(interaction_peer) as peer, httpx.Client() as client:
config: Final = tmp_path / "proxy.yaml"
config.write_text(
yaml.safe_dump(
{
"model_list": [],
"mcp_servers": {
"mrtr": {
"url": peer.url,
"transport": "http",
"protocol_version": "2026-07-28",
"allow_elicitation": True,
}
},
"general_settings": {
"master_key": _KEY,
"store_model_in_db": False,
"mcp_advertised_versions": ["2025-11-25", "2026-07-28"],
},
}
)
)
with owned_proxy(
Gateway(client, _KEY, peer.url),
tmp_path / "proxy",
{
"PYTHONPATH": _PROXY_PYTHONPATH,
"STORE_MODEL_IN_DB": "False",
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "",
},
config=config,
database_setup=(),
remove_environment=("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "REDIS_URL", "REDIS_HOST"),
) as gateway:
for method, params in (
("tools/call", {"name": "mrtr-confirm", "arguments": {}}),
("prompts/get", {"name": "mrtr-confirm", "arguments": {}}),
("resources/read", {"uri": "test://confirm"}),
):
response: Final = rpc(gateway, method, params, status=400)
assert response["error"]["code"] == -32602, response
assert "LITELLM_SALT_KEY" in response["error"]["message"], response

View file

@ -0,0 +1,226 @@
import asyncio
import os
import re
import signal
import threading
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.pdf_document import (
COUNT_REFUSED,
COUNT_TOKENS_TARGET,
LETTER,
messages_body,
pdf_document,
rendered_tokens,
)
from integration._support.process import graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from integration.providers._cache_control_marks_support import anthropic_peer, owned_config
from pydantic import JsonValue, TypeAdapter
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_MODEL: Final = "burst-claude"
_PROVIDER_KEY: Final = "synthetic-provider-key"
_ASK: Final = "Summarize the attached report in one sentence."
_COUNT_BURST: Final = 24
_MESSAGE_BURST: Final = 12
_PAGE_COUNTS: Final = (1, 3, 12)
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
@dataclass(frozen=True, slots=True)
class _Counted:
pages: int
status: int
text: str
input_tokens: int | None
@dataclass(frozen=True, slots=True)
class _Sent:
marker: str
status: int
text: str
call_id: str
def _peer(request: Request) -> Reply:
if request.target == COUNT_TOKENS_TARGET:
return COUNT_REFUSED
return anthropic_peer(request)
def _held(release: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert release.wait(timeout=120), "Held peer was never released"
return _peer(request)
return respond
def _deployment(api_base: str) -> dict[str, JsonValue]:
return {
"model_name": _MODEL,
"litellm_params": {"model": "anthropic/claude-opus-5-5", "api_base": api_base, "api_key": _PROVIDER_KEY},
}
def _count_body(pages: int) -> dict[str, JsonValue]:
blocks: Final[list[JsonValue]] = [pdf_document((LETTER,) * pages)] if pages else []
return {"model": _MODEL, "messages": messages_body(_MODEL, blocks, _ASK)["messages"]}
def _input_tokens(response: httpx.Response) -> int | None:
if response.status_code != 200:
return None
counted: Final = _JSON_OBJECT.validate_json(response.content).get("input_tokens")
return counted if isinstance(counted, int) else None
async def _fire_counts(url: str, key: str, *, tolerate_transport_errors: bool = False) -> tuple[_Counted, ...]:
async def one(client: httpx.AsyncClient, index: int) -> _Counted:
pages: Final = _PAGE_COUNTS[index % len(_PAGE_COUNTS)]
response: Final = await client.post(
"/v1/messages/count_tokens", json=_count_body(pages), headers={"Authorization": f"Bearer {key}"}
)
return _Counted(pages, response.status_code, response.text, _input_tokens(response))
async with httpx.AsyncClient(base_url=url, timeout=90, trust_env=False) as client:
results: Final = await asyncio.gather(
*(one(client, index) for index in range(_COUNT_BURST)), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Counted))
async def _fire_messages(url: str, key: str) -> tuple[_Sent, ...]:
async def one(client: httpx.AsyncClient) -> _Sent:
marker: Final = uuid.uuid4().hex
body: Final = messages_body(_MODEL, [pdf_document((LETTER,))], f"{_ASK} marker-{marker}")
response: Final = await client.post("/v1/messages", json=body, headers={"Authorization": f"Bearer {key}"})
return _Sent(marker, response.status_code, response.text, response.headers.get("x-litellm-call-id", ""))
async with httpx.AsyncClient(base_url=url, timeout=90, trust_env=False) as client:
return tuple(await asyncio.gather(*(one(client) for _ in range(_MESSAGE_BURST))))
def _baseline(gateway: Gateway) -> int:
response: Final = gateway.request("POST", "/v1/messages/count_tokens", _count_body(0))
assert response.status_code == 200, response.text
counted: Final = _input_tokens(response)
assert counted is not None, response.text
return counted
def _assert_exact(counts: tuple[_Counted, ...], baseline: int) -> None:
for item in counts:
assert item.status == 200, (item.pages, item.status, item.text)
assert item.input_tokens is not None and item.input_tokens - baseline == rendered_tokens(
(LETTER,) * item.pages
), (
item.pages,
item.input_tokens,
baseline,
)
def _single_spend_row(item: _Sent) -> None:
assert item.status == 200, (item.status, item.text)
assert item.call_id, item.text
rows: Final = eventually(
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (item.call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert len(rows) == 1, item.call_id
def _held_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(2 * graceful_stop_seconds() + 180)
async def test_pdf_count_burst_across_two_workers_prices_every_page_and_logs_each_message_once(
gateway: Gateway, tmp_path: Path
) -> None:
with wire_server(_peer) as wire:
config: Final = owned_config(tmp_path, [_deployment(wire.url)], litellm_settings={"cache": False})
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
eventually(
lambda: _STARTED_WORKER.findall(owned.log.read_text()),
lambda pids: len(pids) == 2,
seconds=graceful_stop_seconds(),
)
owned_url: Final = str(owned.gateway.client.base_url)
baseline: Final = _baseline(owned.gateway)
wire.drain()
counts, messages = await asyncio.gather(
_fire_counts(owned_url, owned.gateway.key), _fire_messages(owned_url, owned.gateway.key)
)
received: Final = wire.drain()
targets: Final = [request.target for request in received]
assert targets.count(COUNT_TOKENS_TARGET) == _COUNT_BURST, targets
assert targets.count("/v1/messages") == _MESSAGE_BURST, targets
_assert_exact(counts, baseline)
assert {item.marker for item in messages} == {
_JSON_OBJECT.validate_json(item.text)["id"][len("msg_") :] for item in messages if item.status == 200
}
for item in messages:
_single_spend_row(item)
@pytest.mark.timeout(2 * graceful_stop_seconds() + 180)
async def test_worker_sigkill_mid_pdf_count_burst_leaves_the_sibling_pricing_pages(
gateway: Gateway, tmp_path: Path
) -> None:
release: Final = threading.Event()
with wire_server(_held(release)) as wire:
config: Final = owned_config(tmp_path, [_deployment(wire.url)], litellm_settings={"cache": False})
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
workers: Final[tuple[int, ...]] = eventually(
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=graceful_stop_seconds(),
)
owned_url: Final = str(owned.gateway.client.base_url)
burst: Final = asyncio.create_task(
_fire_counts(owned_url, owned.gateway.key, tolerate_transport_errors=True)
)
try:
await asyncio.to_thread(
eventually, lambda: wire.received.qsize(), lambda size: size >= _COUNT_BURST, 90
)
held_by: Final = MappingProxyType({pid: _held_upstream_connections(pid, wire.url) for pid in workers})
victim: Final = max(workers, key=held_by.__getitem__)
os.kill(victim, signal.SIGKILL)
finally:
release.set()
served: Final = await burst
wire.drain()
baseline: Final = _baseline(owned.gateway)
after: Final = await _fire_counts(owned_url, owned.gateway.key)
after_received: Final = wire.drain()
assert sum(held_by.values()) == _COUNT_BURST, held_by
assert held_by[victim] > 0, held_by
assert len(served) == _COUNT_BURST - held_by[victim], (len(served), held_by)
assert len(after) == _COUNT_BURST, len(after)
assert [request.target for request in after_received] == [COUNT_TOKENS_TARGET] * (_COUNT_BURST + 1)
_assert_exact(served, baseline)
_assert_exact(after, baseline)

View file

@ -0,0 +1,273 @@
import io
import json
import uuid
from collections.abc import Sequence
from types import MappingProxyType
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.pdf_document import (
COUNT_REFUSED,
COUNT_TOKENS_TARGET,
LETTER,
NARROW,
Page,
base64_source,
chat_file,
document,
encoded,
messages_body,
pdf_bytes,
pdf_data_url,
pdf_document,
rendered_tokens,
responses_input_file,
text_document,
)
from integration._support.provider import SharedProvider
from integration._support.wire import Reply
from pydantic import JsonValue, TypeAdapter
from pypdf import PdfReader, PdfWriter
_MODEL: Final = "anthropic/claude-opus-5-5"
_ASK: Final = "Summarize the attached report in one sentence."
_TEXT: Final = "The quarterly report covers revenue, margins and headcount."
_TITLE: Final = "Quarterly report"
_CONTEXT: Final = "Board pack, page one"
_PEER_INPUT_TOKENS: Final = 12
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_PNG_HEADER: Final = b"\x89PNG\r\n\x1a\n\x00\x00\x00\x0dIHDR" + (1568).to_bytes(4, "big") * 2 + b"\x08\x06\x00\x00\x00"
def _locked(raw: bytes, user_password: str) -> bytes:
writer: Final = PdfWriter(clone_from=PdfReader(io.BytesIO(raw)))
writer.encrypt(user_password=user_password, owner_password="owner", algorithm="RC4-128")
out: Final = io.BytesIO()
writer.write(out)
return out.getvalue()
_UNREADABLE: Final[MappingProxyType[str, JsonValue]] = MappingProxyType(
{
"garbage-5kb": encoded(bytes(range(256)) * 20),
"cut-mid-stream": encoded(pdf_bytes((LETTER, LETTER))[:200]),
"png-bytes": encoded(_PNG_HEADER),
"user-password": encoded(_locked(pdf_bytes((LETTER, LETTER)), "reader")),
"empty": "",
"int": 1234,
}
)
def _int(value: JsonValue) -> int:
assert isinstance(value, int), value
return value
def _turn(blocks: Sequence[JsonValue]) -> list[JsonValue]:
return messages_body(_MODEL, blocks, _ASK)["messages"]
def _count(gateway: Gateway, provider: SharedProvider, blocks: Sequence[JsonValue]) -> int:
provider.expect(COUNT_REFUSED)
response: Final = gateway.request("POST", "/v1/messages/count_tokens", {"model": _MODEL, "messages": _turn(blocks)})
assert response.status_code == 200, response.text
assert [(request.method, request.target) for request in provider.received()] == [("POST", COUNT_TOKENS_TARGET)]
return _int(_JSON_OBJECT.validate_json(response.content)["input_tokens"])
def _local(gateway: Gateway, parts: Sequence[JsonValue]) -> int:
response: Final = gateway.request(
"POST", "/utils/token_counter", {"model": _MODEL, "messages": _turn(parts)}, params={"call_endpoint": "false"}
)
assert response.status_code == 200, response.text
return _int(_JSON_OBJECT.validate_json(response.content)["total_tokens"])
def _responses_count(gateway: Gateway, provider: SharedProvider, items: Sequence[JsonValue]) -> int:
provider.expect(COUNT_REFUSED)
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": _MODEL, "input": [{"role": "user", "content": [*items, {"type": "input_text", "text": _ASK}]}]},
)
assert response.status_code == 200, response.text
assert [(request.method, request.target) for request in provider.received()] == [("POST", COUNT_TOKENS_TARGET)]
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["object"] == "response.input_tokens", response.text
return _int(payload["input_tokens"])
def _cost(gateway: Gateway, parts: Sequence[JsonValue]) -> float:
payload: Final = gateway.post("/spend/calculate", {"model": _MODEL, "messages": _turn(parts)})
cost: Final = payload["cost"]
assert isinstance(cost, (int, float)), payload
return float(cost)
def _input_price(gateway: Gateway) -> float:
listed: Final = gateway.get("/model/info")["data"]
assert isinstance(listed, list), listed
rows: Final = [object_value(row) for row in listed if isinstance(row, dict) and row.get("model_name") == _MODEL]
assert len(rows) == 1, rows
price: Final = object_value(rows[0]["model_info"])["input_cost_per_token"]
assert isinstance(price, float) and price > 0, price
return price
def _peer_message(identity: str) -> Reply:
return Reply(
body=json.dumps(
{
"id": identity,
"type": "message",
"role": "assistant",
"model": "claude-opus-5-5",
"content": [{"type": "text", "text": "One sentence."}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": _PEER_INPUT_TOKENS, "output_tokens": 3},
}
).encode()
)
def _prompt_tokens(call_id: str) -> int:
rows: Final = eventually(
lambda: read_rows('SELECT prompt_tokens FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
return _int(rows[0]["prompt_tokens"])
@pytest.mark.parametrize("pages", [1, 3, 12])
def test_count_tokens_fallback_prices_each_blank_page_as_a_rendered_image(
gateway: Gateway, provider: SharedProvider, pages: int
) -> None:
letter_pages: Final = (LETTER,) * pages
with_document: Final = _count(gateway, provider, [pdf_document(letter_pages)])
without: Final = _count(gateway, provider, [])
assert with_document - without == rendered_tokens(letter_pages), (with_document, without)
def test_count_tokens_fallback_adds_a_page_text_to_its_rendered_image(
gateway: Gateway, provider: SharedProvider
) -> None:
text_page: Final = _count(gateway, provider, [pdf_document((Page(text=_TEXT),))])
blank_page: Final = _count(gateway, provider, [pdf_document((LETTER,))])
text_source: Final = _count(gateway, provider, [text_document(_TEXT)])
without: Final = _count(gateway, provider, [])
assert text_source - without > 0, (text_source, without)
assert text_page - blank_page == text_source - without, (text_page, blank_page, text_source, without)
def test_count_tokens_fallback_prices_a_page_by_its_rendered_area(gateway: Gateway, provider: SharedProvider) -> None:
mixed: Final = (LETTER, NARROW, LETTER)
without: Final = _count(gateway, provider, [])
narrow: Final = _count(gateway, provider, [pdf_document((NARROW,))])
assert narrow - without == rendered_tokens((NARROW,)), (narrow, without)
assert _count(gateway, provider, [pdf_document(mixed)]) - without == rendered_tokens(mixed)
def test_count_tokens_fallback_prices_every_document_in_the_turn(gateway: Gateway, provider: SharedProvider) -> None:
both: Final = _count(gateway, provider, [pdf_document((LETTER,) * 3), pdf_document((LETTER,))])
without: Final = _count(gateway, provider, [])
assert both - without == rendered_tokens((LETTER,) * 4), (both, without)
def test_count_tokens_fallback_prices_a_pdf_without_pages_as_nothing(
gateway: Gateway, provider: SharedProvider
) -> None:
assert _count(gateway, provider, [pdf_document(())]) == _count(gateway, provider, [])
def test_count_tokens_fallback_adds_title_and_context_once_per_document(
gateway: Gateway, provider: SharedProvider
) -> None:
without: Final = _count(gateway, provider, [])
plain: Final = _count(gateway, provider, [pdf_document((LETTER,))])
annotated: Final = _count(gateway, provider, [pdf_document((LETTER,), title=_TITLE, context=_CONTEXT)])
title: Final = _count(gateway, provider, [text_document(_TITLE)]) - without
context: Final = _count(gateway, provider, [text_document(_CONTEXT)]) - without
assert title > 0 and context > 0, (title, context)
assert annotated - plain == title + context, (annotated, plain, title, context)
@pytest.mark.parametrize("label", list(_UNREADABLE))
def test_count_tokens_fallback_prices_unreadable_pdf_bytes_like_one_image(
gateway: Gateway, provider: SharedProvider, label: str
) -> None:
data: Final = _UNREADABLE[label]
as_pdf: Final = _count(gateway, provider, [document(base64_source(data))])
as_png: Final = _count(gateway, provider, [document(base64_source(data, media_type="image/png"))])
without: Final = _count(gateway, provider, [])
assert as_pdf == as_png, (as_pdf, as_png)
assert 0 <= as_pdf - without < rendered_tokens((LETTER,)), (as_pdf, without)
def test_count_tokens_fallback_reads_a_list_wrapped_base64_string_like_the_bare_string(
gateway: Gateway, provider: SharedProvider
) -> None:
raw: Final = encoded(pdf_bytes((LETTER,)))
wrapped: Final = _count(gateway, provider, [document(base64_source([raw]))])
assert wrapped == _count(gateway, provider, [document(base64_source(raw))]), wrapped
def test_count_tokens_fallback_reads_an_owner_locked_pdf(gateway: Gateway, provider: SharedProvider) -> None:
pages: Final = (LETTER, LETTER)
locked: Final = document(base64_source(encoded(_locked(pdf_bytes(pages), ""))))
assert _count(gateway, provider, [locked]) - _count(gateway, provider, []) == rendered_tokens(pages)
def test_count_tokens_fallback_parses_only_the_pdf_media_type(gateway: Gateway, provider: SharedProvider) -> None:
raw: Final = encoded(pdf_bytes((LETTER,)))
without: Final = _count(gateway, provider, [])
labelled: Final = _count(gateway, provider, [document(base64_source(raw))])
assert labelled - without == rendered_tokens((LETTER,)), (labelled, without)
mislabelled: Final = _count(gateway, provider, [document(base64_source(raw, media_type="application/x-pdf"))])
assert mislabelled == _count(gateway, provider, [document(base64_source(raw, media_type="image/png"))])
def test_utils_token_counter_prices_a_chat_file_by_its_pages(gateway: Gateway) -> None:
one: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,)))])
three: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,) * 3))])
assert three - one == rendered_tokens((LETTER,) * 2), (three, one)
text_page: Final = _local(gateway, [chat_file(pdf_data_url((Page(text=_TEXT),)))])
text_part: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,))), {"type": "text", "text": _TEXT}])
assert text_part - one > 0, (text_part, one)
assert text_page - one == text_part - one, (text_page, text_part, one)
def test_responses_input_tokens_fallback_prices_an_input_file_by_its_pages(
gateway: Gateway, provider: SharedProvider
) -> None:
one: Final = _responses_count(gateway, provider, [responses_input_file(pdf_data_url((LETTER,)))])
three: Final = _responses_count(gateway, provider, [responses_input_file(pdf_data_url((LETTER,) * 3))])
assert three - one == rendered_tokens((LETTER,) * 2), (three, one)
def test_spend_calculate_prices_a_chat_file_by_its_pages(gateway: Gateway) -> None:
twelve: Final = _cost(gateway, [chat_file(pdf_data_url((LETTER,) * 12))])
one: Final = _cost(gateway, [chat_file(pdf_data_url((LETTER,)))])
assert twelve - one == pytest.approx(rendered_tokens((LETTER,) * 11) * _input_price(gateway)), (twelve, one)
def test_a_cached_pdf_message_keeps_the_peer_usage_on_both_spend_rows(
gateway: Gateway, provider: SharedProvider
) -> None:
marker: Final = uuid.uuid4().hex
body: Final = messages_body(_MODEL, [pdf_document((LETTER,))], f"{_ASK} marker-{marker}")
provider.expect(_peer_message(f"msg_{marker}"))
first: Final = gateway.request("POST", "/v1/messages", body)
second: Final = gateway.request("POST", "/v1/messages", body)
assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text)
assert [(request.method, request.target) for request in provider.received()] == [("POST", "/v1/messages")]
assert _JSON_OBJECT.validate_json(second.content)["id"] == f"msg_{marker}", second.text
assert first.headers["x-litellm-call-id"] != second.headers["x-litellm-call-id"]
assert [_prompt_tokens(response.headers["x-litellm-call-id"]) for response in (first, second)] == [
_PEER_INPUT_TOKENS,
_PEER_INPUT_TOKENS,
]

View file

@ -0,0 +1,159 @@
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
import anthropic
import httpx
import openai
from integration._support.client import Gateway
from pydantic import JsonValue, TypeAdapter
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
TOKEN_COUNTING_BETA: Final = "token-counting-2024-11-01"
ANTHROPIC_VERSION: Final = "2023-06-01"
REJECTION: Final = 'messages.0: Unexpected role "system". The Messages API accepts a top-level `system` parameter'
USER_TEXT: Final = "Count this message"
USER: Final[dict[str, JsonValue]] = {"role": "user", "content": USER_TEXT}
ASSISTANT: Final[dict[str, JsonValue]] = {"role": "assistant", "content": "One."}
FOLLOW_UP: Final[dict[str, JsonValue]] = {"role": "user", "content": "Again"}
INSTRUCTION: Final = "You are a terse assistant"
REMINDER: Final = "Answer in one sentence"
LEADING: Final[dict[str, JsonValue]] = {"role": "system", "content": INSTRUCTION}
LIFTED: Final[dict[str, JsonValue]] = {"type": "text", "text": INSTRUCTION}
MID_SYSTEM: Final[dict[str, JsonValue]] = {"role": "system", "content": REMINDER}
EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
ONE_HOUR: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"}
IMAGE_PART: Final[dict[str, JsonValue]] = {
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
},
}
FIVE_KB: Final = "Answer in one sentence. " * 214
CALLER_SYSTEM: Final = "Prefer metric units"
CALLER_BLOCKS: Final[list[JsonValue]] = [
{"type": "text", "text": CALLER_SYSTEM},
{"type": "text", "text": "Never guess", "cache_control": EPHEMERAL},
]
TOOLS: Final[list[JsonValue]] = [
{
"name": "get_weather",
"description": "Look up the current weather for a city",
"input_schema": {
"type": "object",
"properties": {"city": {"type": "string", "description": "City to look up"}},
"required": ["city"],
},
}
]
@dataclass(frozen=True, slots=True)
class LiftCase:
messages: tuple[dict[str, JsonValue], ...]
system: JsonValue | None
lifted: list[JsonValue] | None
LIFT_CASES: Final[Mapping[str, LiftCase]] = MappingProxyType(
{
"string": LiftCase((LEADING, USER), None, [LIFTED]),
"run_with_cache_control": LiftCase(
(
{"role": "system", "content": INSTRUCTION, "cache_control": ONE_HOUR},
{
"role": "system",
"content": [
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
{"type": "text", "text": ""},
IMAGE_PART,
],
},
USER,
),
None,
[
{"type": "text", "text": INSTRUCTION, "cache_control": ONE_HOUR},
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
],
),
"caller_system_string_first": LiftCase(
(LEADING, USER), CALLER_SYSTEM, [{"type": "text", "text": CALLER_SYSTEM}, LIFTED]
),
"caller_system_blocks_first": LiftCase((LEADING, USER), CALLER_BLOCKS, [*CALLER_BLOCKS, LIFTED]),
"caller_empty_system_dropped": LiftCase((LEADING, USER), "", [LIFTED]),
"empty_content_dropped": LiftCase(({"role": "system", "content": ""}, USER), None, None),
"image_only_content_dropped": LiftCase(({"role": "system", "content": [IMAGE_PART]}, USER), None, None),
"integer_content_dropped": LiftCase(({"role": "system", "content": 5}, USER), None, None),
"integer_text_part_dropped": LiftCase(
({"role": "system", "content": [{"type": "text", "text": 7}]}, USER), None, None
),
"five_kb": LiftCase(({"role": "system", "content": FIVE_KB}, USER), None, [{"type": "text", "text": FIVE_KB}]),
}
)
STRING_CASE: Final = LIFT_CASES["string"]
def count_request(model: str, case: LiftCase) -> dict[str, JsonValue]:
return {"model": model, "messages": list(case.messages), **({} if case.system is None else {"system": case.system})}
def expected_count_body(model: str, case: LiftCase) -> dict[str, JsonValue]:
conversation: Final[list[JsonValue]] = [message for message in case.messages if message["role"] != "system"]
return {"model": model, "messages": conversation, **({} if case.lifted is None else {"system": case.lifted})}
def _opening_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
def _conversation_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant", "system")
def _anthropic_tool(tool: JsonValue) -> bool:
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
def accepts_count_body(body: Mapping[str, JsonValue]) -> bool:
# Anthropic, Azure AI Foundry and Bedrock Mantle count_tokens verdicts observed live on 2026-10-09: an empty
# messages list and a system role at messages[0] answer 400 invalid_request_error, a later system role counts
messages: Final = body.get("messages")
tools: Final = body.get("tools", [])
return (
isinstance(messages, list)
and len(messages) > 0
and _opening_message(messages[0])
and all(map(_conversation_message, messages))
and isinstance(body.get("system", ""), (str, list))
and isinstance(tools, list)
and all(map(_anthropic_tool, tools))
)
def anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
def async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
def openai_client(gateway: Gateway) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(trust_env=False),
)
def async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(trust_env=False),
)

View file

@ -0,0 +1,572 @@
import asyncio
import json
import socket
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import AbstractContextManager, ExitStack
from dataclasses import dataclass
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, cast
import httpcore
import httpx
import psutil
import pytest
from anthropic.types import MessageParam
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._count_tokens_system_lift import (
ANTHROPIC_VERSION,
ASSISTANT,
EPHEMERAL,
FOLLOW_UP,
IMAGE_PART,
INSTRUCTION,
JSON_OBJECT,
LEADING,
LIFT_CASES,
LIFTED,
MID_SYSTEM,
REJECTION,
REMINDER,
STRING_CASE,
TOKEN_COUNTING_BETA,
TOOLS,
USER,
USER_TEXT,
LiftCase,
accepts_count_body,
anthropic_client,
async_anthropic_client,
async_openai_client,
count_request,
expected_count_body,
openai_client,
)
from pydantic import JsonValue, TypeAdapter
_MODEL: Final = "claude-opus-5-5"
_API_KEY: Final = "synthetic-count-tokens-key"
_COUNT: Final = 3131
_PROVIDER_TOKENIZERS: Final = frozenset({"anthropic_api", "azure_ai_anthropic_api"})
_SDK_MESSAGES: Final = cast(
list[MessageParam], [LEADING, USER]
) # cast-ok: the SDK types reject the role the proxy lifts
_STRING_BODY: Final = expected_count_body(_MODEL, STRING_CASE)
_PROXY_MODULE: Final = "integration._support.proxy"
_PROBES_PER_ROUND: Final = 8
_CLIENT_ADDRESS: Final = TypeAdapter(tuple[str, int])
_ARGUMENTS: Final = TypeAdapter(tuple[str, ...])
_NAME: Final = TypeAdapter(str)
@dataclass(frozen=True, slots=True)
class _Provider:
prefix: str
target: str
tokenizer: str
azure: bool
@dataclass(frozen=True, slots=True)
class _Deployment:
provider: _Provider
port: int
model: str
_ANTHROPIC: Final = _Provider("anthropic", "/v1/messages/count_tokens", "anthropic_api", False)
_AZURE: Final = _Provider("azure_ai", "/anthropic/v1/messages/count_tokens", "azure_ai_anthropic_api", True)
_PROVIDERS: Final = MappingProxyType({"anthropic": _ANTHROPIC, "azure_ai": _AZURE})
@pytest.fixture(params=_PROVIDERS.keys())
def provider(request: pytest.FixtureRequest) -> _Provider:
name: Final[object] = request.param # pyright: ignore[reportAny] # pytest types the fixture param as Any
return _PROVIDERS[_NAME.validate_python(name)]
def _json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply:
return Reply(status=status, body=json.dumps(payload).encode())
def _rejected(status: int) -> Reply:
return _json_reply(status, {"type": "error", "error": {"type": "invalid_request_error", "message": REJECTION}})
def _counted(request: Request) -> Reply:
accepted: Final = accepts_count_body(JSON_OBJECT.validate_json(request.body))
return _json_reply(200, {"input_tokens": _COUNT}) if accepted else _rejected(400)
def _rejecting(status: int) -> Callable[[Request], Reply]:
def count(_request: Request) -> Reply:
return _rejected(status)
return count
def _holding(held: SimpleQueue[str], release: threading.Event, seconds: float) -> Callable[[Request], Reply]:
def hold(request: Request) -> Reply:
held.put(request.target)
assert release.wait(timeout=seconds), "Held count was never released"
return _counted(request)
return hold
def _peer(provider: _Provider, count: Callable[[Request], Reply] = _counted) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.target == provider.target:
return count(request)
return _json_reply(404, {"error": f"unscripted target {request.target}"})
return respond
def _listening(deployment: _Deployment, count: Callable[[Request], Reply] = _counted) -> AbstractContextManager[Wire]:
return wire_server(_peer(deployment.provider, count), port=deployment.port)
def _reserved_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return _CLIENT_ADDRESS.validate_python(reserve.getsockname())[1]
def _cmdline(process: psutil.Process) -> tuple[str, ...]:
try:
return _ARGUMENTS.validate_python(process.cmdline())
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
return ()
def _serves(cmdline: tuple[str, ...], proxy_port: int) -> bool:
if _PROXY_MODULE not in cmdline or "--port" not in cmdline:
return False
return cmdline[cmdline.index("--port") + 1] == str(proxy_port)
def _listens(process: psutil.Process, proxy_port: int) -> bool:
return any(
connection.status == psutil.CONN_LISTEN and connection.laddr and connection.laddr.port == proxy_port
for connection in process.net_connections(kind="tcp")
)
def _proxy_workers(proxy_port: int) -> frozenset[int]:
(master,) = tuple(process for process in psutil.process_iter() if _serves(_cmdline(process), proxy_port))
spawned: Final = frozenset(child.pid for child in master.children() if _listens(child, proxy_port))
return spawned or frozenset({master.pid})
def _holder(workers: frozenset[int], client_port: int) -> int | None:
def holds(pid: int) -> bool:
return any(
connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == client_port
for connection in psutil.Process(pid).net_connections(kind="tcp")
)
return next((pid for pid in sorted(workers) if holds(pid)), None)
def _client_port(response: httpx.Response) -> int:
stream: Final[object] = response.extensions["network_stream"] # pyright: ignore[reportAny] # httpx types extensions as Any
assert isinstance(stream, httpcore.NetworkStream), stream
return _CLIENT_ADDRESS.validate_python(stream.get_extra_info("client_addr"))[1]
def _probe(gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]) -> tuple[int | None, JsonValue]:
with httpx.Client(base_url=str(gateway.client.base_url), timeout=30, trust_env=False) as client:
response: Final = client.post(
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {gateway.key}"}
)
return _holder(workers, _client_port(response)), JSON_OBJECT.validate_json(response.content).get("input_tokens")
def _round(
gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]
) -> tuple[tuple[int | None, JsonValue], ...]:
def probe(_index: int) -> tuple[int | None, JsonValue]:
return _probe(gateway, body, workers)
with ThreadPoolExecutor(max_workers=_PROBES_PER_ROUND) as pool:
return tuple(pool.map(probe, range(_PROBES_PER_ROUND)))
def _settled_on_every_worker(gateway: Gateway, body: Mapping[str, JsonValue]) -> None:
proxy_port: Final = gateway.client.base_url.port
assert proxy_port is not None, gateway.client.base_url
workers: Final = _proxy_workers(proxy_port)
eventually(
lambda: _round(gateway, body, workers),
lambda observed: (
frozenset(pid for pid, _ in observed) == workers and all(count == _COUNT for _, count in observed)
),
seconds=60,
)
def _deploy(gateway: Gateway, scenario: Scenario, provider: _Provider) -> _Deployment:
port: Final = _reserved_port()
model: Final = scenario.model(
model_info=None, model=f"{provider.prefix}/{_MODEL}", api_base=f"http://127.0.0.1:{port}", api_key=_API_KEY
)
with wire_server(_peer(provider), port=port):
_settled_on_every_worker(gateway, {"model": model, "messages": [USER]})
return _Deployment(provider, port, model)
@pytest.fixture(scope="module")
def deployments() -> Iterator[Mapping[str, _Deployment]]:
with gateway_from_environment() as gateway, gateway.scenario() as scenario:
yield MappingProxyType({chosen.prefix: _deploy(gateway, scenario, chosen) for chosen in (_ANTHROPIC, _AZURE)})
@pytest.fixture
def deployment(provider: _Provider, deployments: Mapping[str, _Deployment]) -> _Deployment:
return deployments[provider.prefix]
def _count(gateway: Gateway, body: Mapping[str, JsonValue], key: str | None = None) -> httpx.Response:
return gateway.request("POST", "/v1/messages/count_tokens", body, key=key)
def _payload(response: httpx.Response) -> dict[str, JsonValue]:
assert response.status_code == 200, response.text
return JSON_OBJECT.validate_json(response.content)
def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int:
response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"})
payload: Final = _payload(response)
total: Final = payload["total_tokens"]
assert payload["tokenizer_type"] not in _PROVIDER_TOKENIZERS, response.text
assert isinstance(total, int) and total > 0 and total != _COUNT, response.text
return total
def _bodies(wire: Wire, provider: _Provider) -> tuple[dict[str, JsonValue], ...]:
received: Final = wire.drain()
for request in received:
assert (request.method, request.target) == ("POST", provider.target), request.target
assert request.headers["anthropic-version"] == ANTHROPIC_VERSION, request.headers
assert TOKEN_COUNTING_BETA in request.headers["anthropic-beta"], request.headers
assert request.headers["content-type"] == "application/json", request.headers
assert request.headers["x-api-key"] == _API_KEY, request.headers
assert (request.headers.get("api-key") == _API_KEY) is provider.azure, request.headers
return tuple(JSON_OBJECT.validate_json(request.body) for request in received)
def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]:
return tuple(
stack.enter_context(httpx.Client(base_url=base_url, timeout=30, trust_env=False)) for _ in range(count)
)
def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]:
response: Final = client.post(
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"}
)
return response.status_code, JSON_OBJECT.validate_json(response.content).get("input_tokens")
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
def test_messages_count_tokens_lifts_the_leading_system_run(
deployment: _Deployment, gateway: Gateway, case: LiftCase
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, count_request(deployment.model, case))
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (expected_count_body(_MODEL, case),)
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_message(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_message(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
counted: Final = asyncio.run(
async_anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_the_provider(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/utils/token_counter",
count_request(deployment.model, STRING_CASE),
params={"call_endpoint": "true"},
)
payload: Final = _payload(response)
expected_tokenizer: Final = deployment.provider.tokenizer
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_COUNT, expected_tokenizer), response.text
assert payload["original_response"] == {"input_tokens": _COUNT}, response.text
assert (payload["request_model"], payload["model_used"]) == (deployment.model, _MODEL), response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_utils_token_counter_local_mode_never_calls_the_peer_for_a_leading_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
assert _local_count(gateway, count_request(deployment.model, STRING_CASE)) > 0
assert peer.drain() == ()
def test_responses_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": deployment.model, "input": USER_TEXT, "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": deployment.model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (
{"model": _MODEL, "messages": [USER], "system": [LIFTED, {"type": "text", "text": REMINDER}]},
)
def test_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = openai_client(gateway).responses.input_tokens.count(
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_async_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = asyncio.run(
async_openai_client(gateway).responses.input_tokens.count(
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
)
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "tools": TOOLS})
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == ({**_STRING_BODY, "tools": TOOLS},)
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {
"model": deployment.model,
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
}
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (
{"model": _MODEL, "messages": [USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], "system": [LIFTED]},
)
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [LEADING]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [], "system": [LIFTED]},)
def test_messages_count_tokens_leaves_the_request_untouched_for_a_non_text_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
local: Final = _local_count(gateway, count_request(deployment.model, STRING_CASE))
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "system": 5})
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [LEADING, USER], "system": 5},)
def test_messages_count_tokens_repeated_request_lifts_each_time(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
answers: Final = tuple(
_payload(_count(gateway, count_request(deployment.model, STRING_CASE))) for _ in range(2)
)
assert answers == ({"input_tokens": _COUNT},) * 2
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * 2
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_peer_call(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [{"role": "system"}, USER]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert peer.drain() == ()
follow_up: Final = _count(gateway, count_request(deployment.model, STRING_CASE))
assert _payload(follow_up) == {"input_tokens": _COUNT}, follow_up.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
first: Final = json.dumps([USER])
last: Final = json.dumps([LEADING, USER])
response: Final = gateway.client.post(
"/v1/messages/count_tokens",
content=f'{{"model": "{deployment.model}", "messages": {first}, "messages": {last}}}',
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
)
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])
def test_messages_count_tokens_falls_back_locally_when_the_peer_rejects_the_lifted_body(
deployment: _Deployment, gateway: Gateway, status: int
) -> None:
with _listening(deployment, _rejecting(status)) as peer:
body: Final = count_request(deployment.model, STRING_CASE)
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_unauthenticated_request_never_reaches_the_peer(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, count_request(deployment.model, STRING_CASE), key="sk-not-a-key")
assert response.status_code == 401, response.text
assert peer.drain() == ()
def test_peer_outage_between_concurrent_waves_falls_back_then_recovers(
deployment: _Deployment, gateway: Gateway
) -> None:
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
body: Final = count_request(deployment.model, STRING_CASE)
local: Final = _local_count(gateway, body)
def count(client: httpx.Client) -> tuple[int, JsonValue]:
return _counted_on(client, gateway.key, body)
with _listening(deployment) as peer:
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
with _listening(deployment) as revived:
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
assert _bodies(revived, deployment.provider) == (_STRING_BODY,) * len(clients)
def test_slow_peer_holds_concurrent_lifted_counts_without_stalling_the_proxy(
deployment: _Deployment, gateway: Gateway
) -> None:
held: Final[SimpleQueue[str]] = SimpleQueue()
release: Final = threading.Event()
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 6)
peer: Final = stack.enter_context(_listening(deployment, _holding(held, release, 20)))
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
stack.callback(release.set)
body: Final = count_request(deployment.model, STRING_CASE)
futures: Final = tuple(pool.submit(_counted_on, client, gateway.key, body) for client in clients)
eventually(held.qsize, lambda size: size == len(clients), seconds=30)
assert gateway.request("GET", "/health/liveliness").status_code == 200
assert _local_count(gateway, body) > 0
assert not any(future.done() for future in futures)
release.set()
assert tuple(future.result(timeout=30) for future in futures) == ((200, _COUNT),) * len(clients)
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
def test_chat_completions_on_the_same_deployment_keeps_a_mid_conversation_system_in_place(
deployments: Mapping[str, _Deployment], gateway: Gateway
) -> None:
identity: Final = f"msg_{uuid.uuid4().hex}"
anthropic_deployment: Final = deployments[_ANTHROPIC.prefix]
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/v1/messages"), request.target
return _json_reply(
200,
{
"id": identity,
"type": "message",
"role": "assistant",
"model": _MODEL,
"content": [{"type": "text", "text": "done"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 3},
},
)
with wire_server(respond, port=anthropic_deployment.port) as peer:
reminder: Final[dict[str, JsonValue]] = {
"role": "system",
"content": [
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
{"type": "text", "text": ""},
IMAGE_PART,
],
}
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": anthropic_deployment.model,
"max_tokens": 16,
"messages": [LEADING, USER, reminder, ASSISTANT, FOLLOW_UP],
},
)
assert response.status_code == 200, response.text
(sent,) = peer.drain()
body: Final = JSON_OBJECT.validate_json(sent.body)
assert body["system"] == [LIFTED], body
assert body["messages"] == [
{"role": "user", "content": [{"type": "text", "text": USER_TEXT}]},
{"role": "system", "content": [{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL}]},
{"role": "assistant", "content": [{"type": "text", "text": "One."}]},
{"role": "user", "content": [{"type": "text", "text": "Again"}]},
], body

View file

@ -11,7 +11,7 @@ from contextlib import ExitStack
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
from typing import Final, cast
import anthropic
import httpx
@ -24,6 +24,25 @@ from integration._support.bedrock_runtime_peer import respond as runtime_generat
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._count_tokens_system_lift import (
ASSISTANT,
FOLLOW_UP,
INSTRUCTION,
LEADING,
LIFT_CASES,
LIFTED,
MID_SYSTEM,
REMINDER,
STRING_CASE,
USER,
USER_TEXT,
LiftCase,
accepts_count_body,
async_openai_client,
count_request,
expected_count_body,
openai_client,
)
from pydantic import JsonValue, TypeAdapter
pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
@ -81,6 +100,10 @@ def _mantle_body(**fields: JsonValue) -> dict[str, JsonValue]:
_MANTLE_BARE: Final = _mantle_body()
_MANTLE_LIFTED: Final = _mantle_body(system=[LIFTED])
_LIFT_SDK_MESSAGES: Final = cast(
list[MessageParam], [LEADING, USER]
) # cast-ok: the SDK types reject the role the proxy lifts
_MANTLE_FULL: Final = _mantle_body(system=_SYSTEM, tools=_TOOLS)
@ -138,25 +161,8 @@ def _rejecting(status: int) -> Callable[[Request], Reply]:
return count
def _anthropic_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
def _anthropic_tool(tool: JsonValue) -> bool:
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
def _strict(request: Request) -> Reply:
body: Final = _JSON_OBJECT.validate_json(request.body)
messages: Final = body.get("messages")
tools: Final = body.get("tools", [])
accepted: Final = (
isinstance(messages, list)
and all(map(_anthropic_message, messages))
and isinstance(body.get("system", ""), (str, list))
and isinstance(tools, list)
and all(map(_anthropic_tool, tools))
)
accepted: Final = accepts_count_body(_JSON_OBJECT.validate_json(request.body))
return _mantle_counted(request) if accepted else _rejected(400)
@ -577,7 +583,7 @@ def test_bedrock_passthrough_count_tokens_still_answers_the_runtime_rejection(
assert mantle.drain() == ()
def test_responses_input_tokens_with_instructions_still_counts_locally(
def test_responses_input_tokens_with_instructions_counts_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
@ -592,14 +598,322 @@ def test_responses_input_tokens_with_instructions_still_counts_locally(
"/v1/responses/input_tokens",
{"model": model, "input": "Count this message", "instructions": "Be terse"},
)
payload: Final = _payload(response)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
(sent,) = _mantle_bodies(mantle)
messages: Final = sent["messages"]
assert isinstance(messages, list) and messages[0] == {"role": "system", "content": "Be terse"}, sent
assert "system" not in sent, sent
local: Final = _local_count(gateway, {"model": model, "messages": messages})
assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text
assert _mantle_bodies(mantle) == (_mantle_body(system=[{"type": "text", "text": "Be terse"}]),)
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
def test_messages_count_tokens_lifts_the_leading_system_run_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int, case: LiftCase
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = _count(gateway, count_request(model, case))
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (expected_count_body(_OPUS_BASE, case),)
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = _anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = asyncio.run(
_async_anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = gateway.request(
"POST", "/utils/token_counter", count_request(model, STRING_CASE), params={"call_endpoint": "true"}
)
payload: Final = _payload(response)
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_MANTLE_COUNT, "bedrock_mantle_api")
assert payload["original_response"] == {"input_tokens": _MANTLE_COUNT}, response.text
assert (payload["request_model"], payload["model_used"]) == (model, _OPUS), response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_openai_sdk_input_tokens_lifts_instructions_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = openai_client(gateway).responses.input_tokens.count(
model=model, input=USER_TEXT, instructions=INSTRUCTION
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_async_openai_sdk_input_tokens_lifts_instructions_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = asyncio.run(
async_openai_client(gateway).responses.input_tokens.count(
model=model, input=USER_TEXT, instructions=INSTRUCTION
)
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED, {"type": "text", "text": REMINDER}]),)
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system_to_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "tools": _TOOLS})
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED], tools=_TOOLS),)
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {
"model": model,
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
}
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (
_mantle_body(messages=[USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], system=[LIFTED]),
)
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [LEADING]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(messages=[], system=[LIFTED]),)
def test_messages_count_tokens_leaves_a_leading_system_in_place_beside_a_non_text_system_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
local: Final = _local_count(gateway, count_request(model, STRING_CASE))
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "system": 5})
assert _payload(response) == {"input_tokens": local}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(messages=[LEADING, USER], system=5),)
def test_messages_count_tokens_repeated_request_lifts_the_leading_system_each_time_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
answers: Final = tuple(_payload(_count(gateway, count_request(model, STRING_CASE))) for _ in range(2))
assert answers == ({"input_tokens": _MANTLE_COUNT},) * 2
assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) * 2
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * 2
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_mantle_call(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "system"}, USER]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert mantle.drain() == ()
follow_up: Final = _count(gateway, count_request(model, STRING_CASE))
assert _payload(follow_up) == {"input_tokens": _MANTLE_COUNT}, follow_up.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
first: Final = json.dumps([USER])
last: Final = json.dumps([LEADING, USER])
response: Final = gateway.client.post(
"/v1/messages/count_tokens",
content=f'{{"model": "{model}", "messages": {first}, "messages": {last}}}',
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
)
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_disabled_token_counter_counts_a_leading_system_through_mantle(
gateway: Gateway, mantle_port: int, tmp_path: Path
) -> None:
with ExitStack() as stack:
runtime: Final = stack.enter_context(wire_server(_runtime))
config: Final = _owned_config(
tmp_path / "disabled-token-counter-lift.yaml", runtime.url, {"disable_token_counter": True}
)
owned: Final = stack.enter_context(
owned_proxy_process(
gateway,
tmp_path,
_mantle_environment(mantle_port),
config=config,
workers=2,
remove_environment=_INHERITED_BEARER,
)
)
with wire_server(_mantle(_strict), port=mantle_port) as strict:
counted: Final = _count(owned.gateway, count_request(_OWNED_OPUS, STRING_CASE))
assert _payload(counted) == {"input_tokens": _MANTLE_COUNT}, counted.text
assert _mantle_bodies(strict) == (_MANTLE_LIFTED,)
assert len(_runtime_count_targets(runtime)) == 1
def test_mantle_outage_between_concurrent_lifted_waves_falls_back_then_recovers(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
runtime: Final = stack.enter_context(wire_server(_runtime))
scenario: Final = stack.enter_context(gateway.scenario())
model: Final = _deployment(scenario, runtime.url)
body: Final = count_request(model, STRING_CASE)
local: Final = _local_count(gateway, body)
def count(client: httpx.Client) -> tuple[int, JsonValue]:
return _counted_on(client, gateway.key, body)
with wire_server(_mantle(_strict), port=mantle_port) as mantle:
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * len(clients)
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
with wire_server(_mantle(_strict), port=mantle_port) as revived:
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
assert _mantle_bodies(revived) == (_MANTLE_LIFTED,) * len(clients)
assert len(_runtime_count_targets(runtime)) == 3 * len(clients)
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])

View file

@ -252,6 +252,23 @@ def test_each_provider_gets_its_own_path_key_and_body_and_is_billed_from_the_cos
assert math.isclose(_number(row["spend"]), expected_spend, rel_tol=1e-9), row
def test_test_connection_evaluation_mode_uses_typesafe_decisions_path(gateway: Gateway) -> None:
provider: Final = _PROVIDERS[1]
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(provider))
model: Final = _deployment(scenario, handle, provider)
response: Final = gateway.request(
"POST",
"/health/test_connection",
{"litellm_params": {"model": model}, "mode": "evaluation"},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success", response.text
(call,) = _upstream_calls(gateway, handle)
assert call["path"] == f"/{handle.scenario_id}{provider.path}"
def test_repeated_identical_requests_each_reach_the_upstream_and_are_each_billed(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))

View file

@ -0,0 +1,292 @@
import os
import socket
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, eventually, gateway_from_environment, object_value
from integration._support.database import read_rows, scratch_database
from integration._support.pdf_document import (
COUNT_REFUSED,
COUNT_TOKENS_TARGET,
LETTER,
chat_body,
chat_file,
messages_body,
pdf_data_url,
pdf_document,
responses_body,
responses_input_file,
)
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._cache_control_marks_support import anthropic_peer, marker_of, owned_config
from pydantic import JsonValue, TypeAdapter
from redis import Redis
_WINDOW: Final = "window-claude"
_ITPM: Final = "itpm-claude"
_BUDGET: Final = "budget-claude"
_AFFINITY: Final = "affinity-claude"
_LIMIT: Final = 5000
_PAGE_COUNT: Final = 12
_PAGES: Final = (LETTER,) * _PAGE_COUNT
_ASK: Final = "Summarize the attached report in one sentence."
_FOLLOW_UPS: Final = 9
_TRANSCRIPT: Final = " ".join(f"Line {number}: revenue, margins and headcount moved." for number in range(160))
_PIN_KEYS: Final = "*:prompt_caching"
_PROVIDER_KEY: Final = "synthetic-provider-key"
_SURFACES: Final = ("messages", "messages-stream", "chat-file", "responses-input-file")
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
@dataclass(frozen=True, slots=True)
class _Rig:
owned: OwnedProxy
wire: Wire
held_port: int
def _marker() -> str:
return uuid.uuid4().hex
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return int(reserve.getsockname()[1])
def _deployment(
name: str,
backend: str,
api_base: str,
*,
model_info: Mapping[str, JsonValue] | None = None,
**params: JsonValue,
) -> dict[str, JsonValue]:
return {
"model_name": name,
"litellm_params": {"model": f"anthropic/{backend}", "api_base": api_base, "api_key": _PROVIDER_KEY, **params},
"model_info": dict(model_info or {}),
}
def _peer(request: Request) -> Reply:
if request.target == COUNT_TOKENS_TARGET:
return COUNT_REFUSED
return anthropic_peer(request)
def _held(release: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert release.wait(timeout=120), "Held peer was never released"
return _peer(request)
return respond
def _request(surface: str, model: str, pages: int, marker: str) -> tuple[str, dict[str, JsonValue]]:
text: Final = f"{_ASK} marker-{marker}"
letter_pages: Final = (LETTER,) * pages
if surface == "messages":
return "/v1/messages", messages_body(model, [pdf_document(letter_pages)], text)
if surface == "messages-stream":
return "/v1/messages", messages_body(model, [pdf_document(letter_pages)], text, stream=True)
if surface == "chat-file":
return "/v1/chat/completions", chat_body(model, [chat_file(pdf_data_url(letter_pages))], text)
return "/v1/responses", responses_body(model, [responses_input_file(pdf_data_url(letter_pages))], text)
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
directory: Final = tmp_path_factory.mktemp("pdf-pre-call")
held_port: Final = _free_port()
with gateway_from_environment() as gateway, wire_server(_peer) as wire:
config: Final = owned_config(
directory,
[
_deployment(_WINDOW, "claude-haiku-5-5", wire.url, model_info={"max_input_tokens": _LIMIT}),
_deployment(_ITPM, "claude-sonnet-4-6", wire.url, itpm=_LIMIT),
_deployment(_BUDGET, "claude-opus-4-8", f"http://127.0.0.1:{held_port}"),
],
litellm_settings={"cache": False},
router_settings={
"enable_pre_call_checks": True,
"optional_pre_call_checks": ["enforce_model_rate_limits"],
},
)
with owned_proxy_process(gateway, directory, {}, config=config, workers=2) as owned:
yield _Rig(owned, wire, held_port)
def _spend_row(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
return rows[0]
@pytest.mark.parametrize("surface", _SURFACES)
def test_context_window_check_rejects_a_pdf_whose_pages_exceed_max_input_tokens(rig: _Rig, surface: str) -> None:
rig.wire.drain()
path, body = _request(surface, _WINDOW, _PAGE_COUNT, _marker())
response: Final = rig.owned.gateway.request("POST", path, body)
assert response.status_code == 400, (response.status_code, response.text)
assert "Context Window exceeded" in response.text, response.text
assert [request.target for request in rig.wire.drain()] == []
def test_context_window_check_passes_a_one_page_pdf_and_forwards_its_bytes(rig: _Rig) -> None:
rig.wire.drain()
path, body = _request("messages", _WINDOW, 1, _marker())
response: Final = rig.owned.gateway.request("POST", path, body)
assert response.status_code == 200, response.text
received: Final = rig.wire.drain()
assert [request.target for request in received] == ["/v1/messages"]
forwarded: Final = _JSON_OBJECT.validate_json(received[0].body)["messages"]
assert isinstance(forwarded, list) and len(forwarded) == 1, forwarded
blocks: Final = object_value(forwarded[0])["content"]
assert isinstance(blocks, list) and pdf_document((LETTER,)) in blocks, blocks
def test_model_itpm_check_rejects_a_pdf_whose_pages_exceed_the_limit(rig: _Rig) -> None:
rig.wire.drain()
path, body = _request("messages", _ITPM, _PAGE_COUNT, _marker())
response: Final = rig.owned.gateway.request("POST", path, body)
assert response.status_code == 429, (response.status_code, response.text)
assert '"error"' in response.text, response.text
assert [request.target for request in rig.wire.drain()] == []
def test_key_budget_reservation_rejects_the_second_concurrent_pdf_request(rig: _Rig) -> None:
release: Final = threading.Event()
requests: Final = tuple(_request("messages", _BUDGET, _PAGE_COUNT, _marker()) for _ in range(2))
with wire_server(_held(release), port=rig.held_port) as wire, rig.owned.gateway.scenario() as scenario:
key: Final = scenario.key(max_budget=0.01)
headers: Final = {"Authorization": f"Bearer {key}"}
with (
httpx.Client(base_url=str(rig.owned.gateway.client.base_url), timeout=90, trust_env=False) as client,
ThreadPoolExecutor(max_workers=2) as pool,
):
futures: Final = tuple(
pool.submit(client.post, path, json=body, headers=headers) for path, body in requests
)
try:
eventually(
lambda: wire.received.qsize() + sum(1 for future in futures if future.done()),
lambda settled: settled >= 2,
seconds=60,
)
finally:
release.set()
responses: Final = tuple(future.result(timeout=90) for future in futures)
received: Final = wire.drain()
assert sorted(response.status_code for response in responses) == [200, 422], [
response.text for response in responses
]
rejected: Final = next(response for response in responses if response.status_code != 200)
assert "udget" in rejected.text, rejected.text
assert [request.target for request in received] == ["/v1/messages"]
def test_a_pre_call_rejection_logs_one_failure_row_and_calls_no_upstream(rig: _Rig) -> None:
path, body = _request("messages", _WINDOW, _PAGE_COUNT, _marker())
response: Final = rig.owned.gateway.request("POST", path, body)
assert response.status_code == 400, (response.status_code, response.text)
assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure"
assert rig.wire.drain() == ()
def _first_turn(session: str, marker: str) -> list[JsonValue]:
return [
{"role": "system", "content": f"Answer from the attached report. session {session}"},
*chat_body(
_AFFINITY,
[{"type": "text", "text": _TRANSCRIPT}, chat_file(pdf_data_url(_PAGES))],
f"{_ASK} marker-{marker}",
)["messages"],
]
def _follow_up(first_turn: list[JsonValue], marker: str) -> list[JsonValue]:
return [
*first_turn,
{"role": "assistant", "content": "One sentence."},
{"role": "user", "content": f"And the margins? marker-{marker}"},
]
def _served(wire: Wire) -> frozenset[str]:
return frozenset(marker_of(request) for request in wire.drain())
def _pin_count(cache: Redis) -> int:
return len(cache.keys(_PIN_KEYS))
@pytest.mark.timeout(2 * graceful_stop_seconds() + 180)
def test_prompt_caching_affinity_pins_follow_ups_of_a_cached_turn_that_carries_a_pdf(
gateway: Gateway, tmp_path: Path
) -> None:
session: Final = _marker()
first_marker: Final = _marker()
follow_up_markers: Final = tuple(_marker() for _ in range(_FOLLOW_UPS))
first_turn: Final = _first_turn(session, first_marker)
with scratch_database() as database_url, wire_server(_peer) as left, wire_server(_peer) as right:
config: Final = owned_config(
tmp_path,
[
_deployment(_AFFINITY, "claude-opus-5-5", left.url, model_info={"id": f"pdf-left-{session}"}),
_deployment(_AFFINITY, "claude-opus-5-5", right.url, model_info={"id": f"pdf-right-{session}"}),
],
litellm_settings={"cache": False},
router_settings={
"optional_pre_call_checks": ["prompt_caching"],
"redis_host": os.environ["REDIS_HOST"],
"redis_port": int(os.environ["REDIS_PORT"]),
},
)
with (
Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache,
owned_proxy_process(
gateway,
tmp_path,
{"DATABASE_URL": database_url},
config=config,
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=2,
) as owned,
owned.gateway.scenario() as scenario,
):
key: Final = scenario.key(metadata={"enable_prompt_caching": True})
pins_before: Final = _pin_count(cache)
first: Final = owned.gateway.request(
"POST", "/v1/chat/completions", {"model": _AFFINITY, "messages": first_turn, "max_tokens": 64}, key=key
)
assert first.status_code == 200, first.text
eventually(lambda: _pin_count(cache), lambda count: count > pins_before, seconds=70)
follow_ups: Final = tuple(
owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": _AFFINITY, "messages": _follow_up(first_turn, marker), "max_tokens": 64},
key=key,
)
for marker in follow_up_markers
)
served: Final = {"left": _served(left), "right": _served(right)}
assert [response.status_code for response in follow_ups] == [200] * _FOLLOW_UPS, [
response.text for response in follow_ups
]
assert served["left"] | served["right"] == {first_marker, *follow_up_markers}, served
first_side: Final = next(side for side in served if first_marker in served[side])
assert served[first_side] == {first_marker, *follow_up_markers}, served

View file

@ -15,7 +15,7 @@ from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.process import owned_proxy
from integration._support.redis_process import owned_redis
from integration._support.wire import Reply, Request, wire_server
@ -24,6 +24,7 @@ from redis import Redis
from redis.exceptions import TimeoutError as RedisTimeoutError
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
MODEL_INFO_ENTRIES: Final = TypeAdapter(tuple[dict[str, JsonValue], ...])
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
OPENAI_MODEL: Final = "gpt-4o-mini"
MASTER_KEY: Final = "sk-integration-usage-routing-redis-reads"
@ -161,6 +162,12 @@ def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
def _registered_deployment_ids(candidate: Gateway, model_name: str) -> frozenset[str]:
entries: Final = MODEL_INFO_ENTRIES.validate_python(candidate.get("/model/info")["data"])
group: Final = tuple(entry for entry in entries if entry["model_name"] == model_name)
return frozenset(string_value(object_value(entry["model_info"])["id"]) for entry in group)
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
@ -202,8 +209,8 @@ def test_proxy_usage_routing_reads_cooldown_tpm_then_rpm_from_redis(
config=config_path,
) as candidate:
eventually(
lambda: wire.received.qsize(),
lambda received: received >= len(deployment_ids),
lambda: _registered_deployment_ids(candidate, model_name),
lambda registered: registered == frozenset(deployment_ids),
seconds=15,
)
wire.drain()

View file

@ -23,7 +23,9 @@ GROUPS: Final = MappingProxyType(
"security": ("security",),
}
)
GITHUB_FILES: Final = frozenset({"tests/integration/database/test_roi_observed.py"})
GITHUB_FILES: Final = frozenset(
{"tests/integration/database/test_roi_observed.py", "tests/integration/mcp/test_interactions.py"}
)
@dataclass(frozen=True, slots=True)

View file

@ -25,7 +25,7 @@ from litellm.rust_bridge.trace.generated.types import Trace, TraceScope
from litellm.rust_bridge.trace.storage import ClickHouseStorage, TraceStorageConfig, span_rows
from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError
from litellm.tracing.types import SpendLogRecord
from scripts.seed_tracing_fixtures import (
from seed_tracing_fixtures import (
TRACE,
TRACE_FIXTURES,
Copies,
@ -499,7 +499,7 @@ class SeededTraceAPI:
@pytest.fixture
def seeded_trace_api(clickhouse_url: str) -> Iterator[SeededTraceAPI]:
from scripts.seed_tracing_fixtures import (
from seed_tracing_fixtures import (
TRACE_FIXTURES,
fixture_replays,
rebase_spend,

View file

@ -61,6 +61,199 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer
_JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage)
@pytest.mark.asyncio
async def test_tool_continuation_reaches_modern_upstream() -> None:
from mcp.types import ElicitResult, TextContent
def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
if payload.method == "server/discover":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"supportedVersions": ["2026-07-28"],
"capabilities": {"tools": {}},
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
},
},
)
if payload.method == "tools/list":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"tools": [{"name": "confirm", "inputSchema": {"type": "object"}}],
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
},
},
)
assert payload.method == "tools/call"
assert payload.params is not None
assert payload.params.get("requestState") == "opaque-upstream-state"
assert payload.params.get("inputResponses") == {"confirmation": {"action": "accept"}}
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"resultType": "complete",
"content": [{"type": "text", "text": "confirmed"}],
"isError": False,
},
},
)
client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", protocol_version="2026-07-28")
result: Final = await client.call_tool(
CallToolRequestParams(
name="confirm",
request_state="opaque-upstream-state",
input_responses={"confirmation": ElicitResult(action="accept")},
),
raise_on_error=True,
)
assert isinstance(result, CallToolResult)
assert result.content == [TextContent(type="text", text="confirmed")]
assert result.is_error is False
@pytest.mark.asyncio
@pytest.mark.parametrize("modern_caller", [False, True])
@pytest.mark.parametrize(
"elicitation_mode,sampling", [("form", False), ("form", True), ("url", False), ("url", True), ("none", True)]
)
async def test_modern_input_request_uses_existing_elicitation_callback(
modern_caller: bool, sampling: bool, elicitation_mode: str
) -> None:
from queue import SimpleQueue
from mcp.types import (
ElicitResult,
ElicitRequestParams,
TextContent,
CreateMessageRequestParams,
CreateMessageResult,
)
from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult
observed: Final[SimpleQueue[str]] = SimpleQueue()
sampled: Final = CreateMessageResult(
role="assistant", content=TextContent(type="text", text="sampled"), model="test"
)
expected_responses: Final = {
**({"consent": {"action": "accept"}} if elicitation_mode != "none" else {}),
**({"sample": sampled.model_dump(by_alias=True, exclude_none=True)} if sampling else {}),
}
async def elicit(context: object, params: ElicitRequestParams) -> ElicitResult:
if params.mode == "url":
assert params.elicitation_id, "Legacy URL input must carry its required elicitation ID"
observed.put(params.message)
return ElicitResult(action="accept")
async def sample(context: object, params: CreateMessageRequestParams) -> CreateMessageResult:
assert params.max_tokens == 10
observed.put("sample")
return sampled
def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
if payload.method == "server/discover":
result = {
"supportedVersions": ["2026-07-28"],
"capabilities": {"tools": {}},
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
}
elif payload.method == "tools/list":
result = {
"tools": [{"name": "confirm", "inputSchema": {"type": "object"}}],
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
}
elif not (payload.params or {}).get("requestState"):
result = {
"resultType": "input_required",
"requestState": "pending",
"inputRequests": {
**(
{
"consent": {
"method": "elicitation/create",
"params": {
"message": "Confirm operation",
**(
{"mode": "form", "requestedSchema": {"type": "object", "properties": {}}}
if elicitation_mode == "form"
else {"mode": "url", "url": "https://example.com/confirm"}
),
},
}
}
if elicitation_mode != "none"
else {}
),
**(
{"sample": {"method": "sampling/createMessage", "params": {"messages": [], "maxTokens": 10}}}
if sampling
else {}
),
},
}
else:
assert (payload.params or {}).get("requestState") == "pending"
assert (payload.params or {}).get("inputResponses") == expected_responses
result = {"resultType": "complete", "content": [{"type": "text", "text": "confirmed"}], "isError": False}
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result})
client: Final = _MockTransportClient(
respond,
server_url="https://example.com/mcp",
protocol_version="2026-07-28",
elicitation_callback=elicit,
sampling_callback=sample,
)
result: Final = await client.call_tool(
CallToolRequestParams(name="confirm", arguments={}), raise_on_error=True, allow_input_required=modern_caller
)
if modern_caller and elicitation_mode != "none":
assert isinstance(result, BoundInputRequiredResult)
assert tuple(result.input_requests or {}) == ("consent",)
assert result.gateway_responses == ({"sample": sampled} if sampling else None)
resumed: Final = await client.call_tool(
CallToolRequestParams(
name="confirm",
arguments={},
request_state=result.request_state,
input_responses={"consent": ElicitResult(action="accept"), **(result.gateway_responses or {})},
),
raise_on_error=True,
allow_input_required=True,
)
assert isinstance(resumed, CallToolResult)
assert resumed.content == [TextContent(type="text", text="confirmed")]
else:
assert isinstance(result, CallToolResult)
assert result.content == [TextContent(type="text", text="confirmed")]
assert sorted(observed.get_nowait() for _ in range(observed.qsize())) == sorted(
([] if modern_caller or elicitation_mode == "none" else ["Confirm operation"])
+ (["sample"] if sampling else [])
)
def _initialized(instructions: str | None = None) -> InitializeResult:
return InitializeResult(
protocol_version=LATEST_HANDSHAKE_VERSION,
@ -3615,3 +3808,79 @@ async def test_optional_discovery_retains_freshness_across_pages(
assert result.next_cursor is None
assert result.ttl_ms == max(0, ttl - cleanup_seconds * 1000)
assert len(await getattr(client, "list_" + kind)(raise_on_error=True)) == 2
def test_prompt_continuation_polling_respects_the_original_deadline() -> None:
from mcp.types import GetPromptRequestParams
def respond(request: httpx2.Request) -> httpx2.Response:
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
assert isinstance(payload, JSONRPCRequest)
result: Final = (
{
"supportedVersions": ["2026-07-28"],
"capabilities": {"prompts": {}},
"resultType": "complete",
"cacheScope": "private",
"ttlMs": 0,
}
if payload.method == "server/discover"
else {"resultType": "input_required", "requestState": "pending"}
)
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result})
loop: Final = _AutojumpClockLoop()
client: Final = _MockTransportClient(
respond, server_url="https://example.com/mcp", protocol_version="2026-07-28", timeout=0.12
)
try:
with pytest.raises(TimeoutError):
loop.run_until_complete(client.get_prompt(GetPromptRequestParams(name="pending")))
assert loop.time() == pytest.approx(0.12)
finally:
loop.close()
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ("form", "url"))
@pytest.mark.parametrize("enabled", (False, True))
async def test_modern_elicitation_honors_server_permission(mode: str, enabled: bool) -> None:
import anyio
from mcp import ClientSession, MCPError
from mcp.shared.message import SessionMessage
from mcp.types import (
ElicitRequest,
ElicitRequestFormParams,
ElicitRequestParams,
ElicitRequestURLParams,
ElicitResult,
InputRequiredResult,
)
async def elicit(context: object, params: ElicitRequestParams) -> ElicitResult:
return ElicitResult(action="accept")
client: Final = MCPClient(server_url="https://example.com/mcp", elicitation_callback=elicit if enabled else None)
request: Final = AsyncMock(
return_value=InputRequiredResult(
request_state="pending",
input_requests={
"consent": ElicitRequest(
params=ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object"})
if mode == "form"
else ElicitRequestURLParams(message="Confirm", url="https://example.com/confirm")
)
},
)
)
send, receive = anyio.create_memory_object_stream[SessionMessage](1)
async with send, receive:
session: Final = ClientSession(receive, send)
if enabled:
result: Final = await client._request_with_interaction(session, request, None, None, True)
assert isinstance(result, InputRequiredResult)
assert result.input_requests["consent"].params.mode == mode
else:
with pytest.raises(MCPError, match="Elicitation is disabled"):
await client._request_with_interaction(session, request, None, None, True)
request.assert_awaited_once_with(None, None)

View file

@ -3,6 +3,7 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -14,6 +15,7 @@ from litellm.integrations.focus.destinations.mavvrik_destination import (
)
VALID_ENDPOINT = "https://api.mavvrik.ai/tenant123"
_FOCUS_NOW: Final = datetime(2026, 2, 1, 0, 0, tzinfo=timezone.utc)
def _make_window() -> FocusTimeWindow:
@ -504,7 +506,6 @@ async def test_export_window_passes_max_rows_as_limit(monkeypatch):
async def test_run_scheduled_export_catches_up_missed_dates():
"""If metricsMarker is 2 days behind, _run_scheduled_export exports missed dates first."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
@ -512,15 +513,15 @@ async def test_run_scheduled_export_catches_up_missed_dates():
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW)
# metricsMarker = 3 days ago → 2 missed dates (day-2 and day-1) + today's run
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
two_days_ago = now - timedelta(days=2)
three_days_ago = now - timedelta(days=3)
now: Final = _FOCUS_NOW
yesterday: Final = now - timedelta(days=1)
two_days_ago: Final = now - timedelta(days=2)
three_days_ago: Final = now - timedelta(days=3)
marker_ts = int(three_days_ago.timestamp())
marker_ts: Final = int(three_days_ago.timestamp())
# Mock destination
dest_mock = MagicMock(spec=FocusMavvrikDestination)
@ -551,7 +552,6 @@ async def test_run_scheduled_export_catches_up_missed_dates():
async def test_run_scheduled_export_no_catchup_when_marker_is_current():
"""If metricsMarker = yesterday, no catch-up needed — just export yesterday."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
@ -559,11 +559,11 @@ async def test_run_scheduled_export_no_catchup_when_marker_is_current():
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW)
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
marker_ts = int(yesterday.timestamp())
now: Final = _FOCUS_NOW
yesterday: Final = now - timedelta(days=1)
marker_ts: Final = int(yesterday.timestamp())
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
@ -592,9 +592,9 @@ async def test_run_scheduled_export_skips_catchup_when_marker_is_unparseable():
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW)
now: Final = _FOCUS_NOW
yesterday: Final = now - timedelta(days=1)
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value="not-a-date")
@ -706,7 +706,6 @@ def test_parse_metrics_marker_returns_none_for_garbage():
async def test_catchup_capped_at_max_catchup_days():
"""Catch-up must not go further back than _MAX_CATCHUP_DAYS."""
import polars as pl
from datetime import datetime, timedelta, timezone
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
@ -714,14 +713,14 @@ async def test_catchup_capped_at_max_catchup_days():
FocusMavvrikDestination,
)
logger = MavvrikFocusLogger()
max_days = MavvrikFocusLogger._MAX_CATCHUP_DAYS
logger: Final = MavvrikFocusLogger(clock=lambda: _FOCUS_NOW)
max_days: Final = MavvrikFocusLogger._MAX_CATCHUP_DAYS
now = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
yesterday = now - timedelta(days=1)
now: Final = _FOCUS_NOW
yesterday: Final = now - timedelta(days=1)
# Marker is 30 days ago — well beyond the cap
thirty_days_ago = now - timedelta(days=30)
marker_ts = int(thirty_days_ago.timestamp())
thirty_days_ago: Final = now - timedelta(days=30)
marker_ts: Final = int(thirty_days_ago.timestamp())
dest_mock = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
@ -740,11 +739,57 @@ async def test_catchup_capped_at_max_catchup_days():
assert db_mock.get_usage_data.call_count <= max_days
# First catch-up date must not be earlier than (yesterday - max_days + 1)
earliest_allowed = yesterday - timedelta(days=max_days - 1)
first_call_start = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"]
earliest_allowed: Final = yesterday - timedelta(days=max_days - 1)
first_call_start: Final = db_mock.get_usage_data.call_args_list[0].kwargs["start_time_utc"]
assert first_call_start.date() >= earliest_allowed.date()
@pytest.mark.asyncio
async def test_run_scheduled_export_uses_one_clock_read_across_midnight():
import polars as pl
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
from litellm.integrations.focus.destinations.mavvrik_destination import (
FocusMavvrikDestination,
)
logger: Final = MavvrikFocusLogger(
clock=iter(
(
datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc),
datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc),
)
).__next__
)
marker_ts: Final = int(datetime(2026, 1, 28, 0, 0, tzinfo=timezone.utc).timestamp())
dest_mock: Final = MagicMock(spec=FocusMavvrikDestination)
dest_mock.get_metrics_marker = AsyncMock(return_value=marker_ts)
db_mock: Final = MagicMock()
db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame())
engine_mock: Final = MagicMock()
engine_mock.database = db_mock
engine_mock.destination = dest_mock
logger._engine = engine_mock
await logger._run_scheduled_export()
windows: Final = tuple(
(call.kwargs["start_time_utc"], call.kwargs["end_time_utc"])
for call in db_mock.get_usage_data.call_args_list
)
assert windows == (
(
datetime(2026, 1, 29, 0, 0, tzinfo=timezone.utc),
datetime(2026, 1, 30, 0, 0, tzinfo=timezone.utc),
),
(
datetime(2026, 1, 30, 0, 0, tzinfo=timezone.utc),
datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc),
),
)
@pytest.mark.asyncio
async def test_register_resets_on_410():
"""_registered flag must be False after a 410 so next run re-registers."""

View file

@ -11,6 +11,7 @@ import litellm
from litellm.litellm_core_utils.prompt_templates.common_utils import encrypted_reasoning_signature
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
CONVERTED_SYSTEM_NOTE,
anthropic_system_blocks,
place_mid_conversation_system,
split_leading_system_run,
)
@ -390,3 +391,23 @@ def test_flagged_placement_keeps_a_system_before_an_assistant_turn_whose_empty_t
)
assert _roles(placed) == ["user", "system", "assistant", "user"]
def test_anthropic_system_blocks_keeps_text_parts_with_their_cache_control_and_drops_the_rest():
run = [
{"role": "system", "content": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
{
"role": "system",
"content": [
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": ""},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aW1hZ2U="}},
],
},
{"role": "system", "content": ""},
]
assert anthropic_system_blocks(run) == (
{"type": "text", "text": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
)

View file

@ -11,6 +11,7 @@ import threading
import time
from collections.abc import Mapping
from concurrent.futures import Future, wait
from itertools import accumulate, chain
from pathlib import Path
from typing import Final
from unittest.mock import MagicMock
@ -1444,7 +1445,7 @@ def _count_user_content(content: list[dict]) -> int:
ids=["base64", "url", "file"],
)
def test_anthropic_document_block_with_opaque_source_is_priced_like_an_image(source: dict[str, str]):
"""A `document` whose bytes can't be tokenized locally is priced like an `image`, not raised on."""
"""A `document` with no readable pages (a bare PDF header, a URL, a file id) is priced like an `image`, not raised on."""
prompt = {"type": "text", "text": "Summarize this file."}
assert _count_user_content([prompt, {"type": "document", "source": source}]) == _count_user_content(
@ -1504,6 +1505,12 @@ def test_openai_file_block_prices_like_the_equivalent_anthropic_document():
assert _count_user_content([prompt, inline_file]) == _count_user_content([prompt, document])
assert _count_user_content([prompt, inline_file]) > _count_user_content([prompt])
readable: Final = _pdf_base64(("Revenue grew eleven percent while churn fell to two percent.",))
readable_file: Final = {"type": "file", "file": {"filename": "report.pdf", "file_data": "data:application/pdf;base64," + readable}}
readable_document: Final = {"type": "document", "title": "report.pdf", "source": _pdf_source(readable)}
assert _count_user_content([prompt, readable_file]) == _count_user_content([prompt, readable_document])
assert _count_user_content([prompt, readable_file]) > _count_user_content([prompt, inline_file])
def test_openai_file_block_without_inline_bytes_counts_what_it_carries():
"""A `file` block naming an uploaded file has no bytes to price, so it adds only the filename's tokens."""
@ -1518,6 +1525,101 @@ def test_openai_file_block_without_inline_bytes_counts_what_it_carries():
)
def _pdf_base64(pages: tuple[str, ...], width: int = 612, height: int = 792) -> str:
def page_objects(index: int, text: str) -> tuple[bytes, bytes]:
escaped: Final = text.replace("\\", "\\\\").replace("(", "\\(").replace(")", "\\)")
stream: Final = f"BT /F1 12 Tf 72 720 Td ({escaped}) Tj ET".encode("latin-1")
content: Final = f"<< /Length {len(stream)} >>\nstream\n".encode() + stream + b"\nendstream"
page: Final = (
f"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 {width} {height}] "
f"/Resources << /Font << /F1 3 0 R >> >> /Contents {4 + 2 * index} 0 R >>"
).encode()
return content, page
kids: Final = " ".join(f"{5 + 2 * index} 0 R" for index in range(len(pages)))
bodies: Final = (
b"<< /Type /Catalog /Pages 2 0 R >>",
f"<< /Type /Pages /Count {len(pages)} /Kids [ {kids} ] >>".encode(),
b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>",
*chain.from_iterable(page_objects(index, text) for index, text in enumerate(pages)),
)
header: Final = b"%PDF-1.4\n"
objects: Final = tuple(
f"{number} 0 obj\n".encode() + body + b"\nendobj\n" for number, body in enumerate(bodies, start=1)
)
offsets: Final = accumulate((len(header), *(len(obj) for obj in objects[:-1])))
xref: Final = f"xref\n0 {len(objects) + 1}\n0000000000 65535 f \n".encode() + b"".join(
f"{offset:010d} 00000 n \n".encode() for offset in offsets
)
trailer: Final = (
f"trailer\n<< /Size {len(objects) + 1} /Root 1 0 R >>\n"
f"startxref\n{len(header) + sum(len(obj) for obj in objects)}\n%%EOF\n"
).encode()
return base64.b64encode(header + b"".join(objects) + xref + trailer).decode()
def _pdf_source(pdf_base64: str) -> dict[str, str]:
return {"type": "base64", "media_type": "application/pdf", "data": pdf_base64}
def test_base64_pdf_document_counts_every_page_text_and_rendering():
"""A base64 PDF is read page by page: each page costs its text plus the image Anthropic renders it to.
Before the fix the whole document was priced as one 85-token image, so a count_tokens call that fell back
to the local counter answered 116 for a 12-page PDF the provider then billed at 35941 input tokens.
"""
prompt: Final = {"type": "text", "text": "Summarize this file."}
first: Final = "Revenue grew eleven percent while churn fell to two percent."
second: Final = "Headcount is flat and the office lease was renewed for three years."
base: Final = _count_user_content([prompt])
blank_page: Final = _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64(("",)))}]) - base
assert blank_page > 0
assert _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first, second)))}]) == (
_count_user_content([prompt, {"type": "text", "text": first}, {"type": "text", "text": second}]) + 2 * blank_page
)
one_page: Final = _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first,)))}]) - base
twelve_pages: Final = (
_count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first,) * 12))}]) - base
)
assert one_page > _count_user_content([prompt, {"type": "text", "text": first}]) - base
assert twelve_pages == 12 * one_page
@pytest.mark.parametrize("fields", [{}, {"title": "Q3 board packet", "context": "Shared by finance"}], ids=["bare", "described"])
def test_pdf_page_rendering_cost_follows_anthropic_image_scaling(fields: dict[str, str]):
"""Anthropic rasterizes each PDF page within its image limits (1568 px long edge, 1.15 MP) and bills
width * height / 750 tokens for it: https://platform.claude.com/docs/en/build-with-claude/pdf-support and
https://platform.claude.com/docs/en/build-with-claude/vision, read 2026-10-07, when Bedrock billed about
1550 tokens per blank Letter page on both Sonnet 4.6 and Opus 4.8.
"""
prompt: Final = {"type": "text", "text": "Summarize this file."}
def blank_page_cost(width: int, height: int) -> int:
with_page: Final = {"type": "document", "source": _pdf_source(_pdf_base64(("",), width, height)), **fields}
without_page: Final = {"type": "document", "source": _pdf_source(_pdf_base64((), width, height)), **fields}
return _count_user_content([prompt, with_page]) - _count_user_content([prompt, without_page])
letter: Final = blank_page_cost(612, 792)
poster: Final = blank_page_cost(2448, 3168)
strip: Final = blank_page_cost(1000, 100)
assert letter == poster == 1534
assert strip == 328
def test_pdf_document_without_pypdf_is_priced_like_an_image(monkeypatch: pytest.MonkeyPatch):
prompt: Final = {"type": "text", "text": "Summarize this file."}
source: Final = _pdf_source(_pdf_base64(("Revenue grew eleven percent while churn fell to two percent.",)))
priced_by_page: Final = _count_user_content([prompt, {"type": "document", "source": source}])
monkeypatch.setitem(sys.modules, "pypdf", None)
priced_as_image: Final = _count_user_content([prompt, {"type": "document", "source": source}])
assert priced_as_image == _count_user_content([prompt, {"type": "image", "source": source}])
assert priced_as_image < priced_by_page
def _png_data_url(width: int, height: int) -> str:
ihdr = b"\x89PNG\r\n\x1a\n" + (13).to_bytes(4, "big") + b"IHDR" + width.to_bytes(4, "big") + height.to_bytes(4, "big")
return "data:image/png;base64," + base64.b64encode(ihdr + b"\x08\x06\x00\x00\x00").decode()

View file

@ -329,3 +329,115 @@ async def test_remote_image_fetch_keeps_counting_handler_event_loop_responsive(
]}],
}
assert messages == original
@pytest.mark.parametrize(
"config_type", (AnthropicCountTokensConfig, AzureAIAnthropicCountTokensConfig)
)
def test_count_lifts_the_leading_system_run_into_system(
config_type: type[AnthropicCountTokensConfig],
) -> None:
"""A Responses ``instructions`` arrives as a leading system-role message. count_tokens answers 400
on that role at the head of ``messages`` and only takes the initial prompt in ``system``, so the
leading run moves there with its cache_control, empty text dropped, and a later reminder stays."""
cache_control: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
messages: Final[list[dict[str, JsonValue]]] = [
{"role": "system", "content": "Be terse", "cache_control": cache_control},
{"role": "system", "content": [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": ""}]},
{"role": "user", "content": "Hello, how are you?"},
{"role": "system", "content": "later reminder"},
{"role": "assistant", "content": "Bonjour."},
]
original: Final = deepcopy(messages)
result: Final = config_type().transform_request_to_count_tokens(model="claude-opus-5-5", messages=messages)
assert result == {
"model": "claude-opus-5-5",
"system": [
{"type": "text", "text": "Be terse", "cache_control": cache_control},
{"type": "text", "text": "Answer in French"},
],
"messages": [
{"role": "user", "content": "Hello, how are you?"},
{"role": "system", "content": "later reminder"},
{"role": "assistant", "content": "Bonjour."},
],
}
assert messages == original
@pytest.mark.parametrize(
("system", "expected_system"),
(
(None, [{"type": "text", "text": "Be terse"}]),
("", [{"type": "text", "text": "Be terse"}]),
("Answer in French", [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": "Be terse"}]),
(
[{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}}],
[
{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": "Be terse"},
],
),
),
ids=["absent", "empty", "string", "blocks"],
)
def test_count_keeps_the_callers_system_ahead_of_the_lifted_run(
system: JsonValue, expected_system: list[dict[str, JsonValue]]
) -> None:
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5",
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "hi"}],
system=system,
)
assert result == {
"model": "claude-opus-5-5",
"system": expected_system,
"messages": [{"role": "user", "content": "hi"}],
}
def test_count_leaves_a_non_text_system_and_its_messages_as_sent() -> None:
"""A malformed ``system`` is the provider's to reject, so nothing is rearranged around it."""
messages: Final[list[dict[str, JsonValue]]] = [
{"role": "system", "content": "Be terse"},
{"role": "user", "content": "hi"},
]
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5", messages=messages, system=5
)
assert result == {"model": "claude-opus-5-5", "system": 5, "messages": messages}
def test_count_drops_a_leading_system_message_without_text() -> None:
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5",
messages=[{"role": "system", "content": ""}, {"role": "user", "content": "hi"}],
)
assert result == {"model": "claude-opus-5-5", "messages": [{"role": "user", "content": "hi"}]}
@pytest.mark.asyncio
async def test_handler_sends_the_leading_system_run_as_system_not_as_a_message(httpx_transport_clients):
"""The wire body is what the provider judges: ``system`` carries the prompt and no message has
``role: "system"``, so a Responses ``instructions`` is counted by Anthropic instead of 400ing."""
with respx.mock:
route = respx.post("https://gateway.example/v1/messages/count_tokens").mock(
return_value=httpx.Response(200, json={"input_tokens": 21})
)
result = await AnthropicCountTokensHandler().handle_count_tokens_request(
model="claude-opus-5-5",
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "Hello, how are you?"}],
auth_header={"x-api-key": "sk-ant-api03-test-key"},
api_base="https://gateway.example",
)
assert result == {"input_tokens": 21}
assert TypeAdapter(dict[str, JsonValue]).validate_json(route.calls.last.request.content) == {
"model": "claude-opus-5-5",
"system": [{"type": "text", "text": "Be terse"}],
"messages": [{"role": "user", "content": "Hello, how are you?"}],
}

View file

@ -33,7 +33,7 @@ def test_discovery_only_exposes_authorized_completed_support(revision, transport
client_extensions=frozenset({"io.modelcontextprotocol/ui"}),
upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}),
)
assert result.supported_versions == [revision]
assert result.supported_versions == ([revision] if transport is MCPTransport.sse else [revision, "2026-07-28"])
assert result.capabilities.tools is not None
assert result.capabilities.prompts is None
assert result.capabilities.resources is None
@ -43,7 +43,7 @@ def test_discovery_only_exposes_authorized_completed_support(revision, transport
assert result.ttl_ms == 0
@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})])
@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"})])
def test_unproven_translation_never_advertises_operations(upstream):
result = build_discovery(
configured=HANDSHAKE_PROTOCOL_VERSIONS,
@ -82,12 +82,12 @@ def test_discovery_results_do_not_share_mutable_capabilities():
assert capabilities.tools.list_changed is not True
def test_modern_candidates_do_not_enable_public_serving():
modern = REVISION_SUPPORT["2026-07-28"]
assert modern.completed is False
def test_modern_support_excludes_legacy_sse() -> None:
modern: Final = REVISION_SUPPORT["2026-07-28"]
assert modern.completed is True
assert "input_required" in modern.results
assert MCPTransport.sse not in modern.transports
assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS)
assert ("2026-07-28", "2026-07-28") in TRANSLATION_PAIRS
@pytest.mark.asyncio
@ -104,3 +104,39 @@ async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted)
with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True):
async with Client(server, mode="legacy"):
pytest.fail("The excluded revision must not initialize")
def test_modern_discovery_requires_opt_in_and_keeps_unsupported_features_disabled() -> None:
from pydantic import TypeAdapter
from litellm.types.mcp import MCPAdvertisedVersions
configured: Final = TypeAdapter(MCPAdvertisedVersions).validate_python(["2026-07-28"])
result: Final = build_discovery(
configured=configured,
revision="2026-07-28",
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=frozenset({"2026-07-28"}),
capabilities=ServerCapabilities(
tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability()
),
)
assert result.supported_versions == ["2026-07-28"]
assert result.capabilities.tools is not None
assert result.capabilities.prompts is not None
assert result.capabilities.resources is not None
assert result.capabilities.tasks is None
assert result.capabilities.extensions is None
@pytest.mark.parametrize("path", ["/mcp/sse", "/mcp/example/sse/"])
def test_modern_protocol_is_rejected_on_legacy_sse_paths(path: str, monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
monkeypatch.setitem(proxy_server.general_settings, "mcp_advertised_versions", ["2025-11-25", "2026-07-28"])
assert unsupported_protocol_version({"path": "/mcp", "headers": [(b"mcp-protocol-version", b"2026-07-28")]}) is None
assert (
unsupported_protocol_version({"path": path, "headers": [(b"mcp-protocol-version", b"2026-07-28")]})
== "2026-07-28"
)

View file

@ -0,0 +1,253 @@
from typing import Final, Literal
import pytest
from mcp import MCPError
from mcp.types import CallToolRequest, CallToolRequestParams, InputRequiredResult
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.interactions import bind_target, open_continuation, seal_continuation
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def test_continuation_is_repeatable_and_bound_to_caller_operation_and_expiry(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", "local-continuation-test")
context: Final = OperationContext(
_caller=UserAPIKeyAuth(user_id="alice", team_id="team"), mcp_servers=("upstream",)
)
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={"amount": 3}))
server: Final = MCPServer(server_id="upstream", name="upstream", url="https://example.com/mcp", transport="http")
sealed: Final = seal_continuation(
bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100
)
retry: Final = operation.model_copy(
update={"params": operation.params.model_copy(update={"request_state": sealed.request_state})}
)
state: Final = open_continuation(retry, context, now=101)
assert state is not None
assert state.upstream_state == "opaque"
assert open_continuation(retry, context, now=102) == state
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(retry, OperationContext(_caller=UserAPIKeyAuth(user_id="bob", team_id="team")), now=101)
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(
retry.model_copy(update={"params": retry.params.model_copy(update={"arguments": {"amount": 4}})}),
context,
now=101,
)
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(retry, context, now=700)
monkeypatch.setenv("LITELLM_SALT_KEY", "rotated-test-salt")
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(retry, context, now=101)
def test_continuation_missing_salt_does_not_reject_initial_request(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice"))
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={}))
assert open_continuation(operation, context, now=100) is None
server: Final = MCPServer(server_id="upstream", name="upstream", url="https://example.com/mcp", transport="http")
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
seal_continuation(bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100)
@pytest.mark.parametrize("identity", [None, UserAPIKeyAuth(), UserAPIKeyAuth(team_id="team")])
def test_continuation_requires_stable_authenticated_principal(
identity: UserAPIKeyAuth | None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={}))
server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http")
with pytest.raises(MCPError, match="authenticated caller identity"):
seal_continuation(
bind_target(InputRequiredResult(request_state="opaque"), server),
operation,
OperationContext(_caller=identity),
now=100,
)
def test_continuation_keeps_original_expiry_and_survives_credential_rotation(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
before: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", api_key="old-test-key"))
after: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", api_key="new-test-key"))
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={"a": 1, "b": 2}))
server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http")
bound: Final = bind_target(InputRequiredResult(request_state="opaque"), server)
first: Final = seal_continuation(bound, operation, before, now=100)
retry: Final = operation.model_copy(
update={
"params": operation.params.model_copy(
update={"request_state": first.request_state, "arguments": {"b": 2, "a": 1}}
)
}
)
state: Final = open_continuation(retry, after, now=200)
assert state is not None
second: Final = seal_continuation(bound, retry, after, now=699, previous=state)
final: Final = retry.model_copy(
update={"params": retry.params.model_copy(update={"request_state": second.request_state})}
)
assert open_continuation(final, after, now=699) == state
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(final, after, now=700)
@pytest.mark.parametrize("change", ["tamper", "team", "org", "principal", "method"])
def test_altered_continuations_are_rejected(change: str, monkeypatch: pytest.MonkeyPatch) -> None:
from mcp.types import GetPromptRequest, GetPromptRequestParams
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", team_id="one", org_id="org"))
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={}))
server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http")
first: Final = seal_continuation(
bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100
)
token: Final = first.request_state
assert token is not None
retry: Final = (
GetPromptRequest(params=GetPromptRequestParams(name="confirm", arguments={}, request_state=token))
if change == "method"
else operation.model_copy(
update={
"params": operation.params.model_copy(
update={"request_state": token + "a" if change == "tamper" else token}
)
}
)
)
caller: Final = UserAPIKeyAuth(
user_id="bob" if change == "principal" else "alice",
team_id="two" if change == "team" else "one",
org_id="other" if change == "org" else "org",
)
with pytest.raises(MCPError) as error:
open_continuation(retry, OperationContext(_caller=caller), now=101)
assert error.value.error.code == -32602
def test_input_responses_without_state_and_unbound_results_are_rejected(monkeypatch: pytest.MonkeyPatch) -> None:
from mcp.types import ElicitResult
from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice"))
operation: Final = CallToolRequest(
params=CallToolRequestParams(name="confirm", input_responses={"consent": ElicitResult(action="accept")})
)
with pytest.raises(MCPError, match="require a gateway continuation"):
open_continuation(operation, context, now=100)
with pytest.raises(MCPError, match="target is unavailable"):
seal_continuation(BoundInputRequiredResult(request_state="opaque"), operation, context, now=100)
def test_authenticated_but_invalid_state_payload_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server.state_tokens import seal_state
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
sealed: Final = seal_state(
{"unsupported": "state-schema"}, purpose="mcp:interaction:repeatable:v1", expires_at=200, now=100
)
assert isinstance(sealed, Ok)
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", request_state=sealed.ok))
with pytest.raises(MCPError, match="Invalid or expired"):
open_continuation(operation, OperationContext(_caller=UserAPIKeyAuth(user_id="alice")), now=101)
@pytest.mark.asyncio
async def test_modern_sampling_errors_abort_without_relaying_form_input() -> None:
import anyio
from mcp import ClientSession
from mcp.shared.message import SessionMessage
from mcp.types import (
CreateMessageRequest,
CreateMessageRequestParams,
ElicitRequest,
ElicitRequestFormParams,
ErrorData,
)
from litellm.proxy._experimental.mcp_server.interactions import ModernClientInteraction
async def sampling(context: object, params: CreateMessageRequestParams) -> ErrorData:
return ErrorData(code=-32603, message="Sampling refused by gateway policy")
send, receive = anyio.create_memory_object_stream[SessionMessage](1)
try:
interaction: Final = ModernClientInteraction(ClientSession(receive, send, sampling_callback=sampling), allow_elicitation=True)
form: Final = ElicitRequest(
params=ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object", "properties": {}})
)
refused: Final = await interaction.request("consent", form)
assert isinstance(refused, ErrorData)
assert refused.code == -32602
with pytest.raises(MCPError, match="Sampling refused by gateway policy"):
await interaction.prepare(
InputRequiredResult(
request_state="opaque",
input_requests={
"consent": form,
"sample": CreateMessageRequest(params=CreateMessageRequestParams(messages=[], max_tokens=1)),
},
)
)
finally:
await send.aclose()
await receive.aclose()
@pytest.mark.asyncio
@pytest.mark.parametrize("invalid_response", ["legacy_state", "gateway_override", "unbound_result"])
async def test_gateway_rejects_invalid_interaction_boundaries(
invalid_response: Literal["legacy_state", "gateway_override", "unbound_result"], monkeypatch: pytest.MonkeyPatch
) -> None:
from unittest.mock import AsyncMock, patch
from mcp.types import CreateMessageResult, ElicitResult, TextContent
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.contracts import WireCompat
from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
context: Final = OperationContext(
_caller=UserAPIKeyAuth(user_id="alice"),
wire_compat=WireCompat.LEGACY if invalid_response == "legacy_state" else WireCompat.MODERN,
)
operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={}))
server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http")
if invalid_response == "legacy_state":
operation.params.request_state = "untrusted-state"
elif invalid_response == "gateway_override":
bound: Final = bind_target(
BoundInputRequiredResult(
request_state="opaque",
gateway_responses={
"sample": CreateMessageResult(
role="assistant", content=TextContent(type="text", text="gateway-owned"), model="test-model"
)
},
),
server,
)
sealed: Final = seal_continuation(bound, operation, context, now=100)
operation.params.request_state = sealed.request_state
operation.params.input_responses = {"sample": ElicitResult(action="accept")}
message: Final = {
"legacy_state": "Continuations require the modern MCP protocol",
"gateway_override": "Cannot replace gateway input responses",
"unbound_result": "MCP continuation target is unavailable",
}[invalid_response]
dispatch: Final = AsyncMock(return_value=InputRequiredResult(request_state="unbound-upstream-state"))
with (
patch.object(operations, "_execute_mcp_server_tool_call", dispatch),
patch.object(operations.global_mcp_server_manager, "get_mcp_server_by_id", return_value=server),
patch.object(operations.time, "time", return_value=101),
):
with pytest.raises(MCPError, match=message) as rejected:
await operations.GatewayOperations().execute(operation, context)
assert rejected.value.error.code == -32602
if invalid_response == "unbound_result":
dispatch.assert_awaited_once()
else:
dispatch.assert_not_awaited()

View file

@ -314,7 +314,12 @@ class TestMCPClientUnitTests:
assert result == mock_result
mock_session_instance.initialize.assert_called_once()
mock_session_instance.call_tool.assert_called_once_with(
name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY, allow_input_required=False
name="test_tool",
arguments={"arg1": "value1"},
input_responses=None,
request_state=None,
progress_callback=ANY,
allow_input_required=True,
)

View file

@ -164,7 +164,9 @@ async def test_elicitation_callback_keeps_initiating_session():
request = AsyncMock(return_value=accepted)
initiating = SimpleNamespace(client_params=SimpleNamespace(capabilities=capabilities), elicit_form=request)
token = legacy_server.active_mcp_session_var.set(initiating)
request_token = active_mcp_request_ctx_var.set(SimpleNamespace(session=initiating, request_id="initiating-call"))
request_token = active_mcp_request_ctx_var.set(
SimpleNamespace(session=initiating, request_id="initiating-call", protocol_version="2025-11-25")
)
try:
callback = _create_elicitation_callback()
legacy_server.active_mcp_session_var.set(SimpleNamespace())
@ -4323,7 +4325,9 @@ class TestMCPServerManager:
mock_create_client.assert_called_once()
called_kwargs = mock_create_client.call_args.kwargs
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"}
mock_client.read_resource.assert_awaited_once_with("https://example.com/resource")
mock_client.read_resource.assert_awaited_once_with(
"https://example.com/resource", input_responses=None, request_state=None, allow_input_required=False
)
assert result is read_result
@pytest.mark.asyncio
@ -19275,3 +19279,32 @@ async def test_repeated_stale_discovery_uses_current_callers_endpoint(endpoint:
original, needed_endpoint=lambda server: getattr(server, endpoint), retry_stale=False,
)
assert resolved is replacement
@pytest.mark.asyncio
async def test_legacy_upstream_elicitation_rejects_modern_downstream_without_consent() -> None:
from types import SimpleNamespace
from mcp.types import ElicitRequestFormParams, ErrorData
from litellm.proxy._experimental.mcp_server import server as legacy_server
from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_elicitation_callback
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
session: Final = SimpleNamespace(client_params=None)
session_token: Final = legacy_server.active_mcp_session_var.set(session)
request_token: Final = active_mcp_request_ctx_var.set(
SimpleNamespace(session=session, request_id="modern-call", protocol_version="2026-07-28")
)
relay: Final = AsyncMock()
try:
callback: Final = create_elicitation_callback()
with patch("litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", relay):
result: Final = await callback(
None, ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object"})
)
assert isinstance(result, ErrorData)
assert result.code == -32602
assert "may have partially completed" in result.message
relay.assert_not_awaited()
finally:
active_mcp_request_ctx_var.reset(request_token)
legacy_server.active_mcp_session_var.reset(session_token)

View file

@ -6,7 +6,7 @@ import json
import os
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Final
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -27,7 +27,7 @@ from mcp.types import (
)
from mcp.types import Tool as MCPTool
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
from pydantic import TypeAdapter
from pydantic import JsonValue, TypeAdapter
from starlette.types import Message, Receive, Scope, Send
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
@ -1086,6 +1086,9 @@ async def test_mcp_get_prompt_success():
extra_headers={"X-Test": "1"},
raw_headers=None,
client_ip=None,
input_responses=None,
request_state=None,
allow_input_required=False,
)
assert result is prompt_result
@ -1149,6 +1152,9 @@ async def test_mcp_read_resource_success():
extra_headers={"X-Test": "1"},
raw_headers=None,
client_ip=None,
input_responses=None,
request_state=None,
allow_input_required=False,
)
assert result is read_result
@ -11149,3 +11155,285 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
context = dispatched.await_args.args[1]
assert context.user_api_key_auth.user_id == "discover-caller"
assert context.mcp_servers == ("allowed",)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"method,params",
(
("tools/call", {"name": "confirm", "arguments": {}}),
("prompts/get", {"name": "confirm"}),
("resources/read", {"uri": "test://confirm"}),
),
)
@pytest.mark.parametrize("continuation", ("initial", "valid", "tampered", "revoked"))
@pytest.mark.parametrize("large_body", (False, True))
@pytest.mark.parametrize("passthrough", (False, True))
@pytest.mark.parametrize("route_name", ("interactive", "friendly"))
async def test_modern_oauth_challenge_follows_continuation_authorization(
method: str,
params: dict[str, JsonValue],
continuation: Literal["initial", "valid", "tampered", "revoked"],
large_body: bool,
passthrough: bool,
route_name: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import time
from pydantic import TypeAdapter
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.interactions import InteractionOperation, bind_target, seal_continuation
from mcp.types import InputRequiredResult
monkeypatch.setenv("LITELLM_SALT_KEY", "challenge-test-salt")
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"mcp_advertised_versions": ["2026-07-28"]})
caller: Final = UserAPIKeyAuth(user_id="alice")
target: Final = _make_oauth2_server("interactive", oauth2_flow="authorization_code").model_copy(
update={
"alias": "friendly",
**(
{"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True}
if passthrough
else {}
),
}
)
operation: Final = TypeAdapter(InteractionOperation).validate_python({"method": method, "params": params})
context: Final = OperationContext(_caller=caller, mcp_servers=(route_name,))
sealed: Final = seal_continuation(
bind_target(InputRequiredResult(request_state="upstream"), target), operation, context, now=int(time.time())
)
request_params: Final = {
**params,
**(
{"requestState": "invalid" if continuation == "tampered" else sealed.request_state}
if continuation != "initial"
else {}
),
}
body: Final = json.dumps(
{
"jsonrpc": "2.0",
"id": 7,
"method": method,
"params": {
**request_params,
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {},
"padding": "x" * (server._MCP_ROUTING_PEEK_MAX_BYTES * 2 if large_body else 0),
},
},
}
).encode()
route_value: Final = params["uri" if method == "resources/read" else "name"]
assert isinstance(route_value, str)
scope: Final[Scope] = {
"type": "http",
"method": "POST",
"path": f"/mcp/{route_name}",
"scheme": "http",
"server": ("localhost", 8000),
"query_string": b"",
"root_path": "",
"headers": [
(b"content-type", b"application/json"),
(b"x-litellm-api-key", b"test-gateway-key"),
(b"authorization", b"Bearer expired-upstream-token"),
(b"mcp-protocol-version", b"2026-07-28"),
(b"mcp-method", method.encode()),
(b"mcp-name", route_value.encode()),
(b"accept", b"application/json, text/event-stream"),
],
}
receive: Final = AsyncMock(
side_effect=[
{"type": "http.request", "body": body[: len(body) // 2], "more_body": True},
{"type": "http.request", "body": body[len(body) // 2 :], "more_body": False},
]
)
send: Final = AsyncMock()
with (
patch.object(
server,
"extract_mcp_auth_context",
AsyncMock(
return_value=(
caller,
None,
[route_name],
None,
{"Authorization": "Bearer expired-upstream-token"},
None,
)
),
),
patch.object(server, "_SESSION_MANAGERS_INITIALIZED", True),
patch.object(server, "_probe_upstream_auth", AsyncMock(return_value=(401, None))) as probe,
patch.object(server.session_manager_stateless, "handle_request", AsyncMock()) as dispatch,
patch.object(
mcp_operations,
"_get_allowed_mcp_servers",
AsyncMock(return_value=[] if continuation == "revoked" else [target]),
),
patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_by_id", return_value=target),
patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=target),
patch.object(
mcp_operations.global_mcp_server_manager, "ensure_oauth_metadata_discovered", AsyncMock(return_value=target)
) as discovery,
patch.object(
mcp_operations.global_mcp_server_manager, "has_user_oauth_token", AsyncMock(return_value=False)
) as token,
):
if continuation in ("initial", "valid"):
with pytest.raises(HTTPException) as rejected:
await server.handle_streamable_http_mcp(scope, receive, send)
assert rejected.value.status_code == 401
if passthrough:
assert "resource_metadata=" in rejected.value.headers["www-authenticate"]
assert "invalid_token" in rejected.value.headers["www-authenticate"]
probe.assert_awaited_once_with(target.url, "Bearer expired-upstream-token")
token.assert_not_awaited()
else:
assert (
f'Bearer authorization_uri="http://localhost:8000/.well-known/oauth-authorization-server/mcp/{route_name}"'
== rejected.value.headers["www-authenticate"]
)
token.assert_awaited_once()
probe.assert_not_awaited()
elif continuation == "revoked":
with pytest.raises(HTTPException) as rejected:
await server.handle_streamable_http_mcp(scope, receive, send)
assert rejected.value.status_code == 403
probe.assert_not_awaited()
discovery.assert_not_awaited()
token.assert_not_awaited()
else:
await server.handle_streamable_http_mcp(scope, receive, send)
response_body: Final = json.loads(
next(
call.args[0]["body"]
for call in send.await_args_list
if call.args[0]["type"] == "http.response.body"
)
)
assert response_body["error"]["code"] == -32602
probe.assert_not_awaited()
discovery.assert_not_awaited()
token.assert_not_awaited()
dispatch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("tools/call", "prompts/get", "resources/read"))
@pytest.mark.parametrize("size_delta", (-1, 0, 1))
@pytest.mark.parametrize("chunked", (False, True))
async def test_modern_preflight_enforces_sdk_body_limit(
method: str, size_delta: int, chunked: bool, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy._experimental.mcp_server import server
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"mcp_advertised_versions": ["2026-07-28"]})
limit: Final = server.session_manager_stateless.max_request_body_size
params: Final = {"uri": "test://confirm"} if method == "resources/read" else {"name": "confirm"}
payload: Final = json.dumps(
{
"jsonrpc": "2.0",
"id": 7,
"method": method,
"params": {
**params,
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {},
"padding": "",
},
},
}
).encode()
body: Final = payload.replace(
b'"padding": ""', b'"padding": "' + b"x" * (limit + size_delta - len(payload)) + b'"', 1
)
assert len(body) == limit + size_delta
chunks: Final = (body[: limit - 1], body[limit - 1 :]) if chunked else (body,)
messages: Final[list[Message]] = [
{"type": "http.request", "body": chunk, "more_body": index < len(chunks) - 1 or size_delta > 0}
for index, chunk in enumerate(chunks)
]
receive: Final = AsyncMock(side_effect=[*messages, {"type": "http.request", "body": b"", "more_body": False}])
scope: Final[Scope] = {
"type": "http",
"method": "POST",
"path": "/mcp",
"scheme": "http",
"server": ("localhost", 8000),
"query_string": b"",
"root_path": "",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer test-key"),
(b"mcp-protocol-version", b"2026-07-28"),
(b"mcp-method", method.encode()),
(b"mcp-name", next(iter(params.values())).encode()),
*([] if chunked else [(b"content-length", str(len(body)).encode())]),
],
}
with (
patch.object(
server, "extract_mcp_auth_context", AsyncMock(return_value=(UserAPIKeyAuth(), None, None, None, None, None))
),
patch.object(server, "_SESSION_MANAGERS_INITIALIZED", True),
patch.object(server.session_manager_stateless, "handle_request", AsyncMock()) as dispatch,
patch.object(server.operations, "_get_allowed_mcp_servers", AsyncMock()) as authorization,
):
if size_delta > 0:
with pytest.raises(HTTPException) as rejected:
await server.handle_streamable_http_mcp(scope, receive, AsyncMock())
assert rejected.value.status_code == 413
assert rejected.value.detail == "Request body too large"
dispatch.assert_not_awaited()
else:
await server.handle_streamable_http_mcp(scope, receive, AsyncMock())
dispatch.assert_awaited_once()
assert receive.await_count == len(chunks)
authorization.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("invalid_request", ("malformed", "other_method", "header_mismatch", "duplicate_header"))
async def test_modern_preflight_leaves_invalid_envelopes_to_sdk_without_upstream_work(
invalid_request: Literal["malformed", "other_method", "header_mismatch", "duplicate_header"],
) -> None:
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
envelope: Final = {
"jsonrpc": "2.0",
"id": 7,
"method": "tools/list" if invalid_request == "other_method" else "tools/call",
"params": {
"name": "confirm",
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {},
},
},
}
headers: Final = [
(b"mcp-protocol-version", b"2026-07-28"),
(b"mcp-method", b"tools/call"),
(b"mcp-name", b"wrong" if invalid_request == "header_mismatch" else b"confirm"),
]
scope: Final[Scope] = {
"type": "http",
"headers": [*headers, *([(b"mcp-method", b"tools/call")] if invalid_request == "duplicate_header" else [])],
}
with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock()) as resolve:
result: Final = await server._preflight_modern_interaction(
scope,
b"{" if invalid_request == "malformed" else json.dumps(envelope).encode(),
OperationContext(_caller=UserAPIKeyAuth(user_id="alice"), mcp_servers=("interactive",)),
)
assert result is None
resolve.assert_not_awaited()

View file

@ -108,6 +108,7 @@ def test_acreate_credential_success(cli_runner, mock_credentials_client):
"test-cred",
{"custom_llm_provider": "azure"},
{"api_key": "test-key"},
display_name=None,
)
@ -215,3 +216,55 @@ def test_aget_credential_success(cli_runner, mock_credentials_client):
output_data = json.loads(result.output)
assert output_data == mock_response
mock_instance.get.assert_called_once_with("test-cred")
def test_list_credentials_table_shows_display_name_and_source(cli_runner, mock_credentials_client):
mock_credentials_client.return_value.list.return_value = {
"credentials": [
{"credential_name": "openai-prod", "display_name": "Prod OpenAI", "source": "db", "credential_info": {}},
{"credential_name": "from-yaml", "display_name": None, "source": "config", "credential_info": {}},
]
}
result = cli_runner.invoke(cli, ["credentials", "list"])
assert result.exit_code == 0
assert "Prod OpenAI" in result.output
assert "config" in result.output
assert "None" not in result.output
def test_create_credential_forwards_the_display_name(cli_runner, mock_credentials_client):
mock_instance = mock_credentials_client.return_value
mock_instance.create.return_value = {"success": True}
result = cli_runner.invoke(
cli,
["credentials", "create", "openai-prod", "--info", "{}", "--values", "{}", "--display-name", "Prod OpenAI"],
)
assert result.exit_code == 0
mock_instance.create.assert_called_once_with("openai-prod", {}, {}, display_name="Prod OpenAI")
@pytest.mark.parametrize(
("flags", "expected"),
[(["--display-name", "Prod OpenAI"], "Prod OpenAI"), (["--clear-display-name"], None)],
)
def test_update_credential_sets_or_clears_the_display_name(cli_runner, mock_credentials_client, flags, expected):
mock_instance = mock_credentials_client.return_value
mock_instance.update_display_name.return_value = {"success": True}
result = cli_runner.invoke(cli, ["credentials", "update", "openai-prod", *flags])
assert result.exit_code == 0, result.output
mock_instance.update_display_name.assert_called_once_with("openai-prod", expected)
@pytest.mark.parametrize("flags", [[], ["--display-name", "Prod", "--clear-display-name"]])
def test_update_credential_requires_exactly_one_display_name_flag(cli_runner, mock_credentials_client, flags):
result = cli_runner.invoke(cli, ["credentials", "update", "openai-prod", *flags])
assert result.exit_code != 0
assert "exactly one" in result.output
mock_credentials_client.return_value.update_display_name.assert_not_called()

View file

@ -130,6 +130,31 @@ def test_create_request(client, base_url, api_key):
}
def test_create_request_includes_the_display_name_only_when_given(client):
labeled = client.create("azure1", {}, {"api_key": "sk-123"}, return_request=True, display_name="Azure EU")
unlabeled = client.create("azure1", {}, {"api_key": "sk-123"}, return_request=True)
assert labeled.json["display_name"] == "Azure EU"
assert "display_name" not in unlabeled.json
@pytest.mark.parametrize("display_name", ["Azure EU", None])
def test_update_display_name_request(client, base_url, api_key, display_name):
request = client.update_display_name("azure1", display_name, return_request=True)
assert request.method == "PATCH"
assert request.url == f"{base_url}/credentials/azure1"
assert request.headers["Authorization"] == f"Bearer {api_key}"
assert request.json == {"display_name": display_name, "credential_info": {}}
@pytest.mark.parametrize(("credential_name", "path"), [("team?prod", "team%3Fprod"), ("a/b#c", "a%2Fb%23c")])
def test_update_display_name_targets_the_whole_credential_name(client, base_url, credential_name, path):
request = client.update_display_name(credential_name, "Label", return_request=True)
assert request.prepare().url == f"{base_url}/credentials/{path}"
@responses.activate
def test_create_mock_response(client):
"""Test create with a mocked successful response"""

View file

@ -14,6 +14,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.credential_endpoints.endpoints import get_llm_router
from litellm.proxy.proxy_server import app
from litellm.models.credentials import CredentialSource
from litellm.types.utils import CredentialItem
client = TestClient(app)
@ -771,7 +772,7 @@ class TestNonAdminCannotPersistWifFieldsOnCredential:
update_mock.assert_awaited_once()
def _wif_credential(name: str = "federated-cred") -> CredentialItem:
def _wif_credential(name: str = "federated-cred", source: CredentialSource = "db") -> CredentialItem:
return CredentialItem(
credential_name=name,
credential_values={
@ -779,6 +780,7 @@ def _wif_credential(name: str = "federated-cred") -> CredentialItem:
"api_key": "sk-old",
},
credential_info={"custom_llm_provider": "anthropic"},
source=source,
)
@ -978,7 +980,7 @@ class TestNonAdminCannotTouchAStoredWifCredential:
def test_non_admin_cannot_delete_a_config_only_wif_credential(self, restore_credential_list, monkeypatch):
"""A ``credential_list`` entry from config.yaml has no DB row, so a gate that consulted
only the DB let a non-admin evict the admin-owned federation settings from memory."""
config_credential = _wif_credential("config-wif")
config_credential = _wif_credential("config-wif", source="config")
monkeypatch.setattr(litellm, "credential_list", [config_credential])
with _repository_holding(None) as repository:
response = _delete_credential("config-wif", auth=_as_non_admin)
@ -989,15 +991,17 @@ class TestNonAdminCannotTouchAStoredWifCredential:
assert litellm.credential_list == [config_credential]
def test_proxy_admin_can_delete_a_config_only_wif_credential(self, restore_credential_list, monkeypatch):
"""The gate lets the admin through to the row delete. The 404 that follows is the rule for
"""The gate lets the admin through to the row delete. The 400 that follows is the rule for
every config-only credential (no row to delete, the entry is back on the next boot), so the
in-memory entry stays put too."""
config_credential = _wif_credential("config-wif")
config_credential = _wif_credential("config-wif", source="config")
monkeypatch.setattr(litellm, "credential_list", [config_credential])
with _repository_holding(None) as repository:
response = _delete_credential("config-wif", auth=_as_admin)
assert response.status_code == 404, response.text
assert response.status_code == 405, response.text
assert response.headers["allow"] == "GET"
assert "defined in config" in response.json()["error"]["message"]
repository.delete_by_name.assert_awaited_once_with("config-wif")
assert litellm.credential_list == [config_credential]
@ -1005,7 +1009,7 @@ class TestNonAdminCannotTouchAStoredWifCredential:
"""POST with the same name carries no WIF field and collides with no DB row, yet
``CredentialAccessor.upsert_credentials`` would replace the admin entry in memory and
the periodic config sync would then make the takeover permanent."""
config_credential = _wif_credential("config-wif")
config_credential = _wif_credential("config-wif", source="config")
monkeypatch.setattr(litellm, "credential_list", [config_credential])
with _repository_holding(None) as repository:
response = _post_credential(
@ -1024,7 +1028,7 @@ class TestNonAdminCannotTouchAStoredWifCredential:
assert litellm.credential_list[0].credential_values["api_key"] == "sk-old"
def test_proxy_admin_can_post_over_a_config_only_wif_credential(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_wif_credential("config-wif")])
monkeypatch.setattr(litellm, "credential_list", [_wif_credential("config-wif", source="config")])
with _repository_holding(None) as repository:
response = _post_credential(
{
@ -1039,50 +1043,6 @@ class TestNonAdminCannotTouchAStoredWifCredential:
repository.create.assert_awaited_once()
assert litellm.credential_list[0].credential_values == {"api_key": "sk-rotated"}
def test_non_admin_cannot_rename_a_credential_onto_a_config_only_wif_credential(
self, restore_credential_list, monkeypatch
):
"""PATCH is the other way to shadow: renaming an ordinary credential onto the WIF
credential's name makes ``_sync_in_memory_credential`` upsert the attacker's values over
the admin entry, with no WIF field in the payload and no DB row to collide with."""
config_credential = _wif_credential("config-wif")
monkeypatch.setattr(litellm, "credential_list", [_plain_credential("mine"), config_credential])
with _repository_holding(_plain_credential("mine")) as repository:
response = _patch_credential(
"mine",
{
"credential_name": "config-wif",
"credential_values": {"api_key": "sk-attacker"},
"credential_info": {},
},
auth=_as_non_admin,
)
assert response.status_code == 403, response.text
assert "anthropic_keycloak_token_url" in response.text
repository.update_by_name.assert_not_awaited()
assert config_credential in litellm.credential_list
assert litellm.credential_list[1].credential_values["api_key"] == "sk-old"
def test_proxy_admin_can_rename_a_credential_onto_a_config_only_wif_credential(
self, restore_credential_list, monkeypatch
):
monkeypatch.setattr(litellm, "credential_list", [_plain_credential("mine"), _wif_credential("config-wif")])
with _repository_holding(_plain_credential("mine")) as repository:
response = _patch_credential(
"mine",
{
"credential_name": "config-wif",
"credential_values": {"api_key": "sk-rotated"},
"credential_info": {},
},
auth=_as_admin,
)
assert response.status_code == 200, response.text
repository.update_by_name.assert_awaited_once()
assert [c.credential_name for c in litellm.credential_list] == ["config-wif"]
def test_non_admin_cannot_post_a_null_wif_field(self, restore_credential_list):
"""Same key-presence rule on the create path: ``{"anthropic_issuer_url": null}`` persists
the key, and the resolver reacts to the key."""
@ -1228,17 +1188,20 @@ def test_delete_credential_still_answers_200_and_drops_the_credential_from_memor
def test_delete_credential_leaves_a_credential_that_only_exists_in_memory_in_place(credential_store):
"""A credential declared in the config yaml is never written to the table, so the delete
matches no row. Reporting success would be the same lie: it comes straight back on the next
proxy boot. ``PATCH /credentials/{name}`` already answers 404 for that credential."""
proxy boot."""
config_only = CredentialItem(
credential_name="from-config-yaml",
credential_values={"api_key": "sk-config"},
credential_info={},
source="config",
)
credential_store(in_memory=(config_only,), delete_by_name=AsyncMock(return_value=None))
response = _delete_credential("from-config-yaml")
assert response.status_code == 404, response.text
assert response.status_code == 405, response.text
assert response.headers["allow"] == "GET"
assert "defined in config" in response.json()["error"]["message"]
assert [credential.credential_name for credential in litellm.credential_list] == ["from-config-yaml"]
@ -1399,3 +1362,222 @@ def test_update_credential_still_accepts_a_body_without_credential_values(creden
written = update_by_name.await_args.kwargs["data"]
assert json.loads(written["credential_info"]) == {"custom_llm_provider": "openai"}
assert set(json.loads(written["credential_values"])) == {"api_key"}, "stored values survive an info-only patch"
def _labeled_credential(name: str = "openai-prod", display_name: str | None = "Prod OpenAI") -> CredentialItem:
return CredentialItem(
credential_name=name,
display_name=display_name,
credential_values={"api_key": "sk-old"},
credential_info={"custom_llm_provider": "openai"},
)
class TestCredentialDisplayName:
def test_create_stores_the_trimmed_display_name_in_the_row_and_in_memory(self, restore_credential_list):
with _repository_holding(None) as repository:
response = _post_credential(
{
"credential_name": "openai-prod",
"display_name": " Prod OpenAI ",
"credential_values": {"api_key": "sk-new"},
"credential_info": {"custom_llm_provider": "openai"},
}
)
assert response.status_code == 200, response.text
assert repository.create.await_args.kwargs["data"]["display_name"] == "Prod OpenAI"
assert [(c.credential_name, c.display_name) for c in litellm.credential_list] == [
("openai-prod", "Prod OpenAI")
]
@pytest.mark.parametrize("display_name", ["", " ", "x" * 256])
def test_create_rejects_an_unusable_display_name(self, restore_credential_list, display_name):
with _repository_holding(None) as repository:
response = _post_credential(
{
"credential_name": "openai-prod",
"display_name": display_name,
"credential_values": {"api_key": "sk-new"},
"credential_info": {},
}
)
assert response.status_code == 400, response.text
assert response.json()["error"]["param"] == "display_name"
repository.create.assert_not_awaited()
assert litellm.credential_list == []
def test_create_accepts_a_display_name_of_exactly_the_maximum_length(self, restore_credential_list):
with _repository_holding(None) as repository:
response = _post_credential(
{
"credential_name": "openai-prod",
"display_name": "x" * 255,
"credential_values": {"api_key": "sk-new"},
"credential_info": {},
}
)
assert response.status_code == 200, response.text
assert repository.create.await_args.kwargs["data"]["display_name"] == "x" * 255
def test_patch_relabels_without_touching_the_name_or_the_values(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential("openai-prod", {"display_name": "Staging OpenAI", "credential_info": {}})
assert response.status_code == 200, response.text
repository.update_by_name.assert_awaited_once()
assert repository.update_by_name.await_args.args[0] == "openai-prod"
written = repository.update_by_name.await_args.kwargs["data"]
assert written["credential_name"] == "openai-prod"
assert written["display_name"] == "Staging OpenAI"
assert json.loads(written["credential_values"]) == {"api_key": "sk-old"}
assert [(c.credential_name, c.display_name, c.credential_values) for c in litellm.credential_list] == [
("openai-prod", "Staging OpenAI", {"api_key": "sk-old"})
]
def test_patch_without_display_name_keeps_the_stored_one(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential("openai-prod", {"credential_info": {"description": "rotated"}})
assert response.status_code == 200, response.text
assert repository.update_by_name.await_args.kwargs["data"]["display_name"] == "Prod OpenAI"
assert litellm.credential_list[0].display_name == "Prod OpenAI"
def test_patch_with_null_display_name_clears_it(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential("openai-prod", {"display_name": None, "credential_info": {}})
assert response.status_code == 200, response.text
written = repository.update_by_name.await_args.kwargs["data"]
assert "display_name" in written
assert written["display_name"] is None
assert litellm.credential_list[0].display_name is None
@pytest.mark.parametrize("display_name", ["", " ", "x" * 256])
def test_patch_rejects_an_unusable_display_name_and_changes_nothing(
self, restore_credential_list, monkeypatch, display_name
):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential("openai-prod", {"display_name": display_name, "credential_info": {}})
assert response.status_code == 400, response.text
assert response.json()["error"]["param"] == "display_name"
repository.update_by_name.assert_not_awaited()
assert litellm.credential_list == [_labeled_credential()]
def test_patch_stores_the_trimmed_display_name(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential("openai-prod", {"display_name": " Staging OpenAI ", "credential_info": {}})
assert response.status_code == 200, response.text
assert repository.update_by_name.await_args.kwargs["data"]["display_name"] == "Staging OpenAI"
assert litellm.credential_list[0].display_name == "Staging OpenAI"
def test_patch_treats_an_empty_credential_name_as_omitted(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential(
"openai-prod", {"credential_name": "", "display_name": "Staging OpenAI", "credential_info": {}}
)
assert response.status_code == 200, response.text
assert repository.update_by_name.await_args.args[0] == "openai-prod"
assert repository.update_by_name.await_args.kwargs["data"]["credential_name"] == "openai-prod"
def test_patch_rejects_a_different_credential_name_and_changes_nothing(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential(
"openai-prod",
{
"credential_name": "openai-production",
"credential_values": {"api_key": "sk-new"},
"credential_info": {},
},
)
assert response.status_code == 400, response.text
assert response.json()["error"]["param"] == "credential_name"
assert "display_name" in response.json()["error"]["message"]
repository.update_by_name.assert_not_awaited()
assert litellm.credential_list == [_labeled_credential()]
def test_patch_accepts_the_same_credential_name_echoed_back(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(_labeled_credential()) as repository:
response = _patch_credential(
"openai-prod",
{"credential_name": "openai-prod", "credential_values": {"api_key": "sk-new"}, "credential_info": {}},
)
assert response.status_code == 200, response.text
assert [(c.credential_name, c.credential_values) for c in litellm.credential_list] == [
("openai-prod", {"api_key": "sk-new"})
]
def test_patch_on_a_config_credential_is_rejected_as_config_owned(self, restore_credential_list, monkeypatch):
config_credential = _labeled_credential(display_name=None).model_copy(update={"source": "config"})
monkeypatch.setattr(litellm, "credential_list", [config_credential])
with _repository_holding(None) as repository:
response = _patch_credential("openai-prod", {"display_name": "Prod", "credential_info": {}})
assert response.status_code == 405, response.text
assert response.headers["allow"] == "GET"
assert "defined in config" in response.json()["error"]["message"]
repository.update_by_name.assert_not_awaited()
assert litellm.credential_list == [config_credential]
def test_patch_on_an_unknown_credential_still_answers_404(self, restore_credential_list):
with _repository_holding(None):
response = _patch_credential("nowhere", {"display_name": "Prod", "credential_info": {}})
assert response.status_code == 404, response.text
@pytest.mark.parametrize("method", ["PATCH", "DELETE"])
def test_a_db_credential_already_gone_from_the_db_answers_404_not_config_owned(
self, restore_credential_list, monkeypatch, method
):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential()])
with _repository_holding(None) as repository:
response = (
_patch_credential("openai-prod", {"display_name": "Prod", "credential_info": {}})
if method == "PATCH"
else _delete_credential("openai-prod")
)
assert response.status_code == 404, response.text
repository.update_by_name.assert_not_awaited()
def test_patch_on_a_name_stored_in_the_db_and_config_applies_to_the_row(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential(display_name=None)])
with _repository_holding(_labeled_credential(display_name=None)) as repository:
response = _patch_credential("openai-prod", {"display_name": "Prod", "credential_info": {}})
assert response.status_code == 200, response.text
assert repository.update_by_name.await_args.kwargs["data"]["display_name"] == "Prod"
def test_reads_return_display_name_and_source(self, restore_credential_list, monkeypatch):
config_credential = CredentialItem(
credential_name="from-config",
credential_values={"api_key": "sk-config-value"},
credential_info={},
source="config",
)
monkeypatch.setattr(litellm, "credential_list", [_labeled_credential(), config_credential])
listed = {entry["credential_name"]: entry for entry in _list_credentials().json()["credentials"]}
by_name = _call_as("GET", "/credentials/by_name/openai-prod").json()
config_by_name = _call_as("GET", "/credentials/by_name/from-config").json()
assert (listed["openai-prod"]["display_name"], listed["openai-prod"]["source"]) == ("Prod OpenAI", "db")
assert (listed["from-config"]["display_name"], listed["from-config"]["source"]) == (None, "config")
assert (by_name["display_name"], by_name["source"]) == ("Prod OpenAI", "db")
assert config_by_name["source"] == "config"
assert listed["openai-prod"]["credential_values"]["api_key"] != "sk-old"

View file

@ -6,6 +6,7 @@ LiteLLM_DailyGatewayRequests.
import asyncio
from collections.abc import Awaitable, Callable
from datetime import datetime, timezone
from typing import Final
import pytest
@ -22,8 +23,12 @@ from litellm.types.proxy.gateway_requests import GatewayRequestCounts, GatewayRe
from litellm.proxy.db.log_db_metrics import record_db_io
def _today() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%d")
_NOW: Final = datetime(2026, 3, 14, 12, 0, tzinfo=timezone.utc)
_DAY: Final = _NOW.strftime("%Y-%m-%d")
def _accumulator() -> GatewayRequestAccumulator:
return GatewayRequestAccumulator(clock=lambda: _NOW)
def _record(accumulator: GatewayRequestAccumulator, status_code: int, **overrides) -> None:
@ -38,32 +43,54 @@ def _record(accumulator: GatewayRequestAccumulator, status_code: int, **override
def test_folds_repeated_requests_into_one_key():
acc = GatewayRequestAccumulator()
acc = _accumulator()
for _ in range(3):
_record(acc, 200)
_record(acc, 500)
snapshot = acc.drain()
assert snapshot == {
GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): (
GatewayRequestKey(date=_DAY, category="llm", route="/chat/completions"): (
GatewayRequestCounts(successful_requests=3, failed_requests=1)
)
}
def test_records_each_request_under_the_date_of_its_clock_read():
accumulator: Final = GatewayRequestAccumulator(
clock=iter(
(
datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc),
datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc),
)
).__next__
)
_record(accumulator, 200)
_record(accumulator, 200)
assert accumulator.drain() == {
GatewayRequestKey(date="2026-01-31", category="llm", route="/chat/completions"): GatewayRequestCounts(
successful_requests=1, failed_requests=0
),
GatewayRequestKey(date="2026-02-01", category="llm", route="/chat/completions"): GatewayRequestCounts(
successful_requests=1, failed_requests=0
),
}
@pytest.mark.parametrize(
"status_code, expected_successful, expected_failed",
[(200, 1, 0), (201, 1, 0), (204, 1, 0), (299, 1, 0), (300, 0, 1), (400, 0, 1), (500, 0, 1)],
)
def test_success_boundary_is_2xx(status_code: int, expected_successful: int, expected_failed: int):
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, status_code)
counts = next(iter(acc.drain().values()))
assert (counts.successful_requests, counts.failed_requests) == (expected_successful, expected_failed)
def test_distinct_dimensions_do_not_merge():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200, route="/chat/completions")
_record(acc, 200, route="/embeddings")
_record(acc, 200, category=BillableCategory.MCP, route="/mcp")
@ -71,14 +98,14 @@ def test_distinct_dimensions_do_not_merge():
def test_drain_empties_the_fold():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
assert len(acc.drain()) == 1
assert acc.drain() == {}
def test_drain_snapshot_is_not_mutated_by_later_records():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
snapshot = acc.drain()
_record(acc, 200)
@ -198,7 +225,7 @@ def test_commit_skips_the_database_entirely_when_nothing_accumulated():
def test_flush_drains_and_commits():
client = FakePrismaClient()
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
asyncio.run(flush_gateway_requests(client, acc))
@ -217,7 +244,7 @@ class ExplodingClient:
def test_flush_swallows_commit_failure_so_the_scheduler_survives():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
asyncio.run(flush_gateway_requests(ExplodingClient(), acc))
@ -225,7 +252,7 @@ def test_flush_swallows_commit_failure_so_the_scheduler_survives():
def test_failed_flush_keeps_counts_for_the_next_attempt():
"""A dropped flush would silently undercount the SGR source of truth."""
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
_record(acc, 500)
@ -234,11 +261,11 @@ def test_failed_flush_keeps_counts_for_the_next_attempt():
client = FakePrismaClient()
asyncio.run(flush_gateway_requests(client, acc))
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 1, 1)]
def test_restored_counts_merge_with_requests_recorded_meanwhile():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
asyncio.run(flush_gateway_requests(ExplodingClient(), acc))
@ -246,7 +273,7 @@ def test_restored_counts_merge_with_requests_recorded_meanwhile():
client = FakePrismaClient()
asyncio.run(flush_gateway_requests(client, acc))
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)]
class ExplodingDBWithInFlightRequest:
@ -266,14 +293,14 @@ class ExplodingClientWithInFlightRequest:
def test_restore_keeps_requests_recorded_while_the_failed_write_was_in_flight():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
asyncio.run(flush_gateway_requests(ExplodingClientWithInFlightRequest(acc), acc))
client = FakePrismaClient()
asyncio.run(flush_gateway_requests(client, acc))
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 1, 1)]
# ── redis buffer ──────────────────────────────────────────────────────────────
@ -340,7 +367,7 @@ def test_non_leader_workers_push_to_redis_and_never_touch_the_database():
redis = FakeRedis()
client = FakePrismaClient()
for _ in range(3):
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
buffer, _ = _buffer(redis, leader=False)
asyncio.run(flush_gateway_requests(client, acc, buffer))
@ -354,21 +381,21 @@ def test_leader_folds_every_workers_snapshot_into_one_statement():
redis = FakeRedis()
client = FakePrismaClient()
for _ in range(50):
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
_record(acc, 500, route="/responses")
buffer, _ = _buffer(redis, leader=False)
asyncio.run(flush_gateway_requests(client, acc, buffer))
leader_acc = GatewayRequestAccumulator()
leader_acc = _accumulator()
_record(leader_acc, 200)
leader, lock = _buffer(redis, leader=True)
asyncio.run(flush_gateway_requests(client, leader_acc, leader))
assert len(client.db.statements) == 1
assert _rows_written(client) == [
(_today(), "llm", "/chat/completions", 51, 0),
(_today(), "llm", "/responses", 0, 50),
(_DAY, "llm", "/chat/completions", 51, 0),
(_DAY, "llm", "/responses", 0, 50),
]
assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
assert lock.held == [GATEWAY_REQUESTS_JOB_NAME]
@ -387,7 +414,7 @@ def test_leader_keeps_the_lease_so_staggered_pods_cost_one_statement_per_interva
for _interval in range(3):
for pod in pods:
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
asyncio.run(flush_gateway_requests(client, acc, pod))
@ -403,16 +430,16 @@ def test_leader_drains_a_backlog_deeper_than_one_capped_pop():
client = FakePrismaClient()
workers = MAX_REDIS_BUFFER_DEQUEUE_COUNT * 2 + 1
for _ in range(workers):
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
buffer, _ = _buffer(redis, leader=False)
asyncio.run(flush_gateway_requests(client, acc, buffer))
leader, _ = _buffer(redis, leader=True)
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
asyncio.run(flush_gateway_requests(client, _accumulator(), leader))
assert len(client.db.statements) == 1
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", workers, 0)]
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", workers, 0)]
assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
@ -421,7 +448,7 @@ def test_leader_with_nothing_buffered_writes_nothing():
client = FakePrismaClient()
leader, lock = _buffer(redis, leader=True)
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
asyncio.run(flush_gateway_requests(client, _accumulator(), leader))
assert client.db.statements == []
assert lock.released == []
@ -430,7 +457,7 @@ def test_leader_with_nothing_buffered_writes_nothing():
def test_leader_requeues_to_redis_when_the_database_commit_fails():
"""Counts popped from Redis are gone from every worker; a failed commit must put them back."""
redis = FakeRedis()
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
_record(acc, 200)
leader, lock = _buffer(redis, leader=True)
@ -443,8 +470,8 @@ def test_leader_requeues_to_redis_when_the_database_commit_fails():
client = FakePrismaClient()
retry, _ = _buffer(redis, leader=True)
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), retry))
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
asyncio.run(flush_gateway_requests(client, _accumulator(), retry))
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)]
class ExplodingRedis(FakeRedis):
@ -467,7 +494,7 @@ class UnwritableRedis(FakeRedis):
def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_requeue_fail():
"""The pop removed the only copy; if Redis will not take it back the leader itself must carry it."""
redis = FakeRedis()
worker_acc = GatewayRequestAccumulator()
worker_acc = _accumulator()
_record(worker_acc, 200)
_record(worker_acc, 200)
worker, _ = _buffer(redis, leader=False)
@ -475,7 +502,7 @@ def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_req
degraded = UnwritableRedis()
degraded.lists = redis.lists
leader_acc = GatewayRequestAccumulator()
leader_acc = _accumulator()
leader, _ = _buffer(degraded, leader=True)
asyncio.run(flush_gateway_requests(ExplodingClient(), leader_acc, leader))
assert degraded.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
@ -483,13 +510,13 @@ def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_req
client = FakePrismaClient()
retry, _ = _buffer(redis, leader=True)
asyncio.run(flush_gateway_requests(client, leader_acc, retry))
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
assert _rows_written(client) == [(_DAY, "llm", "/chat/completions", 2, 0)]
def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush():
"""The scheduler job must not raise, and nothing is popped so nothing needs restoring anywhere."""
redis = UnreadableRedis()
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
client = FakePrismaClient()
leader, _ = _buffer(redis, leader=True)
@ -502,7 +529,7 @@ def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush
def test_failed_redis_push_keeps_counts_locally_for_the_next_flush():
acc = GatewayRequestAccumulator()
acc = _accumulator()
_record(acc, 200)
_record(acc, 500)
buffer, lock = _buffer(ExplodingRedis(), leader=True)
@ -511,7 +538,7 @@ def test_failed_redis_push_keeps_counts_locally_for_the_next_flush():
assert lock.held == []
assert acc.drain() == {
GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): (
GatewayRequestKey(date=_DAY, category="llm", route="/chat/completions"): (
GatewayRequestCounts(successful_requests=1, failed_requests=1)
)
}

View file

@ -2,7 +2,7 @@ import asyncio
import copy
import json
import time
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Awaitable, Iterator, Mapping, Sequence
from contextlib import contextmanager
from datetime import datetime, timedelta
from types import MappingProxyType, SimpleNamespace
@ -705,6 +705,10 @@ def _test_connection_probe(
) -> Iterator[AsyncMock]:
from litellm.types.router import Deployment, LiteLLM_Params
async def run_health_check(awaitable: Awaitable[object], _timeout: float) -> dict[str, str]:
await awaitable
return {"status": "healthy"}
router: Final = MagicMock()
router.get_deployment.side_effect = lambda model_id: (
Deployment(
@ -727,7 +731,7 @@ def _test_connection_probe(
patch("litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", ahealth_check),
patch(
"litellm.proxy.health_endpoints._health_endpoints.run_with_timeout",
AsyncMock(return_value={"status": "healthy"}),
AsyncMock(side_effect=run_health_check),
),
):
yield ahealth_check
@ -874,6 +878,28 @@ async def test_test_model_connection_request_mode_wins_over_resolved_mode():
assert ahealth_check.call_args.kwargs["mode"] == "chat"
@pytest.mark.asyncio
async def test_test_model_connection_evaluation_mode_uses_decisions_handler():
deployment: Final = MappingProxyType(
{
"model_name": "typesafe/jev-latest",
"litellm_params": {"model": "typesafe/jev-latest", "api_key": "fake-typesafe-key"},
"model_info": {"id": "typesafe-jev-id"},
}
)
with _test_connection_probe(deployment) as ahealth_check:
result: Final = await health_test_model_connection(
request=MagicMock(),
mode="evaluation",
litellm_params={"model": "typesafe/jev-latest"},
model_info={"id": "typesafe-jev-id"},
user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"),
)
assert result["status"] == "success"
assert ahealth_check.call_args.kwargs["mode"] == "evaluation"
@pytest.mark.asyncio
async def test_test_model_connection_uses_loaded_deployment_team_id():
"""
@ -4665,6 +4691,41 @@ def test_test_model_connection_accepts_image_edit_mode(monkeypatch):
assert response.json()["status"] == "success"
def test_test_model_connection_accepts_evaluation_mode(monkeypatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
app = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
client = TestClient(app)
with (
patch( # test-quality-ok: endpoint reads the proxy-global DB client and 500s when it is None; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
respx.mock(assert_all_called=True) as respx_mock,
):
upstream = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(
json={
"model": "jev-latest",
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 12, "output_tokens": 1},
}
)
response = client.post(
"/health/test_connection",
json={
"mode": "evaluation",
"litellm_params": {"model": "typesafe/jev-latest", "api_key": "sk-test"},
},
)
assert response.status_code == 200, response.text
assert response.json()["status"] == "success"
assert upstream.called
def _pointfive_admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(token="admin-token", user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN)

View file

@ -2,22 +2,49 @@
Unit Tests for the max parallel request limiter v1 for the proxy
"""
import itertools
from collections.abc import Callable, Iterator
from datetime import datetime
from typing import Final
import pytest
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.parallel_request_limiter import (
PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache, hash_token
from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage
from litellm.types.utils import EmbeddingResponse, ModelResponse, TextCompletionResponse, Usage
FROZEN_INSTANT: Final = datetime(2026, 1, 31, 23, 59, 30)
LAST_MICROSECOND_OF_JANUARY: Final = datetime(2026, 1, 31, 23, 59, 59, 999999)
FIRST_MICROSECOND_OF_FEBRUARY: Final = datetime(2026, 2, 1, 0, 0, 0, 1)
LAST_MINUTE_OF_JANUARY: Final = "2026-01-31-23-59"
FIRST_MINUTE_OF_FEBRUARY: Final = "2026-02-01-00-00"
TORN_MINUTE_OF_JANUARY: Final = "2026-01-31-00-00"
def _frozen_clock() -> datetime:
return FROZEN_INSTANT
def _clock_reading(instants: Iterator[datetime]) -> Callable[[], datetime]:
return lambda: next(instants)
def _clock_rolling_over_after_first_read() -> Callable[[], datetime]:
return _clock_reading(
itertools.chain([LAST_MICROSECOND_OF_JANUARY], itertools.repeat(FIRST_MICROSECOND_OF_FEBRUARY))
)
@pytest.mark.asyncio
async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token():
handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
handler = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()), clock=_frozen_clock
)
session = UserAPIKeyAuth(
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
user_id="alice",
@ -30,7 +57,7 @@ async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_t
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
precise_minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
precise_minute = FROZEN_INSTANT.strftime("%Y-%m-%d-%H-%M")
counted = await handler.internal_usage_cache.async_get_cache(
key=f"cli-session-alice::{precise_minute}::request_count", litellm_parent_otel_span=None
)
@ -63,13 +90,10 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
end_user_id = "customer-1"
parallel_request_handler = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
internal_usage_cache=InternalUsageCache(DualCache()), clock=_frozen_clock
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
precise_minute = FROZEN_INSTANT.strftime("%Y-%m-%d-%H-%M")
scope_ids = [_api_key, user_id, team_id, end_user_id]
for scope_id in scope_ids:
@ -94,8 +118,8 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
await parallel_request_handler.async_log_success_event(
kwargs=kwargs,
response_obj=response_obj,
start_time=datetime.now(),
end_time=datetime.now(),
start_time=FROZEN_INSTANT,
end_time=FROZEN_INSTANT,
)
for scope_id in scope_ids:
@ -107,3 +131,162 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
f"expected 50 tokens counted for {scope_id}, "
f"got {current['current_tpm']}"
)
@pytest.mark.asyncio
async def test_a_pre_call_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read()
)
session: Final = UserAPIKeyAuth(api_key="sk-torn-pre", max_parallel_requests=5)
api_key: Final = session.api_key
await handler.async_pre_call_hook(
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None
) == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1}
for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY):
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None
) is None
@pytest.mark.asyncio
async def test_a_success_event_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read()
)
api_key: Final = hash_token("sk-torn-success")
await internal_usage_cache.async_set_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count",
value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1},
litellm_parent_otel_span=None,
)
await handler.async_log_success_event(
kwargs={
"litellm_params": {
"metadata": {"user_api_key": api_key, "user_api_key_model_max_budget": {}}
}
},
response_obj=ModelResponse(usage=Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7)),
start_time=LAST_MICROSECOND_OF_JANUARY,
end_time=FIRST_MICROSECOND_OF_FEBRUARY,
)
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None
) == {"current_requests": 0, "current_tpm": 7, "current_rpm": 1}
for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY):
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None
) is None
@pytest.mark.asyncio
async def test_a_failure_event_across_a_minute_rollover_lands_in_the_bucket_of_its_first_clock_read():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read()
)
api_key: Final = hash_token("sk-torn-failure")
await internal_usage_cache.async_set_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count",
value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1},
litellm_parent_otel_span=None,
)
await handler.async_log_failure_event(
kwargs={
"litellm_params": {"metadata": {"user_api_key": api_key}},
"exception": Exception("upstream boom"),
},
response_obj=None,
start_time=LAST_MICROSECOND_OF_JANUARY,
end_time=FIRST_MICROSECOND_OF_FEBRUARY,
)
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count", litellm_parent_otel_span=None
) == {"current_requests": 0, "current_tpm": 0, "current_rpm": 1}
for torn_minute in (FIRST_MINUTE_OF_FEBRUARY, TORN_MINUTE_OF_JANUARY):
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::{torn_minute}::request_count", litellm_parent_otel_span=None
) is None
@pytest.mark.asyncio
async def test_a_post_call_headers_read_across_a_minute_rollover_uses_the_bucket_of_its_first_clock_read():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache, clock=_clock_rolling_over_after_first_read()
)
user_api_key_dict: Final = UserAPIKeyAuth(api_key="sk-torn-post", rpm_limit=5, tpm_limit=100)
api_key: Final = user_api_key_dict.api_key
await internal_usage_cache.async_set_cache(
key=f"{api_key}::{LAST_MINUTE_OF_JANUARY}::request_count",
value={"current_requests": 1, "current_tpm": 10, "current_rpm": 1},
litellm_parent_otel_span=None,
)
response: Final = ModelResponse()
response._hidden_params = {}
await handler.async_post_call_success_hook(
data={"model": "gpt-4o-mini"},
user_api_key_dict=user_api_key_dict,
response=response,
)
assert response._hidden_params["additional_headers"] == {
"x-ratelimit-remaining-requests": 4,
"x-ratelimit-limit-requests": 5,
"x-ratelimit-remaining-tokens": 90,
"x-ratelimit-limit-tokens": 100,
}
@pytest.mark.asyncio
async def test_a_request_in_one_minute_is_not_counted_by_a_pre_call_in_the_next_minute():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache,
clock=_clock_reading(iter([datetime(2026, 3, 10, 12, 0, 0), datetime(2026, 3, 10, 12, 1, 0)])),
)
session: Final = UserAPIKeyAuth(api_key="sk-minute-reset", rpm_limit=1)
api_key: Final = session.api_key
await handler.async_pre_call_hook(
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
await handler.async_pre_call_hook(
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
assert await internal_usage_cache.async_get_cache(
key=f"{api_key}::2026-03-10-12-01::request_count", litellm_parent_otel_span=None
) == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1}
@pytest.mark.asyncio
async def test_retry_after_is_the_seconds_until_the_next_minute_of_the_injected_clock():
internal_usage_cache: Final = InternalUsageCache(DualCache())
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=internal_usage_cache,
clock=_clock_reading(itertools.repeat(datetime(2026, 3, 10, 12, 0, 45, 500000))),
)
session: Final = UserAPIKeyAuth(api_key="sk-retry-after", rpm_limit=1)
await handler.async_pre_call_hook(
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
with pytest.raises(ProxyRateLimitError) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
)
assert exc_info.value.headers["retry-after"] == "14.5"

View file

@ -33,6 +33,7 @@ fallback path (unknown model, missing model) for every limiter.
"""
import sys
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -305,7 +306,9 @@ async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit(
Trip the per-key RPM cap and assert the raised exception carries
``model`` / ``llm_provider`` resolved from ``data["model"]``.
"""
handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
handler = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0)
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-test",
max_parallel_requests=10,
@ -404,7 +407,9 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back():
When ``data["model"]`` is unparseable, the resolver falls back to
``litellm_proxy`` — and crucially does *not* leak a secondary exception.
"""
handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
handler = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0)
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-unknown",
max_parallel_requests=10,
@ -438,7 +443,9 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back():
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_missing_model_falls_back():
handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
handler = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()), clock=lambda: datetime(2026, 1, 31, 12, 0, 0)
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-no-model",
max_parallel_requests=10,

View file

@ -8,7 +8,8 @@ middleware is a transparent pass-through when no recorder is injected.
import asyncio
import threading
from typing import List, Optional, Tuple
from datetime import datetime, timezone
from typing import Final, List, Optional, Tuple
import pytest
from starlette.applications import Starlette
@ -28,6 +29,7 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import (
from litellm.proxy.middleware.in_flight_requests_middleware import (
InFlightRequestsMiddleware,
)
from litellm.types.proxy.gateway_requests import GatewayRequestCounts, GatewayRequestKey
class FakeRecorder:
@ -505,14 +507,18 @@ def test_varying_model_ids_fold_into_a_single_persisted_key():
that is. The SGR key is persisted, so it must not carry that dimension: a
caller who could vary it could mint an unbounded number of table rows.
"""
accumulator = GatewayRequestAccumulator()
frozen_now: Final = datetime(2026, 3, 14, 12, 0, tzinfo=timezone.utc)
accumulator = GatewayRequestAccumulator(clock=lambda: frozen_now)
for model_id in ("deploy-1", "deploy-2", "deploy-3"):
client = TestClient(_make_sink_app(None, accumulator, status_code=200, model_id=model_id))
client.post("/v1/chat/completions")
snapshot = accumulator.drain()
assert len(snapshot) == 1
assert next(iter(snapshot.values())).successful_requests == 3
snapshot: Final = accumulator.drain()
assert snapshot == {
GatewayRequestKey(date=frozen_now.strftime("%Y-%m-%d"), category="llm", route="/chat/completions"): (
GatewayRequestCounts(successful_requests=3, failed_requests=0)
)
}
@pytest.mark.parametrize("status_code", [400, 429, 500, 503])

View file

@ -1851,11 +1851,46 @@ def test_ProxyConfig_load_credential_list_returns_items():
dumped = creds[0].model_dump()
assert dumped == {
"credential_name": "openai-key",
"display_name": None,
"credential_info": {"provider": "openai"},
"credential_values": {"api_key": "sk-x"},
}
def test_ProxyConfig_load_credential_list_tags_every_entry_as_config_defined():
creds = ProxyConfig().load_credential_list(
{
"credential_list": [
{"credential_name": "plain", "credential_info": {}, "credential_values": {"api_key": "sk-x"}},
{
"credential_name": "claims-db",
"source": "db",
"credential_info": {},
"credential_values": {"api_key": "sk-y"},
},
]
}
)
assert [(cred.credential_name, cred.source) for cred in creds] == [("plain", "config"), ("claims-db", "config")]
@pytest.mark.parametrize("display_name", [2024, True, "Azure Prod"])
def test_ProxyConfig_load_credential_list_ignores_a_display_name_set_in_config(display_name):
creds = ProxyConfig().load_credential_list(
{
"credential_list": [
{
"credential_name": "azure_cred",
"display_name": display_name,
"credential_info": {},
"credential_values": {"api_key": "sk-x"},
}
]
}
)
assert [(cred.credential_name, cred.display_name) for cred in creds] == [("azure_cred", None)]
def test_ProxyConfig_load_credential_list_invalid_entry_raises():
pc = ProxyConfig()
with pytest.raises(ValidationError):
@ -4021,6 +4056,30 @@ async def test_ProxyConfig_get_credentials_reads_from_writer_not_replica(monkeyp
reader_inner.litellm_credentialstable.find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_ProxyConfig_get_credentials_carries_the_stored_display_name_and_marks_rows_as_db(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
pc = ProxyConfig()
fake_prisma = MagicMock()
fake_prisma.db.litellm_credentialstable.find_many = AsyncMock(
return_value=[{**_encrypted_credential_row("labeled-cred", "sk-labeled"), "display_name": "Prod OpenAI"}]
)
_stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma)
await pc.get_credentials(prisma_client=fake_prisma)
loaded = CredentialAccessor.find_credential("labeled-cred")
assert loaded is not None
assert (loaded.display_name, loaded.source, loaded.credential_values) == (
"Prod OpenAI",
"db",
{"api_key": "sk-labeled"},
)
# ---------------------------------------------------------------------------
# ProxyConfig._reschedule_spend_log_cleanup_job
# ---------------------------------------------------------------------------
@ -5449,14 +5508,24 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f
@pytest.mark.asyncio
@pytest.mark.parametrize("versions", [None, ["2024-11-05"], [], ["2026-07-28"], ["unknown"]])
async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions):
@pytest.mark.parametrize(
("versions", "valid"),
[
(None, True),
(["2024-11-05"], True),
(["2026-07-28"], True),
(["2025-11-25", "2026-07-28"], True),
([], False),
(["unknown"], False),
],
)
async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions, valid):
config = tmp_path / "mcp-versions.yaml"
config.write_text(json.dumps({"model_list": [], "general_settings": {"mcp_advertised_versions": versions}}))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
if versions is None or versions == ["2024-11-05"]:
if valid:
_, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config))
assert settings["mcp_advertised_versions"] == versions
return

View file

@ -3,6 +3,7 @@
import json
import types
from datetime import date, datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -23,6 +24,7 @@ from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
DAY = date(2026, 7, 30)
TODAY = date(2026, 7, 31)
_SCHEDULED_NOW: Final = datetime(2026, 7, 31, 12, 0, tzinfo=timezone.utc)
# The endpoints require ptu_effective_from alongside the count and rate, so a fixture that
@ -1449,10 +1451,10 @@ async def test_scheduled_rollup_backfills_after_pricing_the_day():
"""The catch-up pass runs after the day's own rollup, so it sees yesterday already
priced and does not write it a second time."""
table = _FakeSentinelTable()
yesterday = datetime.now(timezone.utc).date() - timedelta(days=1)
yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1)
prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=2)))], table)
await run_scheduled_ptu_rollup(prisma)
await run_scheduled_ptu_rollup(prisma, clock=lambda: _SCHEDULED_NOW)
yesterday_key = ("t", yesterday.isoformat(), PTU_SENTINEL_API_KEY, "m1")
assert table.upsert_keys.count(yesterday_key) == 1
@ -1475,13 +1477,13 @@ async def test_scheduled_rollup_with_an_explicit_target_date_does_not_backfill()
async def test_scheduled_rollup_holds_one_lock_across_both_phases():
"""Backfill running outside the lock would let another pod's prune race its writes."""
table = _FakeSentinelTable()
yesterday = datetime.now(timezone.utc).date() - timedelta(days=1)
yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1)
prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=3)))], table)
rows_at_release = []
lock = _pod_lock(acquired=True)
lock.release_lock = AsyncMock(side_effect=lambda **kwargs: rows_at_release.append(len(table.rows)))
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=lock)
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=lock, clock=lambda: _SCHEDULED_NOW)
lock.acquire_lock.assert_awaited_once()
assert rows_at_release == [4]
@ -1507,12 +1509,12 @@ async def test_scheduled_rollup_alerts_when_a_backfill_charge_never_landed():
"""An unpriced day that stays unpriced is the silent underbill this work exists to
remove, so it has to reach an operator too."""
table = _FakeSentinelTable()
yesterday = datetime.now(timezone.utc).date() - timedelta(days=1)
yesterday: Final = _SCHEDULED_NOW.date() - timedelta(days=1)
prisma = _prisma_for([_windowed_row(effective_from=_midnight(yesterday - timedelta(days=1)))], table)
prisma.db.litellm_dailyteamspend.upsert = AsyncMock(side_effect=RuntimeError("db down"))
alert = AsyncMock()
await run_scheduled_ptu_rollup(prisma, alert=alert)
await run_scheduled_ptu_rollup(prisma, alert=alert, clock=lambda: _SCHEDULED_NOW)
messages = [call.args[0] for call in alert.await_args_list]
assert any("backfill" in message for message in messages)
@ -1535,23 +1537,41 @@ async def test_a_broken_alert_channel_does_not_fail_the_backfill():
@pytest.mark.asyncio
async def test_scheduled_rollup_with_no_target_date_closes_a_backdated_window():
"""The production call shape from proxy_server.py, on the real clock: no target_date,
"""The production call shape from proxy_server.py, on a frozen clock: no target_date,
a window backdated 30 days, and every elapsed in-window day has to end up priced with
no operator alert raised. Every other rollup test pins target_date, which is exactly
why this regression shipped."""
table = _FakeSentinelTable()
today = datetime.now(timezone.utc).date()
opened_on = today - timedelta(days=30)
today: Final = _SCHEDULED_NOW.date()
opened_on: Final = today - timedelta(days=30)
prisma = _prisma_for([_windowed_row(effective_from=_midnight(opened_on))], table)
alert = AsyncMock()
await run_scheduled_ptu_rollup(prisma, alert=alert)
await run_scheduled_ptu_rollup(prisma, alert=alert, clock=lambda: _SCHEDULED_NOW)
expected = [(opened_on + timedelta(days=offset)).isoformat() for offset in range(30)]
assert _priced_dates(table) == expected
alert.assert_not_awaited()
@pytest.mark.asyncio
async def test_scheduled_rollup_uses_one_day_across_the_midnight_boundary():
table: Final = _FakeSentinelTable()
clock: Final = iter(
(
datetime(2026, 1, 31, 23, 59, 59, 999999, tzinfo=timezone.utc),
datetime(2026, 2, 1, 0, 0, 0, 1, tzinfo=timezone.utc),
)
).__next__
prisma: Final = _prisma_for([_windowed_row(effective_from=_midnight(date(2026, 1, 27)))], table)
result: Final = await run_scheduled_ptu_rollup(prisma, clock=clock)
assert result is not None
assert result.day == date(2026, 1, 30)
assert _priced_dates(table) == ["2026-01-27", "2026-01-28", "2026-01-29", "2026-01-30"]
# --- R8: a rename must not re-price history under the new name ----------------
@ -2066,19 +2086,22 @@ async def test_the_catch_up_pass_reaches_a_config_declared_deployment():
"""The catch-up shares the loader, so config deployments join it without being wired in.
That is what prices the elapsed days of a reservation declared before today."""
table = _FakeSentinelTable()
now = datetime.now(timezone.utc)
started = (now - timedelta(days=3)).strftime("%Y-%m-%dT00:00:00Z")
now: Final = _SCHEDULED_NOW
started: Final = (now - timedelta(days=3)).strftime("%Y-%m-%dT00:00:00Z")
entry = _router_entry(
model_id="cfg-back",
model_info={"ptu_count": 100, "cost_per_ptu_per_hour": 0.02, "team_id": "t", "ptu_effective_from": started},
)
await run_scheduled_ptu_rollup(
_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), router=_router_holding(entry)
_prisma_for([], table),
pod_lock_manager=_pod_lock(acquired=True),
router=_router_holding(entry),
clock=lambda: _SCHEDULED_NOW,
)
charged = sorted(day for (_, day, _, model) in table.rows if model == "cfg-back")
yesterday = (now.date() - timedelta(days=1)).isoformat()
yesterday: Final = (now.date() - timedelta(days=1)).isoformat()
assert len(charged) == 3, charged
assert charged[-1] == yesterday
assert all(row["ptu_flat_cost"] == pytest.approx(48.0) for row in table.rows.values())

View file

@ -387,7 +387,7 @@ def test_change_password_request_passwords_hidden_from_repr():
for rendered in (repr(request), str(request)):
assert "hunter2hunter2" not in rendered
assert "NewP@ssw0rd-2026" not in rendered
@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]])
@pytest.mark.parametrize("versions", [[], ["2099-01-01"]])
def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
from pydantic import ValidationError

View file

@ -6951,19 +6951,13 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
primary_model = "gpt-4"
fallback_model = "gpt-3.5-turbo"
# Freeze the limiter's clock so the per-minute counter key is stable and
# the pre-seeded counter is guaranteed to be the one it reads.
class _FrozenClock(datetime.datetime):
@classmethod
def now(cls, tz=None):
return cls(2026, 1, 1, 12, 30, 0)
precise_minute = "2026-01-01-12-30"
# Real per-key per-model TPM limiter + a key carrying the customer's
# `model_tpm_limit` metadata (only the primary is capped).
limiter = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
internal_usage_cache=InternalUsageCache(DualCache()),
clock=lambda: datetime.datetime(2026, 1, 1, 12, 30, 0),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-lit3890",
@ -7006,30 +7000,27 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
mock_router = MagicMock()
mock_router.fallbacks = [{primary_model: [fallback_model]}]
with patch(
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=real_limiter_pre_call,
):
with patch.object(
processor,
"common_processing_pre_call_logic",
side_effect=real_limiter_pre_call,
):
data, logging_obj = await processor._pre_call_with_fallbacks(
request=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=user_api_key_dict,
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
data, logging_obj = await processor._pre_call_with_fallbacks(
request=MagicMock(),
general_settings={},
proxy_logging_obj=MagicMock(),
user_api_key_dict=user_api_key_dict,
version=None,
proxy_config=MagicMock(),
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
model=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
# The capped primary tripped the real limiter, and the fallback (which
# has no per-model cap) served the request — no 429 to the client.
@ -7038,19 +7029,16 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
# Sanity-check the premise: the limiter genuinely raises a
# ProxyRateLimitError for the capped primary under the frozen clock.
with patch(
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
):
with pytest.raises(ProxyRateLimitError):
await limiter.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
},
call_type="acompletion",
)
with pytest.raises(ProxyRateLimitError):
await limiter.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
},
call_type="acompletion",
)
@staticmethod
def _v3_limiter_rig(

View file

@ -1487,11 +1487,13 @@ class TestCredentialsRepository:
repo._prisma_client.db.litellm_credentialstable._records["my-key"] = {
"credential_id": "cred-1",
"credential_name": "my-key",
"display_name": "My Key",
"credential_values": {"api_key": "encrypted_secret"},
"credential_info": {"provider": "openai"},
}
cred = await repo.find_by_name("my-key")
assert isinstance(cred, CredentialItem)
assert cred.display_name == "My Key"
assert cred.credential_values == {"api_key": "encrypted_secret"}
assert cred.credential_info == {"provider": "openai"}

View file

@ -43,7 +43,7 @@ def classify(category: str, changed: list[str]) -> str:
DOCS = ["README.md", "docs/my_website/index.mdx", "litellm/anywhere.md"]
CLIENT = ["ui/litellm-dashboard/src/App.tsx"]
BACKEND = ["litellm/main.py"]
CI = [".github/workflows/test-litellm-ui-unit.yml"]
CI = [".github/workflows/test-unit.yml"]
@pytest.mark.parametrize(

View file

@ -231,5 +231,5 @@ def test_ui_category_fails_open_when_the_api_fails(tmp_path: Path) -> None:
def test_ui_category_runs_when_the_ui_workflows_themselves_change(tmp_path: Path) -> None:
"""Without this the dashboard jobs would skip on the pull request that edits
them, shipping a workflow change nothing ever exercised."""
decision, _ = _run(tmp_path, files=[".github/workflows/test-litellm-ui-unit.yml"], category="ui")
decision, _ = _run(tmp_path, files=[".github/workflows/test-unit.yml"], category="ui")
assert decision == "decision=run"

View file

@ -1037,3 +1037,64 @@ def test_update_model_cost():
assert litellm.model_cost["gpt-4"]["input_cost_per_token"] == 0.00002
except Exception as e:
pytest.fail(f"An error occurred: {e}")
def test_register_model_scopes_builtin_match_to_the_given_provider():
"""Registering under an id that collides with another provider's catalog
key keeps the entry provider-less under the id instead of merging into the
baseten row."""
colliding_id: Final = "baseten/zai-org/glm-5.2"
builtin_key: Final = "baseten/zai-org/GLM-5.2"
model_cost_entries: Final = _snapshot_model_cost_entries((colliding_id, builtin_key))
builtin_row_before: Final = copy.deepcopy(litellm.model_cost[builtin_key])
try:
litellm.register_model(
{
colliding_id: {
"input_cost_per_token": 0.00000096,
"output_cost_per_token": 0.00000302,
"cache_read_input_token_cost": 0.00000010,
"mode": "chat",
}
},
custom_llm_provider="openai",
)
entry: Final = litellm.model_cost[colliding_id]
assert "litellm_provider" not in entry
assert entry["input_cost_per_token"] == 0.00000096
assert entry["output_cost_per_token"] == 0.00000302
assert entry["cache_read_input_token_cost"] == 0.00000010
assert litellm.model_cost[builtin_key] == builtin_row_before
finally:
_restore_model_cost_entries(model_cost_entries)
def test_per_request_custom_pricing_scopes_a_colliding_deployment_id_to_its_provider():
from litellm.main import _register_custom_pricing_for_request
deployment_id: Final = "baseten/zai-org/glm-5.2"
builtin_key: Final = "baseten/zai-org/GLM-5.2"
shared_key: Final = "openai/zai-org/GLM-5.2"
model_cost_entries: Final = _snapshot_model_cost_entries((deployment_id, builtin_key, shared_key))
builtin_row_before: Final = copy.deepcopy(litellm.model_cost[builtin_key])
openai_models_before: Final = frozenset(litellm.open_ai_chat_completion_models)
try:
_register_custom_pricing_for_request(
model="zai-org/GLM-5.2",
custom_llm_provider="openai",
kwargs={
"input_cost_per_token": 0.00000096,
"output_cost_per_token": 0.00000302,
"metadata": {"model_info": {"id": deployment_id}},
},
model_info={"mode": "chat"},
)
entry: Final = litellm.model_cost[deployment_id]
assert entry["input_cost_per_token"] == 0.00000096
assert entry["output_cost_per_token"] == 0.00000302
assert litellm.model_cost[builtin_key] == builtin_row_before
finally:
_restore_model_cost_entries(model_cost_entries)
litellm.open_ai_chat_completion_models.intersection_update(openai_models_before)

Some files were not shown because too many files have changed in this diff Show more