mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
chore: merge main into litellm_remove_lit002_dict_ban
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
d14859c4c7
135 changed files with 8726 additions and 2396 deletions
|
|
@ -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
|
||||
|
|
|
|||
15
.github/merge-smoke-tests.json
vendored
15
.github/merge-smoke-tests.json
vendored
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
8
.github/scripts/assert_ci_coverage.py
vendored
8
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
140
.github/scripts/run_merge_smoke.py
vendored
140
.github/scripts/run_merge_smoke.py
vendored
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
283
.github/workflows/_test-unit-base.yml
vendored
283
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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
|
||||
134
.github/workflows/check-ui-api-types.yml
vendored
134
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -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."
|
||||
48
.github/workflows/ci-coverage.yml
vendored
48
.github/workflows/ci-coverage.yml
vendored
|
|
@ -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
|
||||
|
|
@ -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: |
|
||||
|
|
|
|||
397
.github/workflows/required-checks-legacy.yml
vendored
Normal file
397
.github/workflows/required-checks-legacy.yml
vendored
Normal 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
|
||||
223
.github/workflows/test-code-quality.yml
vendored
223
.github/workflows/test-code-quality.yml
vendored
|
|
@ -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
|
||||
440
.github/workflows/test-linting.yml
vendored
440
.github/workflows/test-linting.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
105
.github/workflows/test-litellm-ui-lint.yml
vendored
105
.github/workflows/test-litellm-ui-lint.yml
vendored
|
|
@ -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
|
||||
99
.github/workflows/test-litellm-ui-unit.yml
vendored
99
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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
|
||||
24
.github/workflows/test-merge-smoke.yml
vendored
24
.github/workflows/test-merge-smoke.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
103
.github/workflows/test-unit-documentation.yml
vendored
103
.github/workflows/test-unit-documentation.yml
vendored
|
|
@ -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
|
||||
992
.github/workflows/test-unit.yml
vendored
992
.github/workflows/test-unit.yml
vendored
File diff suppressed because it is too large
Load diff
7
Makefile
7
Makefile
|
|
@ -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
|
||||
|
|
|
|||
72
codecov.yaml
72
codecov.yaml
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_CredentialsTable" ADD COLUMN IF NOT EXISTS "display_name" TEXT;
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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] = {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
262
litellm/proxy/_experimental/mcp_server/interactions.py
Normal file
262
litellm/proxy/_experimental/mcp_server/interactions.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2078,6 +2078,7 @@ async def test_model_connection(
|
|||
"responses",
|
||||
"anthropic_messages",
|
||||
"ocr",
|
||||
"evaluation",
|
||||
]
|
||||
| None = fastapi.Body(
|
||||
None,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)")
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
149
tests/integration/_support/pdf_document.py
Normal file
149
tests/integration/_support/pdf_document.py
Normal 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,
|
||||
}
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
129
tests/integration/management/test_credential_display_name.py
Normal file
129
tests/integration/management/test_credential_display_name.py
Normal 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"
|
||||
395
tests/integration/mcp/test_interactions.py
Normal file
395
tests/integration/mcp/test_interactions.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
]
|
||||
159
tests/integration/providers/_count_tokens_system_lift.py
Normal file
159
tests/integration/providers/_count_tokens_system_lift.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"}},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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?"}],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
253
tests/unit/proxy/_experimental/mcp_server/test_interactions.py
Normal file
253
tests/unit/proxy/_experimental/mcp_server/test_interactions.py
Normal 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()
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue