diff --git a/.circleci/config.yml b/.circleci/config.yml index e9c295aca08..0988d824241 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -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 diff --git a/.github/merge-smoke-tests.json b/.github/merge-smoke-tests.json deleted file mode 100644 index 90d3b6a6d59..00000000000 --- a/.github/merge-smoke-tests.json +++ /dev/null @@ -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" - } -} diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index e90589d3c95..04c801cf5cc 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -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", ) diff --git a/.github/scripts/run_merge_smoke.py b/.github/scripts/run_merge_smoke.py index 7e74de324fc..81d47d5b45c 100644 --- a/.github/scripts/run_merge_smoke.py +++ b/.github/scripts/run_merge_smoke.py @@ -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) diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml deleted file mode 100644 index e5847e5fbea..00000000000 --- a/.github/workflows/_test-unit-base.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/check-ui-api-types.yml b/.github/workflows/check-ui-api-types.yml deleted file mode 100644 index 185c20d916d..00000000000 --- a/.github/workflows/check-ui-api-types.yml +++ /dev/null @@ -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." diff --git a/.github/workflows/ci-coverage.yml b/.github/workflows/ci-coverage.yml deleted file mode 100644 index cd36a9a7ed6..00000000000 --- a/.github/workflows/ci-coverage.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/publish-lint-base-counts.yml b/.github/workflows/publish-lint-base-counts.yml index 7af00e9c842..5623ab82b90 100644 --- a/.github/workflows/publish-lint-base-counts.yml +++ b/.github/workflows/publish-lint-base-counts.yml @@ -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: | diff --git a/.github/workflows/required-checks-legacy.yml b/.github/workflows/required-checks-legacy.yml new file mode 100644 index 00000000000..6bba690bd8a --- /dev/null +++ b/.github/workflows/required-checks-legacy.yml @@ -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 diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml deleted file mode 100644 index 8a3a9107b51..00000000000 --- a/.github/workflows/test-code-quality.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index db2fc969ace..95a32a41ad6 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -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 diff --git a/.github/workflows/test-litellm-ui-lint.yml b/.github/workflows/test-litellm-ui-lint.yml deleted file mode 100644 index 9ea5100e21b..00000000000 --- a/.github/workflows/test-litellm-ui-lint.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml deleted file mode 100644 index fcd61cedd50..00000000000 --- a/.github/workflows/test-litellm-ui-unit.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/test-merge-smoke.yml b/.github/workflows/test-merge-smoke.yml index 910763c6af2..2472f2831b6 100644 --- a/.github/workflows/test-merge-smoke.yml +++ b/.github/workflows/test-merge-smoke.yml @@ -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 diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml deleted file mode 100644 index b042e182802..00000000000 --- a/.github/workflows/test-unit-documentation.yml +++ /dev/null @@ -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 diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 21b82db44b7..6dee2bb605f 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -1,4 +1,4 @@ -name: "Unit Tests" +name: unit on: pull_request: @@ -17,26 +17,15 @@ concurrency: group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }} cancel-in-progress: ${{ github.event_name == 'pull_request' }} -# One caller for every tests/test_litellm shard, replacing the nine thin workflow -# files that each wrapped a single call to _test-unit-base.yml. Adding a shard is -# now one matrix entry rather than a new file. -# -# `name` is the shard id, and each check reports as " / Run tests". -# Unit shard names are not required ruleset contexts, so matrix entries can be split freely. -# -# Every entry states its timeouts even when they equal the base workflow's -# defaults. An absent matrix key renders as an empty string, which is not a -# number, so a partially-specified entry would fail the call rather than fall -# back to the default. jobs: rust-bridge: - name: Build the Rust bridge - outputs: - artifact: ${{ steps.rust-bridge-artifact.outputs.name }} - runs-on: ubuntu-latest + name: rust-bridge permissions: contents: read pull-requests: read + outputs: + artifact: ${{ steps.rust-bridge-artifact.outputs.name }} + runs-on: ubuntu-latest env: UV_PYTHON: "3.12" steps: @@ -119,32 +108,46 @@ jobs: RUN_ID: ${{ github.run_id }} RUN_ATTEMPT: ${{ github.run_attempt }} run: echo "name=rust-bridge-${RUN_ID}-${RUN_ATTEMPT}" >> "$GITHUB_OUTPUT" - - unit: - name: ${{ matrix.shard }} - needs: rust-bridge - if: ${{ !cancelled() }} + assert-shard-coverage: permissions: contents: read - id-token: write - pull-requests: write + runs-on: ubuntu-latest + timeout-minutes: 2 + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + - name: Assert every test directory and file is claimed by a shard + run: python3 .github/scripts/assert_ci_coverage.py --shards + unit: + name: ${{ matrix.shard }} + needs: [rust-bridge, assert-shard-coverage] + if: ${{ !cancelled() }} + runs-on: ubuntu-latest + timeout-minutes: ${{ matrix.job-timeout-minutes }} + permissions: + contents: read + pull-requests: read + env: + UV_PYTHON: ${{ matrix.python-version }} + LITELLM_LOCAL_MODEL_COST_MAP: "True" strategy: fail-fast: false matrix: include: - shard: core-utils - artifact-name: core-utils - test-path: >- + test-path: |- tests/unit/decisions tests/unit/litellm_core_utils + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: enterprise-routing - artifact-name: enterprise-routing - test-path: >- + test-path: |- tests/unit/google_genai tests/unit/router_strategy tests/unit/router_utils @@ -160,69 +163,256 @@ jobs: tests/unit/enterprise/proxy/test_file_deletion_blocking.py tests/unit/enterprise/proxy/test_managed_files_access_check.py tests/unit/enterprise/proxy/test_managed_files_hook.py + python-version: "3.12" workers: 4 reruns: 2 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: integrations - artifact-name: integrations - test-path: >- + test-path: |- tests/test_litellm/integrations tests/test_litellm/tracing tests/unit/integrations + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: Vertex AI - artifact-name: llm-vertex-ai + - shard: llms-vertex-ai test-path: >- tests/unit/llms/vertex_ai + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: All Other Providers - artifact-name: llm-other-providers + - shard: llms-anthropic test-path: >- - tests/unit/llms - --ignore=tests/unit/llms/vertex_ai - --ignore=tests/unit/llms/openai - --ignore=tests/unit/llms/meta - --ignore=tests/unit/llms/base_llm/batches/base_batches_config_test.py + tests/unit/llms/anthropic + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: OpenAI and Meta Providers - artifact-name: llm-openai-meta + - shard: llms-bedrock + test-path: |- + tests/unit/llms/bedrock + tests/unit/llms/bedrock_mantle + python-version: "3.12" + workers: 4 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope + + - shard: llms-a-to-g + test-path: |- + tests/unit/llms/test_cache_control_and_reasoning.py + tests/unit/llms/test_custom_llm.py + tests/unit/llms/test_file_content_block.py + tests/unit/llms/test_file_search_responses.py + tests/unit/llms/test_lifecycle_fix.py + tests/unit/llms/test_oss_decision.py + tests/unit/llms/test_polling_url_origin_match.py + tests/unit/llms/test_predibase_transformation.py + tests/unit/llms/a2a + tests/unit/llms/aiml + tests/unit/llms/aiohttp_openai + tests/unit/llms/apiserpent + tests/unit/llms/aws_polly + tests/unit/llms/azure + tests/unit/llms/azure_ai + tests/unit/llms/base_llm + tests/unit/llms/baseten + tests/unit/llms/black_forest_labs + tests/unit/llms/bytez + tests/unit/llms/cerebras + tests/unit/llms/chat + tests/unit/llms/chatgpt + tests/unit/llms/clarifai + tests/unit/llms/claude_code + tests/unit/llms/cloudflare + tests/unit/llms/codex + tests/unit/llms/cohere + tests/unit/llms/cometapi + tests/unit/llms/compactifai + tests/unit/llms/crusoe + tests/unit/llms/custom_httpx + tests/unit/llms/dashscope + tests/unit/llms/databricks + tests/unit/llms/dataforseo + tests/unit/llms/datarobot + tests/unit/llms/deepagents + tests/unit/llms/deepgram + tests/unit/llms/deepinfra + tests/unit/llms/deepseek + tests/unit/llms/docker_model_runner + tests/unit/llms/duckduckgo + tests/unit/llms/e2b + tests/unit/llms/edenai + tests/unit/llms/elevenlabs + tests/unit/llms/exa_ai + tests/unit/llms/fal_ai + tests/unit/llms/fastcrw + tests/unit/llms/featherless_ai + tests/unit/llms/firecrawl + tests/unit/llms/fireworks_ai + tests/unit/llms/gdc + tests/unit/llms/gemini + tests/unit/llms/gigachat + tests/unit/llms/github_copilot + tests/unit/llms/gradient_ai + tests/unit/llms/groq + --ignore=tests/unit/llms/base_llm/batches/base_batches_config_test.py + python-version: "3.12" + workers: 4 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope + + - shard: llms-h-to-z + test-path: |- + tests/unit/llms/heroku + tests/unit/llms/hosted_vllm + tests/unit/llms/huggingface + tests/unit/llms/hyperbolic + tests/unit/llms/inception + tests/unit/llms/infinity + tests/unit/llms/jina_ai + tests/unit/llms/lambda_ai + tests/unit/llms/langflow + tests/unit/llms/langgraph + tests/unit/llms/laya + tests/unit/llms/lemonade + tests/unit/llms/linkup + tests/unit/llms/litellm_proxy + tests/unit/llms/llamafile + tests/unit/llms/lm_studio + tests/unit/llms/manus + tests/unit/llms/meta_llama + tests/unit/llms/minimax + tests/unit/llms/mistral + tests/unit/llms/modelscope + tests/unit/llms/mongodb + tests/unit/llms/moonshot + tests/unit/llms/nadir + tests/unit/llms/nebius + tests/unit/llms/neosantara + tests/unit/llms/nimble + tests/unit/llms/novita + tests/unit/llms/nscale + tests/unit/llms/nvidia_nim + tests/unit/llms/nvidia_riva + tests/unit/llms/oci + tests/unit/llms/ocr + tests/unit/llms/ollama + tests/unit/llms/oobabooga + tests/unit/llms/openai_like + tests/unit/llms/opencode + tests/unit/llms/openrouter + tests/unit/llms/ovhcloud + tests/unit/llms/parallel_ai + tests/unit/llms/parasail + tests/unit/llms/pass_through + tests/unit/llms/perplexity + tests/unit/llms/pg_vector + tests/unit/llms/publicai + tests/unit/llms/ragflow + tests/unit/llms/recraft + tests/unit/llms/reducto + tests/unit/llms/replicate + tests/unit/llms/runwayml + tests/unit/llms/s3_vectors + tests/unit/llms/sagemaker + tests/unit/llms/sail + tests/unit/llms/sambanova + tests/unit/llms/sap + tests/unit/llms/scaleway + tests/unit/llms/searchapi + tests/unit/llms/searxng + tests/unit/llms/serper + tests/unit/llms/snowflake + tests/unit/llms/soniox + tests/unit/llms/stability + tests/unit/llms/tavily + tests/unit/llms/tencent + tests/unit/llms/tinyfish + tests/unit/llms/together_ai + tests/unit/llms/tool_loop + tests/unit/llms/triton + tests/unit/llms/v0 + tests/unit/llms/valkey + tests/unit/llms/vercel_ai_gateway + tests/unit/llms/volcengine + tests/unit/llms/voyage + tests/unit/llms/wandb + tests/unit/llms/watsonx + tests/unit/llms/xai + tests/unit/llms/xinference + tests/unit/llms/you_com + tests/unit/llms/zai + python-version: "3.12" + workers: 4 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope + + - shard: llms-openai-meta test-path: >- tests/unit/llms/openai tests/unit/llms/meta + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: misc - artifact-name: misc + - shard: root test-path: >- tests/test_litellm/test_*.py tests/unit/test_*.py + python-version: "3.12" workers: 4 reruns: 2 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: misc-dirs - artifact-name: misc-dirs + - shard: router test-path: >- tests/unit/test_router + python-version: "3.12" + workers: 4 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope + + - shard: router-unit-tests + test-path: >- + tests/router_unit_tests + python-version: "3.12" + workers: 0 + reruns: 0 + timeout-minutes: 10 + job-timeout-minutes: 60 + dist: loadscope + + - shard: endpoints + test-path: >- tests/unit/a2a_protocol + tests/unit/anthropic_interface tests/unit/batches tests/unit/chat_completions tests/unit/completion_extras @@ -230,25 +420,58 @@ jobs: tests/unit/embeddings tests/unit/endpoints tests/unit/files - tests/unit/harness tests/unit/images tests/unit/interactions tests/unit/messages + tests/unit/ocr + tests/unit/passthrough tests/unit/rag + tests/unit/realtime_api tests/unit/rerank_api - tests/unit/rust_bridge - tests/unit/secret_managers tests/unit/vector_stores tests/unit/videos - --ignore=tests/unit/rust_bridge/native_route_wheel_test.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope + + - shard: rust-bridge-harness + test-path: >- + tests/unit/compression + tests/unit/harness + tests/unit/rust_bridge + tests/unit/sandbox + --ignore=tests/unit/rust_bridge/native_route_wheel_test.py + python-version: "3.12" + workers: 4 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope + + - shard: enterprise-repositories-secrets + test-path: >- + tests/proxy_behavior/lens/test_connection.py + tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py + tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py + tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py + tests/unit/integration_support + tests/unit/models + tests/unit/repositories + tests/unit/secret_managers + tests/unit/skills/test_skills_main.py + tests/unit/tracing + python-version: "3.12" + workers: 2 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + dist: loadscope - shard: proxy-auth - artifact-name: proxy-auth - test-path: >- + test-path: |- tests/unit/proxy/auth --ignore=tests/unit/proxy/auth/test_auth_checks.py --ignore=tests/unit/proxy/auth/test_user_api_key_auth.py @@ -257,27 +480,29 @@ jobs: --ignore=tests/unit/proxy/auth/test_models_fallback_endpoint.py --ignore=tests/unit/proxy/auth/test_multipart_bypass_repro.py --ignore=tests/unit/proxy/auth/test_proxy_routes.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-hooks-client - artifact-name: proxy-hooks-client - test-path: >- + test-path: |- tests/unit/proxy/hooks tests/unit/proxy/policy_engine tests/unit/proxy/client --ignore=tests/unit/proxy/hooks/test_banned_keyword_list.py --ignore=tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-endpoints - artifact-name: proxy-endpoints - test-path: >- + test-path: |- tests/unit/proxy/management_endpoints tests/unit/proxy/management_helpers tests/unit/proxy/list_api @@ -292,14 +517,15 @@ jobs: --ignore=tests/unit/proxy/management_endpoints/test_key_generate_prisma.py --ignore=tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py --ignore=tests/unit/proxy/management_helpers/test_audit_logs_proxy.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-feature-endpoints - artifact-name: proxy-feature-endpoints - test-path: >- + test-path: |- tests/unit/proxy/guardrails --ignore=tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py --ignore=tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py @@ -328,32 +554,43 @@ jobs: tests/unit/proxy/ui_crud_endpoints tests/unit/proxy/config_resolvers tests/unit/proxy/utils + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-server - artifact-name: proxy-server - test-path: "tests/unit/proxy/proxy_server" + test-path: >- + tests/unit/proxy/proxy_server + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 60 job-timeout-minutes: 100 + dist: loadscope - shard: mcp-elicitation - artifact-name: mcp-elicitation - test-path: >- + test-path: |- + --override-ini=pythonpath=tests + tests/unit/experimental_mcp_client/test_mcp_client.py + tests/unit/proxy/_experimental/mcp_server/test_capabilities.py + tests/unit/proxy/_experimental/mcp_server/test_interactions.py tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py + tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py + tests/unit/proxy/_experimental/mcp_server/test_operations.py + tests/integration/mcp/test_interactions.py + python-version: "3.12" workers: 2 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-infra - artifact-name: proxy-infra - test-path: >- + test-path: |- tests/unit/proxy/db --ignore=tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py --ignore=tests/unit/proxy/db/test_update_daily_tag_spend.py @@ -362,8 +599,6 @@ jobs: tests/unit/proxy/spend_tracking --ignore=tests/unit/proxy/spend_tracking/test_search_api_logging.py tests/unit/proxy/pass_through_endpoints - tests/unit/proxy/_experimental - --ignore=tests/unit/proxy/_experimental/mcp_server tests/unit/proxy/experimental tests/unit/proxy/common_utils --ignore=tests/unit/proxy/common_utils/test_cache_aware_routing.py @@ -378,14 +613,15 @@ jobs: tests/unit/proxy/management tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py tests/unit/proxy/roi_calculator + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-infra-root - artifact-name: proxy-infra-root - test-path: >- + test-path: |- tests/unit/proxy/test_*.py --ignore=tests/unit/proxy/test_aproxy_startup.py --ignore=tests/unit/proxy/test_credential_slot_registry.py @@ -412,32 +648,35 @@ jobs: --ignore=tests/unit/proxy/test_unit_test_proxy_hooks.py --ignore=tests/unit/proxy/test_update_spend.py --ignore=tests/unit/proxy/test_zero_cost_model_budget_bypass.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: caching-local - artifact-name: caching-local test-path: >- tests/unit/caching + python-version: "3.12" workers: 2 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: proxy-extras - artifact-name: proxy-extras test-path: >- tests/unit/litellm_proxy_extras + python-version: "3.12" workers: 2 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: enterprise-package - artifact-name: enterprise-package - test-path: >- + test-path: |- tests/unit/enterprise/integrations tests/unit/enterprise/proxy/auth tests/unit/enterprise/proxy/guardrails @@ -445,15 +684,17 @@ jobs: tests/unit/enterprise/proxy/test_audit_logging_endpoints.py tests/unit/enterprise/proxy/test_liteadmin.py tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - shard: enterprise-managed-files - artifact-name: enterprise-managed-files test-path: >- tests/unit/enterprise/proxy/hooks + python-version: "3.12" workers: 4 reruns: 0 timeout-minutes: 20 @@ -461,181 +702,148 @@ jobs: dist: load - shard: responses-caching-types - artifact-name: responses-caching-types - test-path: >- - tests/unit/responses + test-path: |- + tests/unit/responses/litellm_completion_transformation + tests/unit/responses/test_additional_tools.py + tests/unit/responses/test_custom_tool_call.py + tests/unit/responses/test_dispatch.py + tests/unit/responses/test_metadata_codex_callback.py + tests/unit/responses/test_no_duplicate_spend_logs.py + tests/unit/responses/test_null_test_fix.py + tests/unit/responses/test_responses_api_bridge_flag.py + tests/unit/responses/test_responses_api_lifecycle.py + tests/unit/responses/test_responses_api_request_body.py + tests/unit/responses/test_responses_prompt_management.py + tests/unit/responses/test_responses_router_cooldown.py + tests/unit/responses/test_responses_streaming_iterator.py + tests/unit/responses/test_responses_supported_endpoints_passthrough.py + tests/unit/responses/test_responses_utils.py + tests/unit/responses/test_responses_websocket_all_providers.py + tests/unit/responses/test_rust_bridge_websocket.py + tests/unit/responses/test_sse_output_recovery.py + tests/unit/responses/test_streaming_iterator.py + tests/unit/responses/test_streaming_iterator_error_events.py + tests/unit/responses/test_text_format_conversion.py tests/unit/types - --ignore=tests/unit/responses/mcp + python-version: "3.12" workers: 2 reruns: 0 timeout-minutes: 20 job-timeout-minutes: 60 + dist: loadscope - - shard: unit - artifact-name: unit - test-path: >- - tests/unit/anthropic_interface - tests/unit/compression - tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py - tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py - tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py - tests/unit/integration_support - tests/unit/models - tests/unit/ocr - tests/unit/passthrough - tests/unit/realtime_api - tests/unit/repositories - tests/unit/sandbox - tests/unit/skills/test_skills_main.py - tests/unit/tracing - tests/proxy_behavior/lens/test_connection.py - workers: 2 - reruns: 0 - timeout-minutes: 20 - job-timeout-minutes: 60 - uses: ./.github/workflows/_test-unit-base.yml - with: - rust-bridge-artifact: ${{ needs.rust-bridge.outputs.artifact }} - test-path: ${{ matrix.test-path }} - workers: ${{ matrix.workers }} - reruns: ${{ matrix.reruns }} - timeout-minutes: ${{ matrix.timeout-minutes }} - job-timeout-minutes: ${{ matrix.job-timeout-minutes }} - dist: ${{ matrix.dist || 'loadscope' }} - artifact-name: ${{ matrix.artifact-name }} - - lens-python-310: - name: Lens Python 3.10 - permissions: - contents: read - id-token: write - pull-requests: write - uses: ./.github/workflows/_test-unit-base.yml - with: - python-version: "3.10" - test-path: tests/unit/proxy/lens/test_inference.py - workers: 0 - reruns: 0 - timeout-minutes: 5 - artifact-name: lens-python-310 - - # Fast guard — fails the workflow when a test directory or file inside a sharded - # tree is claimed by no shard. The semantic-shard design has no catch-all bucket, - # so an unassigned child runs nowhere; assert_ci_coverage.py holds the tree list - # and reads the same test-path keys the coverage census does. - assert-shard-coverage: - runs-on: ubuntu-latest - timeout-minutes: 2 - permissions: - contents: read - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - name: Assert every test directory and file is claimed by a shard - run: python3 .github/scripts/assert_ci_coverage.py --shards - - proxy-db: - needs: assert-shard-coverage - # Display only the semantic shard name in the checks UI instead of GHA's - # default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)" - # which includes every matrix field and gets truncated past the test-path. - name: ${{ matrix.test-group }} - permissions: - contents: read - id-token: write - pull-requests: write - strategy: - fail-fast: false - matrix: - include: - # Must run serially — event-loop conflict with the logging worker. - - test-group: key-generation + - shard: key-generation test-path: >- tests/unit/proxy/management_endpoints/test_key_generate_prisma.py + python-version: "3.12" workers: 0 + reruns: 2 + timeout-minutes: 20 + job-timeout-minutes: 60 dist: loadscope - timeout: 20 - # ---- auth: split into 2 shards ---- - - test-group: auth-checks - test-path: >- + - shard: auth-checks + test-path: |- tests/unit/proxy/auth/test_auth_checks.py tests/unit/proxy/auth/test_user_api_key_auth.py tests/unit/proxy/test_credential_slot_registry.py tests/unit/proxy/test_deprecated_key_grace_period.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: jwt-and-keys - test-path: >- + + - shard: jwt-and-keys + test-path: |- tests/unit/proxy/auth/test_jwt.py tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py tests/unit/proxy/test_proxy_custom_auth.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - # ---- test_proxy_utils.py, single shard, worksteal distribution ---- - - test-group: proxy-utils + - shard: proxy-utils test-path: >- tests/unit/proxy/test_proxy_utils.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: worksteal - timeout: 15 - # ---- proxy server: split into 2 shards ---- - - test-group: proxy-server-core - test-path: >- + - shard: proxy-server-core + test-path: |- tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py tests/unit/proxy/test__lazy_features.py tests/unit/proxy/test_aproxy_startup.py tests/unit/proxy/test_proxy_server.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: proxy-runtime - test-path: >- + + - shard: proxy-runtime + test-path: |- tests/unit/proxy/auth/test_multipart_bypass_repro.py tests/unit/proxy/auth/test_proxy_routes.py tests/unit/proxy/middleware/test_request_size_limit_middleware.py tests/unit/proxy/test_proxy_config_unit_test.py tests/unit/proxy/test_proxy_token_counter.py tests/unit/proxy/test_server_root_path.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: mcp-oauth - test-path: >- + - shard: mcp-oauth + test-path: |- tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py tests/unit/proxy/_experimental/mcp_server/outbound_credentials + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - # ---- logging: split into 2 shards ---- - - test-group: custom-logging - test-path: >- + - shard: custom-logging + test-path: |- tests/proxy_unit_tests/test_proxy_custom_logger.py tests/unit/proxy/test_custom_callback_input.py tests/unit/proxy/test_custom_logger_s3_gcs.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: logging-misc - test-path: >- + + - shard: logging-misc + test-path: |- tests/unit/proxy/management_helpers/test_audit_logs_proxy.py tests/unit/proxy/spend_tracking/test_search_api_logging.py tests/unit/proxy/test_proxy_reject_logging.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: db-and-spend - test-path: >- + - shard: db-and-spend + test-path: |- tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py tests/unit/proxy/db/test_update_daily_tag_spend.py @@ -643,30 +851,50 @@ jobs: tests/unit/proxy/test_prisma_client_backoff_retry.py tests/unit/proxy/test_update_spend.py tests/unit/skills/test_skills_db.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - # ---- guardrails + budget + hooks: split into 2 ---- - - test-group: guardrails-hooks - test-path: >- + - shard: guardrails-hooks + test-path: |- tests/unit/proxy/hooks/test_banned_keyword_list.py tests/unit/proxy/test_proxy_setting_guardrails.py tests/unit/proxy/test_unit_test_proxy_hooks.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: budgets + + - shard: guardrails-tests test-path: >- + tests/guardrails_tests + --deselect=tests/guardrails_tests/test_custom_guardrail.py::test_get_guardrail_dynamic_request_body_params + python-version: "3.12" + workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 + dist: loadscope + + - shard: budgets + test-path: |- tests/unit/proxy/auth/test_default_end_user_budget_simple.py tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py tests/unit/proxy/test_zero_cost_model_budget_bypass.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - - test-group: endpoints-and-responses - test-path: >- + - shard: endpoints-and-responses + test-path: |- tests/proxy_unit_tests/test_proxy_exception_mapping.py tests/unit/proxy/lens tests/unit/proxy/auth/test_models_fallback_endpoint.py @@ -685,25 +913,379 @@ jobs: tests/unit/proxy/test_reducto_ocr_route.py tests/unit/proxy/test_response_polling_pre_call_checks.py tests/unit/proxy/test_ui_path_detection.py + python-version: "3.12" workers: 4 + reruns: 2 + timeout-minutes: 15 + job-timeout-minutes: 60 dist: loadscope - timeout: 15 - uses: ./.github/workflows/_test-unit-base.yml - with: - test-path: ${{ matrix.test-path }} - workers: ${{ matrix.workers }} - reruns: 2 - timeout-minutes: ${{ matrix.timeout }} - dist: ${{ matrix.dist }} - artifact-name: proxy-db-${{ matrix.test-group }} + - shard: lens-python-310 + test-path: >- + tests/unit/proxy/lens/test_inference.py + python-version: "3.10" + workers: 0 + reruns: 0 + timeout-minutes: 5 + job-timeout-minutes: 60 + dist: loadscope + 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: ${{ needs.rust-bridge.outputs.artifact }} + + - name: Install dependencies + if: steps.changes.outputs.decision != 'skip' + timeout-minutes: 8 + env: + RUST_BRIDGE_ARTIFACT: ${{ needs.rust-bridge.outputs.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: ${{ matrix.timeout-minutes }} + env: + TEST_PATH: ${{ matrix.test-path }} + MAX_FAILURES: "10" + WORKERS: ${{ matrix.workers }} + RERUNS: ${{ matrix.reruns }} + TEST_TIMEOUT_SECONDS: "120" + DIST: ${{ matrix.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-${{ matrix.shard }}-${{ github.run_id }}-${{ github.run_attempt }} + path: coverage.xml + retention-days: 1 + ui-unit: + name: ui-unit + permissions: + contents: read + pull-requests: read + 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 + docs: + name: docs + permissions: + contents: read + pull-requests: read + runs-on: ubuntu-latest + timeout-minutes: 10 + 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 + + - 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 + coverage: + name: coverage + needs: unit + if: always() + runs-on: ubuntu-latest + timeout-minutes: 10 + 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-*-${{ github.run_id }}-* + path: coverage-reports + merge-multiple: false + + - 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: unit + 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: unit + fail_ci_if_error: false unit-passed: name: unit passed - needs: [rust-bridge, unit, lens-python-310, assert-shard-coverage, proxy-db] + permissions: {} + needs: [rust-bridge, assert-shard-coverage, unit, ui-unit, docs] if: always() runs-on: ubuntu-latest timeout-minutes: 2 - permissions: {} steps: - name: Require every unit job to succeed env: diff --git a/Makefile b/Makefile index b8614ca4eec..da6cebedede 100644 --- a/Makefile +++ b/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 diff --git a/codecov.yaml b/codecov.yaml index 4d93c18f3ac..ae72f6d5286 100644 --- a/codecov.yaml +++ b/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: diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261009000000_add_credential_display_name/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261009000000_add_credential_display_name/migration.sql new file mode 100644 index 00000000000..264bd88a807 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261009000000_add_credential_display_name/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_CredentialsTable" ADD COLUMN IF NOT EXISTS "display_name" TEXT; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 038dfdeaca5..59ffb037177 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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") diff --git a/litellm/constants.py b/litellm/constants.py index 74273b9ecb9..699872f32cc 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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, diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 960b9dde53f..729a6acb20d 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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) diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py index 6e73387326a..f9e54c30e89 100644 --- a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -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) diff --git a/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py index 2169dfcad39..cb4bbf00c85 100644 --- a/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py +++ b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py @@ -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] = { diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 151830f1589..7ac3e9c80c5 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -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( diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index ef53c30efa3..fd19ae3341b 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -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} diff --git a/litellm/main.py b/litellm/main.py index b02ccca9ecb..5c2e4d2b2a1 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, ) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7930a823e07..15ab7f160b0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/models/credentials.py b/litellm/models/credentials.py index b91ced275ff..730778a323b 100644 --- a/litellm/models/credentials.py +++ b/litellm/models/credentials.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/capabilities.py b/litellm/proxy/_experimental/mcp_server/capabilities.py index bfd00327eb4..1ee263633be 100644 --- a/litellm/proxy/_experimental/mcp_server/capabilities.py +++ b/litellm/proxy/_experimental/mcp_server/capabilities.py @@ -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, ...]: diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index f3800c93a77..bf0fae1a210 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/interactions.py b/litellm/proxy/_experimental/mcp_server/interactions.py new file mode 100644 index 00000000000..42addcfb65b --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/interactions.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index 13c3d265429..17bb84e066a 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8185ca508c0..be63e2c169c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 6ec5bdf6af7..039fbfdb916 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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): diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 0c0300fa0e5..55cbc0d6498 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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": diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 92578aa43b9..00a05307a89 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bed2503d211..d46dd3ab3e9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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, diff --git a/litellm/proxy/client/cli/commands/credentials.py b/litellm/proxy/client/cli/commands/credentials.py index 2c4080dbeb2..2028bdb5599 100644 --- a/litellm/proxy/client/cli/commands/credentials.py +++ b/litellm/proxy/client/cli/commands/credentials.py @@ -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) diff --git a/litellm/proxy/client/credentials.py b/litellm/proxy/client/credentials.py index d9edecd2eb7..ebcd79457a6 100644 --- a/litellm/proxy/client/credentials.py +++ b/litellm/proxy/client/credentials.py @@ -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, diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 6d25307db96..8df22f458cf 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -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: diff --git a/litellm/proxy/db/gateway_request_tracking.py b/litellm/proxy/db/gateway_request_tracking.py index b236acdb75d..ec8b1103333 100644 --- a/litellm/proxy/db/gateway_request_tracking.py +++ b/litellm/proxy/db/gateway_request_tracking.py @@ -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: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 4180ca1a656..557ff2680b4 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -2078,6 +2078,7 @@ async def test_model_connection( "responses", "anthropic_messages", "ocr", + "evaluation", ] | None = fastapi.Body( None, diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 0f884387673..7c02ac308f3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8627565321c..a5af4dc3618 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 038dfdeaca5..59ffb037177 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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") diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index b2bf6f46b3a..e4370fe304d 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -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 diff --git a/litellm/repositories/credentials_repository.py b/litellm/repositories/credentials_repository.py index ddb9767b2b9..fa6aeca534d 100644 --- a/litellm/repositories/credentials_repository.py +++ b/litellm/repositories/credentials_repository.py @@ -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 {}, } diff --git a/litellm/router.py b/litellm/router.py index c8640356f7e..5e93025b479 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 556cb6712ae..e104f74faf7 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -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, diff --git a/litellm/utils.py b/litellm/utils.py index a2992bda8b6..d97ee5f2988 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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"] diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 7930a823e07..15ab7f160b0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/schema.prisma b/schema.prisma index 038dfdeaca5..59ffb037177 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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") diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index 1c793705dc6..edb69f24e29 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -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 diff --git a/tests/code_coverage_tests/check_workflow_startup_safety.py b/tests/code_coverage_tests/check_workflow_startup_safety.py index cf150daef4c..935ab7b3bef 100644 --- a/tests/code_coverage_tests/check_workflow_startup_safety.py +++ b/tests/code_coverage_tests/check_workflow_startup_safety.py @@ -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.*?)\}\}", re.DOTALL) QUOTED: Final = re.compile(r"'[^']*'") ARITHMETIC: Final = re.compile(r"[+*]") -MATRIX_REF: Final = re.compile(r"^\$\{\{\s*matrix\.(?P[\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", "") + 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)") diff --git a/tests/code_coverage_tests/test_e2e_metadata.py b/tests/code_coverage_tests/test_e2e_metadata.py index 57d04262969..22fe84f72a5 100644 --- a/tests/code_coverage_tests/test_e2e_metadata.py +++ b/tests/code_coverage_tests/test_e2e_metadata.py @@ -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. diff --git a/tests/code_coverage_tests/test_merge_smoke.py b/tests/code_coverage_tests/test_merge_smoke.py index e3b96e222b9..550817910c0 100644 --- a/tests/code_coverage_tests/test_merge_smoke.py +++ b/tests/code_coverage_tests/test_merge_smoke.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)) diff --git a/tests/code_coverage_tests/test_unit_passed_gate.py b/tests/code_coverage_tests/test_unit_passed_gate.py index 7b511538d40..5a0d04681d4 100644 --- a/tests/code_coverage_tests/test_unit_passed_gate.py +++ b/tests/code_coverage_tests/test_unit_passed_gate.py @@ -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"] == {} diff --git a/tests/e2e/ui/tests/usage/usageActivityTabs.spec.ts b/tests/e2e/ui/tests/usage/usageActivityTabs.spec.ts index 2ee5ae3e392..522025e22f3 100644 --- a/tests/e2e/ui/tests/usage/usageActivityTabs.spec.ts +++ b/tests/e2e/ui/tests/usage/usageActivityTabs.spec.ts @@ -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 " $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 { await navigateToPage(page, Page.NewUsage); diff --git a/tests/e2e_harness/AGENTS.md b/tests/e2e_harness/AGENTS.md index 92882850b67..212e46d754c 100644 --- a/tests/e2e_harness/AGENTS.md +++ b/tests/e2e_harness/AGENTS.md @@ -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` diff --git a/tests/integration/_support/pdf_document.py b/tests/integration/_support/pdf_document.py new file mode 100644 index 00000000000..5c6ad82b85f --- /dev/null +++ b/tests/integration/_support/pdf_document.py @@ -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, + } diff --git a/tests/integration/database/test_request_log_indexes_at_boot.py b/tests/integration/database/test_request_log_indexes_at_boot.py index fe35ec2f1f9..e72477de901 100644 --- a/tests/integration/database/test_request_log_indexes_at_boot.py +++ b/tests/integration/database/test_request_log_indexes_at_boot.py @@ -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" diff --git a/tests/integration/management/test_credential_display_name.py b/tests/integration/management/test_credential_display_name.py new file mode 100644 index 00000000000..0e72dcf8b2f --- /dev/null +++ b/tests/integration/management/test_credential_display_name.py @@ -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" diff --git a/tests/integration/mcp/test_interactions.py b/tests/integration/mcp/test_interactions.py new file mode 100644 index 00000000000..42af8cf5dba --- /dev/null +++ b/tests/integration/mcp/test_interactions.py @@ -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 diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py new file mode 100644 index 00000000000..f196018f117 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py @@ -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) diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py new file mode 100644 index 00000000000..41d5f3d8685 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py @@ -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, + ] diff --git a/tests/integration/providers/_count_tokens_system_lift.py b/tests/integration/providers/_count_tokens_system_lift.py new file mode 100644 index 00000000000..56619d1c963 --- /dev/null +++ b/tests/integration/providers/_count_tokens_system_lift.py @@ -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), + ) diff --git a/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py b/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py new file mode 100644 index 00000000000..13091ccd14a --- /dev/null +++ b/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py @@ -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 diff --git a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py index f9d4f7f91bb..933fc49c149 100644 --- a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py @@ -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]) diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py index db6629a43d2..07b02480fc8 100644 --- a/tests/integration/providers/test_decisions_wire.py +++ b/tests/integration/providers/test_decisions_wire.py @@ -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)) diff --git a/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py b/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py new file mode 100644 index 00000000000..9fad73f8080 --- /dev/null +++ b/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py @@ -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 diff --git a/tests/integration/routing/test_usage_based_routing_redis_reads.py b/tests/integration/routing/test_usage_based_routing_redis_reads.py index f4801eb3318..f3ee7855395 100644 --- a/tests/integration/routing/test_usage_based_routing_redis_reads.py +++ b/tests/integration/routing/test_usage_based_routing_redis_reads.py @@ -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() diff --git a/tests/integration/run.py b/tests/integration/run.py index f7b2dead197..aa5ead066f2 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -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) diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 3596d1d343c..3f7e03a09a0 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -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, diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 9a30e1990b1..ff192c9279c 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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) diff --git a/tests/unit/integrations/focus/test_mavvrik_destination.py b/tests/unit/integrations/focus/test_mavvrik_destination.py index 1be72b4409b..13f462c4ef3 100644 --- a/tests/unit/integrations/focus/test_mavvrik_destination.py +++ b/tests/unit/integrations/focus/test_mavvrik_destination.py @@ -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.""" diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py index d1e23a17747..9e3704a9df0 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py @@ -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"}}, + ) diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index 163897c2609..a646135efc3 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -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() diff --git a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py index 2a31ac75d1d..9af7497bc82 100644 --- a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py +++ b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py @@ -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?"}], + } diff --git a/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py b/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py index f104e655704..f8e3655b2f2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py @@ -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" + ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_interactions.py b/tests/unit/proxy/_experimental/mcp_server/test_interactions.py new file mode 100644 index 00000000000..8286269525a --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_interactions.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py index 1a592aa1c9a..a270ea16a39 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py @@ -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, ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index c7b477a6cc6..b67e4d24f47 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 97c6e2703ad..49d21924d9c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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() diff --git a/tests/unit/proxy/client/cli/test_credentials_commands.py b/tests/unit/proxy/client/cli/test_credentials_commands.py index fb9d749dd02..fe9bc333f37 100644 --- a/tests/unit/proxy/client/cli/test_credentials_commands.py +++ b/tests/unit/proxy/client/cli/test_credentials_commands.py @@ -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() diff --git a/tests/unit/proxy/client/test_credentials.py b/tests/unit/proxy/client/test_credentials.py index 666c5dac2b0..9de990da487 100644 --- a/tests/unit/proxy/client/test_credentials.py +++ b/tests/unit/proxy/client/test_credentials.py @@ -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""" diff --git a/tests/unit/proxy/credential_endpoints/test_endpoints.py b/tests/unit/proxy/credential_endpoints/test_endpoints.py index aff7babe424..83d6167256a 100644 --- a/tests/unit/proxy/credential_endpoints/test_endpoints.py +++ b/tests/unit/proxy/credential_endpoints/test_endpoints.py @@ -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" diff --git a/tests/unit/proxy/db/test_gateway_request_tracking.py b/tests/unit/proxy/db/test_gateway_request_tracking.py index a6689b38039..5566a5cb9bc 100644 --- a/tests/unit/proxy/db/test_gateway_request_tracking.py +++ b/tests/unit/proxy/db/test_gateway_request_tracking.py @@ -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) ) } diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index 266fd06333c..5aa763fcc9c 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -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) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index 1772f41a9d3..a2e43b3bc9c 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -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" diff --git a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py index 16b5406bc21..187100c24d8 100644 --- a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -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, diff --git a/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py b/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py index 333906884a8..f2515905707 100644 --- a/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py +++ b/tests/unit/proxy/middleware/test_billable_request_metrics_middleware.py @@ -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]) diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index d1b80627b3a..c8c89662583 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -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 diff --git a/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index 8f25cffecf5..decebcaa2bc 100644 --- a/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/unit/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -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()) diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index dd7bd86982e..50a4c3a6ca7 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -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 diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 1aed60ee5e2..cc7cce937d7 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -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( diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index f2e033c9141..cfea6f880fd 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -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"} diff --git a/tests/unit/test_circleci_path_filter.py b/tests/unit/test_circleci_path_filter.py index 873258d26a2..2a4084bd8c2 100644 --- a/tests/unit/test_circleci_path_filter.py +++ b/tests/unit/test_circleci_path_filter.py @@ -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( diff --git a/tests/unit/test_detect_changes.py b/tests/unit/test_detect_changes.py index d8feb2371a3..44a4bfdf448 100644 --- a/tests/unit/test_detect_changes.py +++ b/tests/unit/test_detect_changes.py @@ -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" diff --git a/tests/unit/test_register_model_custom_pricing.py b/tests/unit/test_register_model_custom_pricing.py index 288ac70d81d..02c90e0d97b 100644 --- a/tests/unit/test_register_model_custom_pricing.py +++ b/tests/unit/test_register_model_custom_pricing.py @@ -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) diff --git a/tests/unit/test_router_model_cost_isolation.py b/tests/unit/test_router_model_cost_isolation.py index 4aed84d05c8..ecc1f3ab078 100644 --- a/tests/unit/test_router_model_cost_isolation.py +++ b/tests/unit/test_router_model_cost_isolation.py @@ -26,7 +26,13 @@ from litellm.litellm_core_utils.ptu_pricing import ptu_config_error from litellm.litellm_core_utils.llm_cost_calc.utils import SERVICE_TIER_COST_KEY_SUFFIXES from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS -from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo +from litellm.types.router import ( + Deployment, + DeploymentTypedDict, + LiteLLM_Params, + LiteLLMParamsTypedDict, + ModelInfo, +) from litellm.utils import ( _invalidate_model_cost_lowercase_map, reapply_runtime_model_cost_registrations, @@ -3227,3 +3233,89 @@ def test_price_data_reload_refreshes_the_cached_model_group_and_deployment_info( assert router.cached_model_group_info("grp").input_cost_per_token == new_price assert router.cached_deployment_model_info("dep-a", "openai/gpt-4o")["input_cost_per_token"] == new_price + + +_COLLIDING_BUILTIN_KEY: Final = "baseten/zai-org/GLM-5.2" +_COLLIDING_SHARED_OPENAI_KEY: Final = "openai/zai-org/GLM-5.2" +_COLLIDING_MODEL_INFO: Final = { + "input_cost_per_token": 0.00000096, + "output_cost_per_token": 0.00000302, + "cache_read_input_token_cost": 0.00000010, + "mode": "chat", +} + + +def _colliding_id_deployment(model_id: str, custom_llm_provider: str) -> DeploymentTypedDict: + litellm_params: Final = LiteLLMParamsTypedDict( + model="zai-org/GLM-5.2", + api_base="https://inference.baseten.co/v1", + custom_llm_provider=custom_llm_provider, + ) + model_info: Final[dict] = {"id": model_id, **_COLLIDING_MODEL_INFO} + return DeploymentTypedDict( + model_name="nvidia/zai-org/glm-5.2", + litellm_params=litellm_params, + model_info=model_info, + ) + + +@pytest.mark.parametrize("model_id", ("baseten/zai-org/glm-5.2",)) +def test_deployment_id_colliding_with_another_providers_catalog_key_keeps_its_own_pricing(model_id: str) -> None: + """A deployment id equal to another provider's catalog key must not merge + into that row: the merge leaves litellm_provider=baseten on the entry, which + _check_provider_match then rejects for the openai request, billing $0.""" + builtin_row_before: Final = copy.deepcopy(litellm.model_cost[_COLLIDING_BUILTIN_KEY]) + model_cost_entries: Final = { + key: copy.deepcopy(litellm.model_cost.get(key)) + for key in (model_id, _COLLIDING_BUILTIN_KEY, _COLLIDING_SHARED_OPENAI_KEY) + } + try: + router: Final = Router(model_list=[_colliding_id_deployment(model_id, "openai")]) + + response: Final = router.completion( + model="nvidia/zai-org/glm-5.2", + messages=[{"role": "user", "content": "colliding id pricing"}], + mock_response=litellm.ModelResponse( + model="zai-org/GLM-5.2", + usage=litellm.Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500), + ), + ) + + assert isinstance(response, litellm.ModelResponse) + assert response._hidden_params["response_cost"] == pytest.approx( + 1000 * 0.00000096 + 500 * 0.00000302 + ) + assert litellm.model_cost[_COLLIDING_BUILTIN_KEY] == builtin_row_before + finally: + _restore_model_cost_entries(model_cost_entries) + litellm.get_model_info.cache_clear() + + +def test_deployment_id_matching_its_own_providers_catalog_key_still_merges() -> None: + """Same-provider collision keeps today's merge behavior: the deployment's + custom prices land on the baseten row the openai-compatible baseten request + matches against.""" + model_id: Final = "baseten/zai-org/glm-5.2" + model_cost_entries: Final = { + key: copy.deepcopy(litellm.model_cost.get(key)) + for key in (model_id, _COLLIDING_BUILTIN_KEY) + } + try: + router: Final = Router(model_list=[_colliding_id_deployment(model_id, "baseten")]) + + response: Final = router.completion( + model="nvidia/zai-org/glm-5.2", + messages=[{"role": "user", "content": "same provider pricing"}], + mock_response=litellm.ModelResponse( + model="zai-org/GLM-5.2", + usage=litellm.Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500), + ), + ) + + assert isinstance(response, litellm.ModelResponse) + assert response._hidden_params["response_cost"] == pytest.approx( + 1000 * 0.00000096 + 500 * 0.00000302 + ) + finally: + _restore_model_cost_entries(model_cost_entries) + litellm.get_model_info.cache_clear() diff --git a/tests/unit/test_select_ui_test_scope.py b/tests/unit/test_select_ui_test_scope.py index ebfb7701e6a..4605aab956c 100644 --- a/tests/unit/test_select_ui_test_scope.py +++ b/tests/unit/test_select_ui_test_scope.py @@ -1,6 +1,6 @@ """Regression tests for the UI unit-test scope decision. -`.github/workflows/test-litellm-ui-unit.yml` narrows the dashboard's Vitest run to +`.github/workflows/test-unit.yml` narrows the dashboard's Vitest run to `vitest related ` so a pull request only pays for the tests it can affect. `related` resolves a file to the tests that import it, so a file no test imports resolves to nothing, and with `--passWithNoTests` the job then goes green @@ -26,7 +26,7 @@ import yaml REPO_ROOT = Path(__file__).resolve().parents[2] SCOPE_SCRIPT = REPO_ROOT / ".github" / "scripts" / "select_ui_test_scope.sh" -WORKFLOW = REPO_ROOT / ".github" / "workflows" / "test-litellm-ui-unit.yml" +WORKFLOW = REPO_ROOT / ".github" / "workflows" / "test-unit.yml" STEP_NAME = "Run UI unit tests (Vitest)" FULL_SUITE_ARGV = ["run", "test", "--", "--run", "--pool", "forks", "--maxWorkers=14"] @@ -75,7 +75,7 @@ def test_an_empty_change_set_fails_open_to_the_full_suite() -> None: def _step_script() -> str: workflow = yaml.safe_load(WORKFLOW.read_text()) - steps = workflow["jobs"]["ui-unit-tests"]["steps"] + steps = workflow["jobs"]["ui-unit"]["steps"] script = next(step["run"] for step in steps if step.get("name") == STEP_NAME) resolved = script.replace("${{ github.repository }}", "BerriAI/litellm") assert "${{" not in resolved, "the step uses an Actions expression this harness does not resolve" diff --git a/tests/unit/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index 4fa9c5bd3c1..725551e19d3 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -9,7 +9,7 @@ import pytest import yaml _REPO_ROOT: Final = Path(__file__).resolve().parents[2] -_BASE_WORKFLOW: Final = _REPO_ROOT / ".github" / "workflows" / "_test-unit-base.yml" +_UNIT_WORKFLOW: Final = _REPO_ROOT / ".github" / "workflows" / "test-unit.yml" _SHARD_ENV: Final = MappingProxyType( {"MAX_FAILURES": "10", "RERUNS": "0", "DIST": "loadscope", "TEST_TIMEOUT_SECONDS": "60", "COVERAGE_CORE": "sysmon"} ) @@ -19,8 +19,8 @@ _FAILING_TEST: Final = "def test_fails():\n assert False\n" def _run_tests_script() -> str: - workflow: Final = yaml.safe_load(_BASE_WORKFLOW.read_text()) - return next(step["run"] for step in workflow["jobs"]["run"]["steps"] if step.get("name") == "Run tests") + workflow: Final = yaml.safe_load(_UNIT_WORKFLOW.read_text()) + return next(step["run"] for step in workflow["jobs"]["unit"]["steps"] if step.get("name") == "Run tests") def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.CompletedProcess[str]: diff --git a/tests/unit/test_unit_shard_per_test_timeout.py b/tests/unit/test_unit_shard_per_test_timeout.py index 8096124ca8c..830b46b723d 100644 --- a/tests/unit/test_unit_shard_per_test_timeout.py +++ b/tests/unit/test_unit_shard_per_test_timeout.py @@ -10,7 +10,7 @@ import pytest import yaml _REPO_ROOT: Final = Path(__file__).resolve().parents[2] -_BASE_WORKFLOW: Final = _REPO_ROOT / ".github" / "workflows" / "_test-unit-base.yml" +_UNIT_WORKFLOW: Final = _REPO_ROOT / ".github" / "workflows" / "test-unit.yml" _SHARD_ENV: Final = MappingProxyType({"WORKERS": "2", "RERUNS": "2", "DIST": "loadscope", "TEST_TIMEOUT_SECONDS": "1"}) _HANG_GUARD_FLAGS: Final = frozenset(("-n", "--dist", "--reruns", "--reruns-delay", "--timeout", "--rerun-except")) _HUNG_TEST_MODULE: Final = """ @@ -39,8 +39,8 @@ def test_passes(): def _run_tests_script() -> str: - workflow: Final = yaml.safe_load(_BASE_WORKFLOW.read_text()) - return next(step["run"] for step in workflow["jobs"]["run"]["steps"] if step.get("name") == "Run tests") + workflow: Final = yaml.safe_load(_UNIT_WORKFLOW.read_text()) + return next(step["run"] for step in workflow["jobs"]["unit"]["steps"] if step.get("name") == "Run tests") def _pytest_invocations(script: str) -> tuple[tuple[str, ...], ...]: diff --git a/ui/litellm-dashboard/AGENTS.md b/ui/litellm-dashboard/AGENTS.md index c890c33ce9b..bfba65b2860 100644 --- a/ui/litellm-dashboard/AGENTS.md +++ b/ui/litellm-dashboard/AGENTS.md @@ -4,7 +4,7 @@ Never put LiteLLM tokens or API keys in `localStorage`. `localStorage` survives When you fix lint violations that are grandfathered in `eslint-suppressions.json`, run `eslint . --prune-suppressions` and commit the updated baseline so the gate ratchets down instead of leaving a stale suppression -`src/lib/http/schema.d.ts` is generated from the proxy's OpenAPI spec; never hand-edit it. After changing a backend route or response model that the dashboard consumes, run `npm run gen:api` and commit the result (CI `Check UI API Types Sync` enforces this) +`src/lib/http/schema.d.ts` is generated from the proxy's OpenAPI spec; never hand-edit it. After changing a backend route or response model that the dashboard consumes, run `npm run gen:api` and commit the result (CI `lint / ui-api-types` enforces this) Tests come in three tiers, named by the standard definitions. `Foo.test.tsx` is a unit test: one module, collaborators replaced by doubles, no multi-component tree, and it should run in milliseconds. `Foo.integration.test.tsx` renders a real component tree with real children and only stubs the network boundary; it costs seconds per case, so it earns its place by proving wiring that a unit test cannot reach. Browser-level tests live in `tests/e2e/ui/` as Playwright specs against a live proxy diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.test.ts index ee903628f08..7732ddc8fb3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.test.ts @@ -143,6 +143,16 @@ describe("useCredentials", () => { expect(credentialListCall).not.toHaveBeenCalled(); }); + it("does not call the API when the caller disables the query", async () => { + (credentialListCall as any).mockResolvedValue(mockCredentialsResponse); + + const { result } = renderHook(() => useCredentials({ enabled: false }), { wrapper }); + + expect(result.current.isFetched).toBe(false); + expect(result.current.data).toBeUndefined(); + expect(credentialListCall).not.toHaveBeenCalled(); + }); + it("should return empty credentials array when API returns empty data", async () => { // Mock API returning empty credentials array (credentialListCall as any).mockResolvedValue({ credentials: [] }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts index bdbf4445514..cd2ec717ffb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts @@ -5,11 +5,11 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export const credentialsKeys = createQueryKeys("credentials"); -export const useCredentials = () => { +export const useCredentials = ({ enabled = true }: { enabled?: boolean } = {}) => { const { accessToken } = useAuthorized(); return useQuery({ queryKey: credentialsKeys.list({}), queryFn: async () => await credentialListCall(accessToken!), - enabled: Boolean(accessToken), + enabled: enabled && Boolean(accessToken), }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx index 0c93bb234a5..09e00db9fb0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx @@ -82,6 +82,14 @@ vi.mock("../../hooks/models/useModelCostMap", () => ({ })); const mockTeams = [{ team_id: "team-1", team_alias: "Engineering" }]; +const mockCredentials = [ + { credential_name: "openai-prod", display_name: "Prod OpenAI", credential_values: {}, credential_info: {} }, +]; +const mockUseCredentials = vi.hoisted(() => vi.fn()); +vi.mock("../../hooks/credentials/useCredentials", () => ({ + useCredentials: mockUseCredentials, +})); + vi.mock("../../hooks/teams/useTeams", () => ({ useTeams: () => ({ data: mockTeams, isLoading: false, error: null, refetch: vi.fn() }), })); @@ -152,6 +160,35 @@ describe("AllModelsTab", () => { modelsInfoCalls.length = 0; setModelsInfo([makeRow()]); vi.spyOn(useAuthorizedModule, "default").mockReturnValue(MOCK_AUTHORIZED); + mockUseCredentials.mockImplementation(({ enabled = true }: { enabled?: boolean } = {}) => ({ + data: enabled ? { credentials: mockCredentials } : undefined, + isLoading: false, + })); + }); + + describe("credential labels", () => { + const rowWithCredential = () => { + const row = makeRow(); + return { ...row, litellm_params: { ...row.litellm_params, litellm_credential_name: "openai-prod" } }; + }; + + it("shows the credential's display name for a proxy admin", async () => { + setModelsInfo([rowWithCredential()]); + renderWithProviders(); + + expect(await screen.findByText("Prod OpenAI")).toBeInTheDocument(); + expect(mockUseCredentials).toHaveBeenLastCalledWith({ enabled: true }); + }); + + it("skips the admin-only credential list for other roles and shows the raw name", async () => { + vi.spyOn(useAuthorizedModule, "default").mockReturnValue({ ...MOCK_AUTHORIZED, userRole: "Internal User" }); + setModelsInfo([rowWithCredential()]); + renderWithProviders(); + + expect(await screen.findByText("openai-prod")).toBeInTheDocument(); + expect(screen.queryByText("Prod OpenAI")).not.toBeInTheDocument(); + expect(mockUseCredentials).toHaveBeenLastCalledWith({ enabled: false }); + }); }); it("renders the fetched models and the server row count", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx index 2217bca0fa0..ea502396072 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx @@ -2,11 +2,14 @@ import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentials"; +import { credentialLabelsByName } from "@/components/shared/credentialOptions"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal"; import { ModelData } from "@/components/model_dashboard/types"; import { toast } from "@/lib/toast"; +import { isProxyAdminRole } from "@/utils/roles"; import { uiHref } from "@/utils/uiHref"; import { modelDeleteCall, modelPatchUpdateCall } from "@/components/networking"; import { useQueryClient } from "@tanstack/react-query"; @@ -80,6 +83,11 @@ const AllModelsTab = ({ const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); const { accessToken, userId, userRole, isViewOnly } = useAuthorized(); const { data: teams, isLoading: isLoadingTeams } = useTeams(); + const { data: credentialsResponse } = useCredentials({ enabled: isProxyAdminRole(userRole ?? "") }); + const credentialLabels = useMemo( + () => credentialLabelsByName(credentialsResponse?.credentials ?? []), + [credentialsResponse], + ); const queryClient = useQueryClient(); const [tableState, setTableState] = useQueryStates(TABLE_STATE); @@ -324,6 +332,7 @@ const AllModelsTab = ({ onDeleteClick={handleDeleteClick} onTogglePauseClick={handleTogglePause} pausingModelId={pausingModelId} + credentialLabels={credentialLabels} /> {modelViewMode === "current_team" && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx index be6130b0288..adbeb70a9c5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx @@ -157,6 +157,19 @@ describe("AllModelsTable", () => { expect(screen.getByText("Manual")).toBeInTheDocument(); }); + it("renders the credential's display name when one is set and keeps the name as the tooltip", () => { + render( + , + ); + expect(screen.getByText("Prod OpenAI")).toBeInTheDocument(); + expect(screen.queryByText("openai-prod")).not.toBeInTheDocument(); + expect(screen.getByTitle("openai-prod")).toBeInTheDocument(); + }); + it("shows 'Defined in config' for a config model and the creator for a DB model", () => { const { rerender } = render(); expect(screen.getByText("alice")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx index 2a52bdfb46e..e9eca183a7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.tsx @@ -80,6 +80,7 @@ interface AllModelsTableProps { onDeleteClick: (modelId: string) => void; onTogglePauseClick: (modelId: string, blocked: boolean) => void | Promise; pausingModelId: string | null; + credentialLabels?: ReadonlyMap; } function EmptyState() { @@ -128,6 +129,7 @@ export function AllModelsTable({ onDeleteClick, onTogglePauseClick, pausingModelId, + credentialLabels, }: AllModelsTableProps) { const [filtersOpen, setFiltersOpen] = useState(false); @@ -141,9 +143,20 @@ export function AllModelsTable({ onDeleteClick, onTogglePauseClick, pausingModelId, + credentialLabels, }; return getModelsTableColumns(columnDeps); - }, [userRole, userID, isViewOnly, onModelIdClick, onTeamIdClick, onDeleteClick, onTogglePauseClick, pausingModelId]); + }, [ + userRole, + userID, + isViewOnly, + onModelIdClick, + onTeamIdClick, + onDeleteClick, + onTogglePauseClick, + pausingModelId, + credentialLabels, + ]); const modelGroupOptions = useMemo( () => [ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx index 9581d3db198..c277f527f86 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx @@ -160,7 +160,7 @@ function CredentialsHeader() { ); } -function CredentialsCell({ credentialName }: { credentialName: string | undefined }) { +function CredentialsCell({ credentialName, label }: { credentialName: string | undefined; label: string | undefined }) { if (!credentialName) { return ( @@ -173,7 +173,7 @@ function CredentialsCell({ credentialName }: { credentialName: string | undefine return ( - {credentialName} + {label ?? credentialName} ); } @@ -364,6 +364,7 @@ export interface ModelsTableColumnDeps { onDeleteClick?: (modelId: string) => void; onTogglePauseClick?: (modelId: string, blocked: boolean) => void | Promise; pausingModelId?: string | null; + credentialLabels?: ReadonlyMap; } export const getModelsTableColumns = ({ @@ -375,6 +376,7 @@ export const getModelsTableColumns = ({ onDeleteClick, onTogglePauseClick, pausingModelId, + credentialLabels, }: ModelsTableColumnDeps): ColumnDef[] => [ { id: MODEL_ID_COLUMN_ID, @@ -412,7 +414,15 @@ export const getModelsTableColumns = ({ enableSorting: false, size: 180, minSize: 110, - cell: ({ row }) => , + cell: ({ row }) => { + const credentialName = row.original.litellm_params?.litellm_credential_name; + return ( + + ); + }, }, { id: CREATED_BY_COLUMN_ID, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.integration.test.tsx index 3b9c714d354..d7ec325653f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.integration.test.tsx @@ -30,7 +30,14 @@ const renderForm = () => onCancel={vi.fn()} onSuccess={onSuccess} accessToken="test-token" - credentials={[{ credential_name: "bedrock-prod", credential_info: {}, credential_values: {} }]} + credentials={[ + { + credential_name: "bedrock-prod", + display_name: "Bedrock Prod", + credential_info: {}, + credential_values: {}, + }, + ]} />, ); @@ -166,6 +173,43 @@ describe("VectorStoreForm submit payload", () => { }); }); + it("shows the credential display name beside the name, searches by it, and still submits the name", async () => { + const user = setupUser(); + renderForm(); + + const picker = screen.getByPlaceholderText("Select or search for existing credentials"); + await user.click(picker); + + const option = await screen.findByRole("option", { name: /bedrock-prod/ }); + expect(option).toHaveTextContent("bedrock-prod"); + expect(option).toHaveTextContent("Bedrock Prod"); + + await user.type(picker, "bedrock prod"); + await user.click(await screen.findByRole("option", { name: /bedrock-prod/ })); + await user.type(screen.getByPlaceholderText("Enter vector store ID from your provider"), "vs-labeled"); + await submit(user); + + await vi.waitFor(() => expect(mockCreate).toHaveBeenCalledTimes(1)); + expect(createdPayload().litellm_credential_name).toBe("bedrock-prod"); + }); + + it("sends no credential when None is picked after a credential", async () => { + const user = setupUser(); + renderForm(); + + const picker = screen.getByPlaceholderText("Select or search for existing credentials"); + await user.click(picker); + await user.click(await screen.findByRole("option", { name: /bedrock-prod/ })); + expect(picker).toHaveValue("Bedrock Prod"); + await user.click(picker); + await user.click(await screen.findByRole("option", { name: "None" })); + await user.type(screen.getByPlaceholderText("Enter vector store ID from your provider"), "vs-none"); + await submit(user); + + await vi.waitFor(() => expect(mockCreate).toHaveBeenCalledTimes(1)); + expect(createdPayload().litellm_credential_name).toBeUndefined(); + }); + it("blocks the request and reports invalid metadata JSON instead of submitting", async () => { const user = setupUser(); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx index 01abfe82920..d82f56f3870 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx @@ -33,6 +33,8 @@ import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@ import { Textarea } from "@/components/ui/textarea"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { useZodForm } from "@/lib/forms/useZodForm"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { credentialOptions } from "@/components/shared/credentialOptions"; const EMBEDDING_MODEL_RENAME_PROVIDERS = new Set(["milvus", "valkey", "mongodb"]); @@ -159,11 +161,6 @@ const EMPTY_VALUES: VectorStoreFormValues = { valkey_embedding_field: "embedding", }; -interface CredentialOption { - label: string; - value: string | null; -} - const labelWithHint = (label: string, hint: string): React.ReactNode => ( <> {label} @@ -225,14 +222,6 @@ const VectorStoreForm: React.FC = ({ loadModels(); }, [accessToken]); - const credentialOptions: CredentialOption[] = [ - { value: null, label: "None" }, - ...credentials.map((credential) => ({ - value: credential.credential_name, - label: credential.credential_name, - })), - ]; - const makeProviderChangeHandler = (onChange: (provider: string) => void) => (provider: string | null) => { if (provider === null) return; onChange(provider); @@ -495,35 +484,14 @@ const VectorStoreForm: React.FC = ({ "Optionally select API provider credentials for this vector store eg. Bedrock API KEY", )} > - {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( - option.value === value) ?? null} - onValueChange={(option: CredentialOption | null) => onChange(option ? option.value : undefined)} - itemToStringLabel={(option: CredentialOption) => option.label} - isItemEqualToValue={(option: CredentialOption, selected: CredentialOption) => - option.value === selected.value - } - > - - - No matching credentials - - {(option: CredentialOption) => ( - - {option.label} - - )} - - - + {({ id, value, onChange }) => ( + onChange(selected === "" || selected === null ? undefined : selected)} + /> )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx index b8d82eee0a2..71d6a4545f5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx @@ -20,20 +20,14 @@ import { StatusBadge } from "@/components/shared/table_cells"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Card, CardContent } from "@/components/ui/card"; -import { - Combobox, - ComboboxContent, - ComboboxEmpty, - ComboboxInput, - ComboboxItem, - ComboboxList, -} from "@/components/ui/combobox"; import { Input } from "@/components/ui/input"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Textarea } from "@/components/ui/textarea"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { useZodForm } from "@/lib/forms/useZodForm"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { credentialOptions } from "@/components/shared/credentialOptions"; interface VectorStoreInfoViewProps { vectorStoreId: string; @@ -68,11 +62,6 @@ const toFormValues = (vectorStore: VectorStore): VectorStoreEditValues => ({ litellm_credential_name: vectorStore.litellm_credential_name, }); -interface CredentialOption { - label: string; - value: string | null; -} - const labelWithHint = (label: string, hint: string): React.ReactNode => ( <> {label} @@ -175,14 +164,6 @@ const VectorStoreInfoView: React.FC = ({ } }; - const credentialOptions: CredentialOption[] = [ - { value: null, label: "None" }, - ...credentials.map((credential) => ({ - value: credential.credential_name, - label: credential.credential_name, - })), - ]; - if (loadFailed) { return (
@@ -321,43 +302,16 @@ const VectorStoreInfoView: React.FC = ({

- {({ - id, - value, - onChange, - "aria-invalid": ariaInvalid, - "aria-describedby": ariaDescribedBy, - }) => ( - option.value === value) ?? null} - onValueChange={(option: CredentialOption | null) => - onChange(option ? option.value : undefined) + {({ id, value, onChange }) => ( + + onChange(selected === "" || selected === null ? undefined : selected) } - itemToStringLabel={(option: CredentialOption) => option.label} - isItemEqualToValue={(option: CredentialOption, selected: CredentialOption) => - option.value === selected.value - } - > - - - No matching credentials - - {(option: CredentialOption) => ( - - {option.label} - - )} - - - + /> )} diff --git a/ui/litellm-dashboard/src/components/ModelInfoEditForm.test.tsx b/ui/litellm-dashboard/src/components/ModelInfoEditForm.test.tsx new file mode 100644 index 00000000000..1a05374d883 --- /dev/null +++ b/ui/litellm-dashboard/src/components/ModelInfoEditForm.test.tsx @@ -0,0 +1,117 @@ +import { render, screen } from "@testing-library/react"; +import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import type { CredentialItem } from "@/components/networking"; + +import ModelInfoEditForm from "./ModelInfoEditForm"; + +vi.mock("@/components/networking", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + vectorStoreListCall: vi.fn().mockResolvedValue({ data: [] }), + }; +}); + +const credentialsList: CredentialItem[] = [ + { + credential_name: "openai-main", + display_name: "Main OpenAI", + credential_values: {}, + credential_info: { custom_llm_provider: "openai" }, + }, +]; + +const renderForm = ({ + onSubmit = vi.fn().mockResolvedValue(undefined), + isEditing = true, + litellmParams = { model: "gpt-4o" } as Record, +} = {}) => + render( + , + ); + +describe("ModelInfoEditForm existing-credentials picker", () => { + it("shows the display name beside the credential name, searches by it, and submits the name", async () => { + const onSubmit = vi.fn().mockResolvedValue(undefined); + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + renderForm({ onSubmit }); + + const picker = await screen.findByPlaceholderText("Select or search for existing credentials"); + await user.click(picker); + + const option = await screen.findByRole("option", { name: /openai-main/ }); + expect(option).toHaveTextContent("openai-main"); + expect(option).toHaveTextContent("Main OpenAI"); + + await user.type(picker, "main openai"); + await user.click(await screen.findByRole("option", { name: /openai-main/ })); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await vi.waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0].litellm_credential_name).toBe("openai-main"); + }); + + it("keeps the attached credential when the search text is emptied and only clears through None", async () => { + const onSubmit = vi.fn().mockResolvedValue(undefined); + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + renderForm({ onSubmit, litellmParams: { model: "gpt-4o", litellm_credential_name: "openai-main" } }); + + const picker = await screen.findByPlaceholderText("Select or search for existing credentials"); + expect(picker).toHaveValue("Main OpenAI"); + await user.clear(picker); + await user.tab(); + await user.click(screen.getByRole("button", { name: /save changes/i })); + await vi.waitFor(() => expect(onSubmit).toHaveBeenCalledTimes(1)); + expect(onSubmit.mock.calls[0][0].litellm_credential_name).toBe("openai-main"); + + await user.click(picker); + await user.click(await screen.findByRole("option", { name: "None" })); + await user.click(screen.getByRole("button", { name: /save changes/i })); + await vi.waitFor(() => expect(onSubmit).toHaveBeenCalledTimes(2)); + expect(onSubmit.mock.calls[1][0].litellm_credential_name).toBeNull(); + }); + + it("shows the attached credential's display name when not editing", () => { + renderForm({ isEditing: false, litellmParams: { model: "gpt-4o", litellm_credential_name: "openai-main" } }); + + expect(screen.getByText("Main OpenAI")).toBeInTheDocument(); + expect(screen.queryByText("openai-main")).not.toBeInTheDocument(); + }); + + it("falls back to the credential name, then Manual, when not editing", () => { + const { unmount } = renderForm({ + isEditing: false, + litellmParams: { model: "gpt-4o", litellm_credential_name: "not-listed" }, + }); + expect(screen.getByText("not-listed")).toBeInTheDocument(); + unmount(); + + renderForm({ isEditing: false }); + expect(screen.getByText("Manual")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx b/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx index 9ee49c262d8..2a451711407 100644 --- a/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx +++ b/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx @@ -14,6 +14,8 @@ import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { credentialLabelsByName, credentialOptions } from "@/components/shared/credentialOptions"; import { Switch } from "@/components/ui/switch"; import { Textarea } from "@/components/ui/textarea"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; @@ -345,6 +347,9 @@ const ChipList: React.FC<{ values: unknown; emptyLabel: string }> = ({ values, e ); }; +const attachedCredentialLabel = (credentialName: string | null | undefined, credentials: CredentialItem[]): string => + credentialName ? credentialLabelsByName(credentials).get(credentialName) ?? credentialName : "Manual"; + const ModelInfoEditForm: React.FC = ({ localModelData, modelData, @@ -634,36 +639,23 @@ const ModelInfoEditForm: React.FC = ({ Existing Credentials {isEditing ? ( - {({ id, value, onChange, onBlur }) => { - const items: { value: string | null; label: string }[] = [ - { value: null, label: "None" }, - ...credentialsList.map((credential) => ({ - value: credential.credential_name, - label: credential.credential_name, - })), - ]; - return ( - - ); - }} + {({ id, value, onChange }) => ( + { + if (selected !== null) onChange(selected === "" ? null : selected); + }} + /> + )} ) : ( - {localModelData.litellm_params?.litellm_credential_name || "Manual"} + + {attachedCredentialLabel(localModelData.litellm_params?.litellm_credential_name, credentialsList)} + )}
diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx index 93a43d5e533..7ab757cfed9 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.test.tsx @@ -1,6 +1,6 @@ import type { ProxyModel } from "@/app/(dashboard)/hooks/models/useModels"; import type { Organization } from "@/components/networking"; -import { screen } from "@testing-library/react"; +import { screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; @@ -660,4 +660,181 @@ describe("ModelSelect", () => { expect(screen.getByLabelText("model-4")).toBeInTheDocument(); expect(screen.queryByLabelText("model-5")).not.toBeInTheDocument(); }); + + it("should let a selected model that is no longer offered be found and deselected", async () => { + const user = userEvent.setup(); + const liveModels: ProxyModel[] = Array.from({ length: 6 }, (_, i) => ({ + id: `model-${i}`, + object: "model", + created: 1234567890, + owned_by: "test", + })); + mockUseAllProxyModels.mockReturnValue({ + data: { data: liveModels }, + isLoading: false, + } as unknown as ReturnType); + const liveIds = liveModels.map((m) => m.id); + + renderWithProviders(); + + await openModelList(user); + await user.type(screen.getAllByRole("combobox")[0], "retired"); + const retired = await screen.findByRole("option", { name: "retired-model" }); + expect(retired).toHaveAttribute("aria-selected", "true"); + + await user.click(retired); + + expect(mockOnChange).toHaveBeenCalledWith(liveIds); + }); + + it("should list selections that are no longer offered in an Unavailable group ahead of every other group", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + + const groups = within(screen.getByRole("listbox")).getAllByRole("group"); + expect( + groups.map((group) => + within(group) + .getAllByRole("option") + .map((option) => option.textContent), + ), + ).toEqual([ + ["retired-b", "retired-a", "retired-c", "retired-d", "retired-e", "retired-f"], + ["All Proxy Models", "No Default Models"], + ["All Openai models", "All Anthropic models"], + ["gpt-4", "claude-3"], + ]); + expect(screen.getByText("Unavailable")).toBeInTheDocument(); + }); + + it("should keep an unavailable selection removable while a special option is selected", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + const retired = screen.getByRole("option", { name: "retired-model" }); + expect(retired).not.toHaveAttribute("aria-disabled", "true"); + + await user.click(retired); + + expect(mockOnChange).toHaveBeenCalledWith(["all-proxy-models"]); + }); + + it("should not mark selections Unavailable when the model list could not be loaded", async () => { + const user = userEvent.setup(); + mockUseAllProxyModels.mockReturnValue({ + data: undefined, + isLoading: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + + expectOffered("All Proxy Models"); + expectNotOffered("Unavailable"); + }); + + it("should list selections outside the organization's model ceiling as Unavailable", async () => { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ + data: createMockOrganization(["other-model"]), + isLoading: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + await user.click(screen.getByRole("option", { name: "claude-3" })); + + expect(mockOnChange).toHaveBeenCalledWith(["gpt-4"]); + }); + + it("should keep the other selections when removing one unavailable model alongside a special option", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + await user.click(screen.getByRole("option", { name: "retired-a" })); + + expect(mockOnChange).toHaveBeenCalledWith(["all-proxy-models", "retired-b"]); + }); + + it("should not mark team selections Unavailable while the organization's model ceiling is unknown", async () => { + const user = userEvent.setup(); + mockUseOrganization.mockReturnValue({ + data: undefined, + isLoading: false, + } as unknown as ReturnType); + mockUseTeam.mockReturnValue({ + data: { team_id: "team-1", organization_models: null }, + isLoading: false, + isFetching: false, + } as unknown as ReturnType); + + renderWithProviders( + , + ); + + await openModelList(user); + + expectOffered("No Default Models"); + expectNotOffered("Unavailable"); + }); + + it("should not show an Unavailable group when every selection is offered", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await openModelList(user); + + expect(screen.getByRole("option", { name: "gpt-4" })).toHaveAttribute("aria-selected", "true"); + expectNotOffered("Unavailable"); + }); }); diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx index c29eba9d997..f169382b4ed 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx @@ -126,6 +126,24 @@ const filterModels = ( return filterFn(filterArgs); }; +const isOfferedListKnown = ( + proxyModelsLoaded: boolean, + context: ModelSelectProps["context"], + organizationID: string | undefined, + organizationModels: string[] | undefined, +) => proxyModelsLoaded && !(context === "team" && organizationID !== undefined && organizationModels === undefined); + +const unavailableGroups = ( + selectedOptions: ModelOption[], + offeredByValue: Map, + offeredListKnown: boolean, +): ModelOptionGroup[] => { + if (!offeredListKnown) return []; + const items = selectedOptions.filter((option) => !offeredByValue.has(option.value)); + if (items.length === 0) return []; + return [{ label: "Unavailable", items }]; +}; + export const ModelSelect = (props: ModelSelectProps) => { const anchor = useComboboxAnchor(); const { id, teamID, organizationID, options, context, dataTestId, value = [], onChange, style } = props; @@ -151,17 +169,10 @@ export const ModelSelect = (props: ModelSelectProps) => { const handleChange = (selected: ModelOption[]) => { const values = selected.map((option) => option.value); - const specialValues = values.filter(isSpecialOption); + const addedSpecialValues = values.filter((v) => isSpecialOption(v) && !value.includes(v)); + const addedSpecial = addedSpecialValues[addedSpecialValues.length - 1]; - let finalValues: string[]; - if (specialValues.length > 0) { - const lastSelectedSpecial = specialValues[specialValues.length - 1]; - finalValues = [lastSelectedSpecial]; - } else { - finalValues = values; - } - - onChange(finalValues); + onChange(addedSpecial === undefined ? values : [addedSpecial]); }; const filteredModels = filterModels(allProxyModels?.data ?? [], props, { @@ -171,7 +182,7 @@ export const ModelSelect = (props: ModelSelectProps) => { const { wildcard, regular } = splitWildcardModels(filteredModels); - const groups: ModelOptionGroup[] = [ + const offeredGroups: ModelOptionGroup[] = [ ...(includeSpecialOptions ? [ { @@ -228,8 +239,16 @@ export const ModelSelect = (props: ModelSelectProps) => { }, ]; - const optionsByValue = new Map(groups.flatMap((group) => group.items).map((option) => [option.value, option])); - const selectedOptions = value.map((v) => optionsByValue.get(v) ?? { label: v, value: v }); + const offeredByValue = new Map(offeredGroups.flatMap((group) => group.items).map((option) => [option.value, option])); + const selectedOptions = value.map((v) => offeredByValue.get(v) ?? { label: v, value: v }); + const groups: ModelOptionGroup[] = [ + ...unavailableGroups( + selectedOptions, + offeredByValue, + isOfferedListKnown(allProxyModels !== undefined, context, organizationID, organizationModels), + ), + ...offeredGroups, + ]; const overflowOptions = selectedOptions.slice(MAX_VISIBLE_MODEL_CHIPS); return ( diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx index 220d46647f6..45014d06a3a 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx @@ -150,6 +150,7 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi const credentials: CredentialItem[] = [ { credential_name: "test-credential", + display_name: "Prod OpenAI", credential_values: {}, credential_info: { custom_llm_provider: "openai", @@ -299,6 +300,47 @@ describe("AddModelForm", () => { expect(screen.queryByRole("switch")).not.toBeInTheDocument(); }); + describe("the existing-credentials picker", () => { + const openPicker = async () => { + const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); + mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true)); + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const props = createTestProps(); + renderWithProviders(); + const input = await screen.findByPlaceholderText("Select or search for existing credentials"); + await user.click(input); + return { user, input, form: props.form }; + }; + + it("shows each credential's display name and credential name together", async () => { + await openPicker(); + + const option = await screen.findByRole("option", { name: /test-credential/ }); + expect(option).toHaveTextContent("test-credential"); + expect(option).toHaveTextContent("Prod OpenAI"); + }); + + it("filters by the display name, not only the credential name", async () => { + const { user, input } = await openPicker(); + + await user.type(input, "prod open"); + + expect(await screen.findByRole("option", { name: /test-credential/ })).toBeInTheDocument(); + }); + + it("selecting by display name sets litellm_credential_name to the credential name", async () => { + const { user, input, form } = await openPicker(); + expect(screen.getByText("OR")).toBeInTheDocument(); + + await user.type(input, "prod open"); + await user.click(await screen.findByRole("option", { name: /test-credential/ })); + + expect(input).toHaveValue("Prod OpenAI"); + expect(form.getValues("litellm_credential_name")).toBe("test-credential"); + expect(screen.queryByText("OR")).not.toBeInTheDocument(); + }); + }); + it("should display the provider field and the Test Connect / Add Model buttons", async () => { const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true)); @@ -312,6 +354,18 @@ describe("AddModelForm", () => { expect(await screen.findByRole("button", { name: "Add Model" })).toBeInTheDocument(); }); + it("offers the Evaluation decisions mode", async () => { + const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); + mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true)); + + renderWithProviders(); + + await screen.findByText("Provider"); + await userEvent.click(screen.getByRole("combobox", { name: "Mode" })); + + expect(await screen.findByRole("option", { name: "Evaluation - /v1/decisions", exact: true })).toBeInTheDocument(); + }); + it("shows only the Close button in the connection test dialog footer", async () => { const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true)); diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index b5b6c439dc2..2b3593e0149 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -8,6 +8,7 @@ import { Field, FieldLabel } from "@/components/ui/field"; import { Card, CardContent } from "@/components/ui/card"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect"; +import { credentialOptions } from "@/components/shared/credentialOptions"; import { SimpleTooltip } from "@/components/ui/tooltip"; import { Info } from "lucide-react"; import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; @@ -151,16 +152,7 @@ const AddModelForm: React.FC = ({ [sortedProviderMetadata], ); - const credentialOptions: SearchSelectOption[] = useMemo( - () => [ - { label: "None", value: "" }, - ...credentials.map((credential) => ({ - label: credential.credential_name, - value: credential.credential_name, - })), - ], - [credentials], - ); + const credentialSelectOptions: SearchSelectOption[] = useMemo(() => credentialOptions(credentials), [credentials]); const applyProviderSelection = (provider: string | null) => { setSelectedProvider(provider); @@ -323,7 +315,7 @@ const AddModelForm: React.FC = ({ control.onChange(value === "" ? null : value)} /> diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx index 05da3f96109..289fd390c62 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx @@ -13,6 +13,7 @@ export const TEST_MODES = [ { value: "batch", label: "Batch - /batch" }, { value: "anthropic_messages", label: "Anthropic Messages - /v1/messages" }, { value: "ocr", label: "OCR - /ocr" }, + { value: "evaluation", label: "Evaluation - /v1/decisions" }, ]; // Define the available auto router routing strategies diff --git a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx index 923cf2aa0e3..6190b6aef78 100644 --- a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx @@ -102,6 +102,20 @@ describe("prepareModelAddRequest", () => { expect(deployment.litellmParamsObj.timeout).toBe(5); }); + it("saves the selected mode under model_info", async () => { + const formValues = { + model_mappings: [{ public_name: "Jev", litellm_model: "typesafe/jev-latest" }], + mode: "evaluation", + }; + + const deployments = await prepareModelAddRequest({ ...formValues }, "token", null); + + expect(deployments).toHaveLength(1); + const [deployment] = deployments!; + expect(deployment.modelInfoObj.mode).toBe("evaluation"); + expect(deployment.litellmParamsObj).not.toHaveProperty("mode"); + }); + it.each([ ["OpenAI", "openai/*"], ["Azure_AI_Studio", "azure_ai/*"], diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx index 8d94ae0d213..4fa3a3cc6bd 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx @@ -1,5 +1,6 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render, screen, waitFor } from "@testing-library/react"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import { Providers } from "../provider_info_helpers"; import { CredentialItem } from "../networking"; @@ -125,4 +126,142 @@ describe("CredentialModal", () => { expect(screen.getByLabelText("Credential Name:")).toBeDisabled(); }); }); + + describe("display name", () => { + const fillRequiredAddFields = async (user: ReturnType) => { + fireEvent.change(screen.getByLabelText("Credential Name:"), { target: { value: "new-cred" } }); + const providerInput = await screen.findByPlaceholderText("Select a provider"); + await user.click(providerInput); + await user.click(await screen.findByText("OpenAI")); + fireEvent.change(await screen.findByLabelText("OpenAI API Key"), { target: { value: "sk-test" } }); + }; + + it("submits the display name with the credential name in add mode", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const onSubmit = vi.fn(); + renderModal({ mode: "add", onSubmit }); + + await fillRequiredAddFields(user); + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: "Prod" } }); + fireEvent.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0]).toMatchObject({ credential_name: "new-cred", display_name: "Prod" }); + }); + + it("submits no display_name key when the display name is left blank in add mode", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const onSubmit = vi.fn(); + renderModal({ mode: "add", onSubmit }); + + await fillRequiredAddFields(user); + fireEvent.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("display_name"); + }); + + it("prefills the display name and submits an edited one", async () => { + const onSubmit = vi.fn(); + renderModal({ + mode: "edit", + onSubmit, + existingCredential: { ...mockCredential, display_name: "Prod" }, + }); + + const displayNameInput = screen.getByLabelText("Display Name:") as HTMLInputElement; + expect(displayNameInput.value).toBe("Prod"); + + fireEvent.change(displayNameInput, { target: { value: "Staging" } }); + fireEvent.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0]).toMatchObject({ + credential_name: "test-credential", + display_name: "Staging", + }); + }); + + it("submits display_name: null when the display name is cleared in edit mode", async () => { + const onSubmit = vi.fn(); + renderModal({ + mode: "edit", + onSubmit, + existingCredential: { ...mockCredential, display_name: "Prod" }, + }); + + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: "" } }); + fireEvent.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0].display_name).toBeNull(); + }); + + it("trims the display name before submitting in add mode", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const onSubmit = vi.fn(); + renderModal({ mode: "add", onSubmit }); + + await fillRequiredAddFields(user); + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: " Prod " } }); + fireEvent.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0]).toMatchObject({ credential_name: "new-cred", display_name: "Prod" }); + }); + + it("submits display_name: null when the display name is whitespace-only in edit mode", async () => { + const onSubmit = vi.fn(); + renderModal({ + mode: "edit", + onSubmit, + existingCredential: { ...mockCredential, display_name: "Prod OpenAI" }, + }); + + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: " " } }); + fireEvent.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0].display_name).toBeNull(); + }); + + it("leaves display_name out of the edit when it was not changed", async () => { + const onSubmit = vi.fn(); + renderModal({ + mode: "edit", + onSubmit, + existingCredential: { ...mockCredential, display_name: "Prod" }, + }); + + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: " Prod " } }); + fireEvent.click(screen.getByRole("button", { name: "Update Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("display_name"); + }); + + it("accepts a 255-character display name and blocks a 256-character one", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const onSubmit = vi.fn(); + renderModal({ mode: "add", onSubmit }); + + await fillRequiredAddFields(user); + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: "x".repeat(256) } }); + fireEvent.click(screen.getByRole("button", { name: "Add Credential" })); + expect(await screen.findByText("Display name must be at most 255 characters")).toBeInTheDocument(); + expect(onSubmit).not.toHaveBeenCalled(); + + fireEvent.change(screen.getByLabelText("Display Name:"), { target: { value: "x".repeat(255) } }); + fireEvent.click(screen.getByRole("button", { name: "Add Credential" })); + await waitFor(() => expect(onSubmit).toHaveBeenCalledTimes(1)); + expect(onSubmit.mock.calls[0][0].display_name).toBe("x".repeat(255)); + }); + + it("keeps the credential name read-only while the display name stays editable in edit mode", () => { + renderModal({ mode: "edit", existingCredential: { ...mockCredential, display_name: "Prod" } }); + + expect(screen.getByLabelText("Credential Name:")).toBeDisabled(); + expect(screen.getByLabelText("Display Name:")).toBeEnabled(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx index 202274d2aa2..61afd5eea99 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx @@ -62,6 +62,22 @@ interface CredentialModalProps { const sameProvider = (left: string | null | undefined, right: string | null | undefined): boolean => (left ?? "").toLowerCase() === (right ?? "").toLowerCase(); +const DISPLAY_NAME_MAX_LENGTH = 255; + +const displayNameChange = ( + value: unknown, + existingCredential: CredentialItem | null | undefined, +): { display_name?: string | null } => { + const trimmed = typeof value === "string" ? value.trim() : ""; + if (!existingCredential) { + return trimmed ? { display_name: trimmed } : {}; + } + if (trimmed === (existingCredential.display_name ?? "")) { + return {}; + } + return { display_name: trimmed || null }; +}; + const initialFormValues = ( existingCredential: CredentialItem | null | undefined, initialProvider: string | null | undefined, @@ -69,6 +85,7 @@ const initialFormValues = ( if (existingCredential) { return { credential_name: existingCredential.credential_name, + display_name: existingCredential.display_name ?? "", custom_llm_provider: existingCredential.credential_info.custom_llm_provider, ...Object.fromEntries( Object.entries(existingCredential.credential_values || {}).map(([key, value]) => [key, value ?? null]), @@ -127,6 +144,7 @@ export default function CredentialModal({ const meta = { credential_name: values.credential_name, custom_llm_provider: values.custom_llm_provider, + ...displayNameChange(values.display_name, existingCredential), }; if (!isEdit) { onSubmit({ ...meta, ...buildCreateCredentialValues(withoutRestrictedFields(values), selection) }, []); @@ -170,12 +188,36 @@ export default function CredentialModal({ value={typeof control.value === "string" ? control.value : ""} onChange={control.onChange} onBlur={control.onBlur} - placeholder="Enter a friendly name for these credentials" + placeholder="Unique name that models reference this credential by" disabled={isEdit} /> )} + + typeof value !== "string" || + value.trim().length <= DISPLAY_NAME_MAX_LENGTH || + `Display name must be at most ${DISPLAY_NAME_MAX_LENGTH} characters`, + }, + }} + className="mb-4" + > + {(control) => ( + + )} + + { expect(await window.navigator.clipboard.readText()).toBe("b-openai-key"); }); + it("should show the display name with the credential name beneath it, and the bare name when unset", () => { + const credentials: CredentialItem[] = [ + { + credential_name: "openai-prod", + display_name: "Prod OpenAI", + credential_values: {}, + credential_info: { custom_llm_provider: "openai" }, + }, + { credential_name: "plain-key", credential_values: {}, credential_info: { custom_llm_provider: "openai" } }, + ]; + render(); + + const labeledRow = screen.getByRole("row", { name: /Prod OpenAI/ }); + expect(within(labeledRow).getByText("openai-prod")).toBeInTheDocument(); + const plainRow = screen.getByRole("row", { name: /plain-key/ }); + expect(within(plainRow).getAllByText("plain-key")).toHaveLength(1); + }); + + it("should sort by the display name when one is set", () => { + const credentials: CredentialItem[] = [ + { credential_name: "a-key", display_name: "zulu", credential_values: {}, credential_info: {} }, + { credential_name: "b-key", credential_values: {}, credential_info: {} }, + ]; + render(); + + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("b-key")).toBeInTheDocument(); + expect(within(rows[1]).getByText("zulu")).toBeInTheDocument(); + }); + + it("should badge a config credential and block editing and deleting it", async () => { + const user = userEvent.setup(); + const credentials: CredentialItem[] = [ + { credential_name: "from-config", source: "config", credential_values: {}, credential_info: {} }, + { credential_name: "from-db", source: "db", credential_values: {}, credential_info: {} }, + ]; + render(); + + expect(within(screen.getByRole("row", { name: /from-config/ })).getByText("Config")).toBeInTheDocument(); + expect(within(screen.getByRole("row", { name: /from-db/ })).queryByText("Config")).not.toBeInTheDocument(); + + await user.click(screen.getByTestId("credential-actions-from-config")); + expect(await screen.findByTestId("credential-config-owned-hint")).toBeInTheDocument(); + const edit = screen.getByTestId("credential-action-edit"); + const remove = screen.getByTestId("credential-action-delete"); + expect(edit).toHaveAttribute("data-disabled"); + expect(remove).toHaveAttribute("data-disabled"); + await user.click(edit); + await user.click(remove); + expect(mockOnEdit).not.toHaveBeenCalled(); + expect(mockOnDelete).not.toHaveBeenCalled(); + }); + + it("should keep editing enabled for a DB credential", async () => { + const user = userEvent.setup(); + const credentials: CredentialItem[] = [ + { credential_name: "from-db", source: "db", credential_values: {}, credential_info: {} }, + ]; + render(); + + await user.click(screen.getByTestId("credential-actions-from-db")); + expect(await screen.findByTestId("credential-action-edit")).not.toHaveAttribute("data-disabled"); + expect(screen.queryByTestId("credential-config-owned-hint")).not.toBeInTheDocument(); + }); + it("should not render the actions menu when the user cannot modify credentials", () => { render(); // Read parity: names still render... diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialsTableColumns.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialsTableColumns.tsx index 69ea0acd330..581b4d3da32 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialsTableColumns.tsx @@ -6,13 +6,16 @@ import { Copy, MoreHorizontal, Pencil, Trash2 } from "lucide-react"; import { CredentialItem } from "@/components/networking"; import { getProviderLogoAndName } from "@/components/provider_info_helpers"; import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { credentialLabel } from "@/components/shared/credentialOptions"; import { IdentityCell } from "@/components/shared/table_cells"; import { Badge } from "@/components/ui/badge"; import { buttonVariants } from "@/components/ui/button"; import { DropdownMenu, DropdownMenuContent, + DropdownMenuGroup, DropdownMenuItem, + DropdownMenuLabel, DropdownMenuSeparator, DropdownMenuTrigger, } from "@/components/ui/dropdown-menu"; @@ -51,6 +54,7 @@ interface CredentialRowActionsProps { } function CredentialRowActions({ credential, onEdit, onDelete }: CredentialRowActionsProps) { + const configOwned = credential.source === "config"; return ( - onEdit(credential)}> + {configOwned && ( + <> + + + Defined in config.yaml. Edit the file to change or delete it. + + + + + )} + onEdit(credential)} + > Edit @@ -76,6 +94,7 @@ function CredentialRowActions({ credential, onEdit, onDelete }: CredentialRowAct onDelete(credential)} > @@ -100,13 +119,19 @@ export const getCredentialsTableColumns = ({ const dataColumns: ColumnDef[] = [ { id: "credential_name", - accessorKey: "credential_name", + accessorFn: credentialLabel, meta: { title: "Credential Name" }, header: ({ column }) => , size: 260, enableSorting: true, cell: ({ row }) => ( - + Config
: undefined} + className="max-w-72" + titleClassName="font-medium" + /> ), }, { diff --git a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts index d7218f6f4d3..efad20dddaf 100644 --- a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts +++ b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.test.ts @@ -1,6 +1,10 @@ import { describe, expect, it, vi } from "vitest"; import { Providers } from "../provider_info_helpers"; -import { resetCredentialFormOnProviderChange } from "./credential_form_helpers"; +import { + buildCredential, + resetCredentialFormOnProviderChange, + withoutRestrictedFields, +} from "./credential_form_helpers"; /** * Build a minimal FormInstance stub that records calls. We don't depend @@ -58,6 +62,14 @@ describe("resetCredentialFormOnProviderChange", () => { expect(fields.credential_name).toBe("my-prod-key"); }); + it("preserves a display name typed before the provider was picked", () => { + const { stub, fields } = makeFormStub({ credential_name: "my-prod-key", display_name: "Prod OpenAI" }); + + resetCredentialFormOnProviderChange(stub, Providers.OpenAI, vi.fn()); + + expect(fields.display_name).toBe("Prod OpenAI"); + }); + it("updates custom_llm_provider and selectedProvider state to the new value", () => { const { stub, fields } = makeFormStub({ credential_name: "x" }); const setSelectedProvider = vi.fn(); @@ -81,3 +93,23 @@ describe("resetCredentialFormOnProviderChange", () => { expect(credentialNameCalls).toHaveLength(0); }); }); + +describe("buildCredential", () => { + const values = { credential_name: "openai-prod", custom_llm_provider: "openai", api_key: "sk-test" }; + + it.each([ + ["a label", "Prod OpenAI"], + ["a cleared label", null], + ])("sends %s as a top-level display_name, never as a credential value", (_, displayName) => { + const formValues = { ...values, display_name: displayName }; + + const credential = buildCredential(formValues, withoutRestrictedFields(formValues)); + + expect(credential.display_name).toBe(displayName); + expect(credential.credential_values).toEqual({ api_key: "sk-test" }); + }); + + it("leaves display_name out when the form never set it", () => { + expect(buildCredential(values, withoutRestrictedFields(values))).not.toHaveProperty("display_name"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts index 4db1bb78371..7de443ae7cc 100644 --- a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts +++ b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts @@ -21,10 +21,11 @@ interface CredentialFormAdapter { * The credential name is preserved because it's a user-supplied label * that shouldn't reset just because the admin re-selected a provider. */ -const restrictedFields: readonly string[] = ["credential_name", "custom_llm_provider"]; +const restrictedFields: readonly string[] = ["credential_name", "display_name", "custom_llm_provider"]; export const buildCredential = (values: Record, credentialValues: Record) => ({ credential_name: values.credential_name as string, + ...(values.display_name !== undefined ? { display_name: values.display_name as string | null } : {}), credential_values: credentialValues, credential_info: { custom_llm_provider: values.custom_llm_provider as string, @@ -40,10 +41,14 @@ export function resetCredentialFormOnProviderChange( setSelectedProvider: (p: string | null) => void, ): void { const preservedName = form.getFieldValue("credential_name"); + const preservedDisplayName = form.getFieldValue("display_name"); form.resetFields(); if (preservedName !== undefined) { form.setFieldValue("credential_name", preservedName); } + if (preservedDisplayName !== undefined) { + form.setFieldValue("display_name", preservedDisplayName); + } setSelectedProvider(newProvider); form.setFieldValue("custom_llm_provider", newProvider); } diff --git a/ui/litellm-dashboard/src/components/model_info_view.test.tsx b/ui/litellm-dashboard/src/components/model_info_view.test.tsx index 037ebc4040e..f555911eb5c 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.test.tsx @@ -1542,6 +1542,27 @@ describe("ModelInfoView", () => { expect(screen.getByTestId("reuse-credentials-button")).toBeInTheDocument(); }); + it("names the attached credential by its display name in the re-use dialog", async () => { + mockCredentialListCall.mockResolvedValue({ + credentials: [ + { + credential_name: "selected-credential", + display_name: "Selected Label", + credential_values: {}, + credential_info: {}, + }, + ], + } as never); + const user = userEvent.setup(); + render(, { wrapper }); + + await user.click(await screen.findByTestId("reuse-credentials-button")); + + const dialog = await screen.findByRole("dialog", { name: "Using Existing Credential" }); + await vi.waitFor(() => expect(dialog).toHaveTextContent("Selected Label")); + expect(dialog).not.toHaveTextContent("selected-credential"); + }); + it.each([["auto_router/adaptive_router"], ["auto_router/quality_router"]])( "offers no Test Connection for %s, whose targets it cannot build", async (model) => { @@ -1577,13 +1598,13 @@ describe("ModelInfoView", () => { const openCredentialSelect = async (user: ReturnType, triggerText?: string) => { const trigger = screen .getAllByRole("combobox") - .filter((element) => element.getAttribute("data-slot") === "select-trigger") - .find((element) => triggerText === undefined || element.textContent?.includes(triggerText)); + .filter((element) => element.getAttribute("placeholder") === "Select or search for existing credentials") + .find((element) => triggerText === undefined || (element as HTMLInputElement).value.includes(triggerText)); if (trigger === undefined) { throw new Error(`Could not find credential selector${triggerText ? ` with ${triggerText}` : ""}`); } await user.click(trigger); - await screen.findByRole("combobox", { expanded: true }); + await screen.findByRole("option", { name: "None" }); }; const save = async (user: ReturnType) => { @@ -1837,8 +1858,8 @@ describe("ModelInfoView", () => { const user = userEvent.setup(); await enterEditMode(user); - await openSelect(user, "selected-credential"); - await user.click(await screen.findByText("other-credential")); + await openCredentialSelect(user, "selected-credential"); + await user.click(await screen.findByRole("option", { name: "other-credential" })); const payload = await save(user); @@ -1902,11 +1923,10 @@ describe("ModelInfoView", () => { await user.click(screen.getByRole("button", { name: /cancel/i })); await user.click(await screen.findByRole("button", { name: /edit settings/i })); - const credentialTrigger: HTMLElement = screen + const credentialTrigger = screen .getAllByRole("combobox") - .filter((element) => element.getAttribute("data-slot") === "select-trigger") - .at(0) as HTMLElement; - expect(credentialTrigger).toHaveTextContent("selected-credential"); + .find((element) => element.getAttribute("placeholder") === "Select or search for existing credentials"); + expect(credentialTrigger).toHaveValue("selected-credential"); }); it("shows Manual in read mode after saving None", async () => { diff --git a/ui/litellm-dashboard/src/components/model_info_view.tsx b/ui/litellm-dashboard/src/components/model_info_view.tsx index 4fe49a28936..26c6308e3ac 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.tsx @@ -28,6 +28,7 @@ import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import DeleteResourceModal from "./common_components/DeleteResourceModal"; import EditAutoRouterModal from "./edit_auto_router/edit_auto_router_modal"; import ReuseCredentialsModal from "./model_add/reuse_credentials"; +import { credentialLabelsByName } from "./shared/credentialOptions"; import { toast } from "@/lib/toast"; import { CredentialItem, @@ -808,7 +809,10 @@ export default function ModelInfoView({ Using Existing Credential -

{modelData.litellm_params.litellm_credential_name}

+

+ {credentialLabelsByName(credentialsList).get(modelData.litellm_params.litellm_credential_name ?? "") ?? + modelData.litellm_params.litellm_credential_name} +