From 4d6fc36fa0c5e018bcc2a4cc9e4492f05bd19514 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Fri, 26 Jun 2026 19:21:55 -0700 Subject: [PATCH 01/37] test(e2e): move rust OCR e2e into llm_translation on the shared harness The rust OCR smoke lived under tests/e2e/gateway and spoke raw httpx with ad-hoc dataclasses, diverging from the rest of tests/e2e. Move it to tests/e2e/llm_translation and rebuild it on the shared harness: typed pydantic bodies in models.py (OcrDocument/OcrBody/OcrPage/OcrResponse), a Gateway.ocr() route through the shared transport, Result/unwrap for outcomes, the e2e marker, and the client/scoped_key fixtures. No test touches httpx or requests directly now. Behavior preserved: the config-presence check still reads gateway/litellm-config.yml without a proxy, /model/info confirms the proxy loaded every rust-ocr deployment, and each provider case asserts a well-formed OCR document over /v1/ocr. Also add tests/e2e/CONTRIBUTING.md documenting the end-to-end testing flow so new features land with coverage that walks the feature like production does. --- tests/e2e/CONTRIBUTING.md | 94 +++++++++++ tests/e2e/e2e_gateway.py | 18 ++- tests/e2e/gateway/test_ocr_rust_e2e.py | 148 ------------------ .../e2e/llm_translation/test_ocr_rust_e2e.py | 117 ++++++++++++++ tests/e2e/models.py | 38 ++++- 5 files changed, 254 insertions(+), 161 deletions(-) create mode 100644 tests/e2e/CONTRIBUTING.md delete mode 100644 tests/e2e/gateway/test_ocr_rust_e2e.py create mode 100644 tests/e2e/llm_translation/test_ocr_rust_e2e.py diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md new file mode 100644 index 00000000000..c0c08453b25 --- /dev/null +++ b/tests/e2e/CONTRIBUTING.md @@ -0,0 +1,94 @@ +# Contributing e2e tests + +This directory holds the live end-to-end suites that prove product correctness against a real running proxy and real provider APIs. The goal of this guide is simple: when you ship a feature, you add e2e coverage that walks that feature the way production does, across every route and edge case it touches, so a later change that breaks it fails here first + +Read this before adding a test. The harness already encodes most of the rules; your job is to fill in the feature-specific behavior, not to reinvent plumbing + +## What a complete test looks like + +A feature test is complete only when it walks the feature end to end, in this order + +1. CREATE the resource (key / team / budget / ...) and immediately queue its deletion +2. CONFIGURE the feature's setting on it (assign the budget, turn on compression, set the limit) +3. ACT; drive real traffic through the gateway exactly like prod does (right model, real auth headers, enough calls to actually trigger the behavior) +4. SETTLE; poll the DB / spend logs until the write lands. Writes are eventually consistent (spend flushes on proxy_batch_write_at, ~60s), so poll to a deadline. Never sleep once +5. ASSERT the recorded state the feature promises (spend > budget, cost reduced, tag attributed, ...) +6. ASSERT the enforced behavior the gateway returns (429 budget_exceeded, block, refusal, ...) +7. TEARDOWN; every resource you created is deleted + +### The one rule that makes it complete + +It must assert BOTH sides: the recorded state (step 5) AND the enforced behavior (step 6) + +A test that only checks "the call went through", or only checks spend without checking the 429, is not complete; it is checking plumbing, not the product promise + +### Example: budget enforcement + +``` +create a key -> (1) +assign a budget -> (2) +send a bunch of calls -> (3) +poll for db spend -> (4) +assert spend > budget -> (5) +assert status_code == 429 -> (6) ("budget_exceeded") +key auto-deleted on teardown -> (7) +``` + +### The skeleton every test fills in + +``` +setup -> create the resource + queue cleanup +configure -> apply the feature's knob +act -> send real calls like production +settle -> poll the DB until the write lands +assert -> recorded state is correct (the feature happened) +assert -> gateway enforced it (the product promise held) +teardown -> delete everything you created +``` + +If a step is missing, the test is not done. That is the whole pattern + +## Style: lay the pattern down in a class + +Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being + +```python +# pseudo-code to convey intent, not the real API +class TestPromptCompression: + def test_prompt_compression_add_to_virtual_key(self): + new_key = self.resources.create_key(user_id, compression=True) # turn the feature on + resources._defer(new_key) # queue key deletion + + def test_prompt_compression_accumulate_spend(self, key_id, user_id): + for _ in range(10): + response = self.resources.gateway.post("gemini-2.5-flash", key_id, user_id) + compressed_value = ... + assert response.cost == compressed_value # the cost was actually reduced +``` + +That snippet only conveys intent. What you actually write uses the real harness: the `client` fixture for your suite, the `scoped_key` fixture for an auto-deleted key, typed pydantic bodies from `models.py`, and `unwrap(...)` on the tagged-union result. `tests/e2e/llm_translation/test_custom_pricing_e2e.py` is the reference to copy from; it creates a scoped key, drives a real gemini call, polls `/spend/logs` to a deadline for the cost-breakdown row, then asserts the input and output costs match the configured custom rates and that a sibling deployment kept its own price. Read it before writing yours + +## Use the shared transport; never touch requests directly + +Every HTTP call goes through the shared transport, never through `requests.*` in a test. `e2e_http.py` is the only module permitted to call `requests.*`, and that is enforced in CI by `tests/code_coverage_tests/check_e2e_no_raw_requests.py`. A test that imports requests will fail the check + +The shape is layered so tests stay declarative + +`transport.py` exposes a `Transport` Protocol with `post`, `get`, `delete`, `send`, `stream`, `probe`, plus `bearer(key)` and the `master` header. `HttpTransport` fulfils it, and `SplitTransport` routes each call by path to the data plane or the control plane so a split control-plane/data-plane deployment works without any change in the test + +`e2e_gateway.py` holds `Gateway`, a frozen dataclass that wraps a `Transport` and adds the operations tests reuse: `generate_key` / `delete_key` / `key_info`, `model_info`, the LLM calls `chat` / `chat_stream` / `embed` / `ocr`, the spend read-back `spend_logs`, and the poll helpers `poll_logs_for_key` / `poll_logs_for_request_id` that loop to `poll_timeout` instead of sleeping once. Add a new route as a method here so other suites get it for free + +Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture + +Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip + +Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache + +## Pre-commit steps + +Before you push + +- Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets +- Bring the stack up with docker-compose from this directory and run your suite locally against it, so you exercise the same skip-vs-fail path CI does +- Use the config at `tests/e2e/gateway/litellm-config.yml` if your feature needs a model, pricing override, guardrail, or other proxy setting declared up front; add the deployment there and read it back in the test rather than hardcoding values +- Capture screenshots of the tests passing and attach them to the PR as proof of fix diff --git a/tests/e2e/e2e_gateway.py b/tests/e2e/e2e_gateway.py index b700145d434..a67ea594a71 100644 --- a/tests/e2e/e2e_gateway.py +++ b/tests/e2e/e2e_gateway.py @@ -34,6 +34,8 @@ from models import ( KeyInfoResponse, ModelInfoEntry, ModelInfoResponse, + OcrBody, + OcrResponse, SpendLogRow, SpendLogs, SpendLogsParams, @@ -120,9 +122,7 @@ class Gateway: ) def chat_stream(self, key: str, body: ChatBody) -> StreamingResponse: - return self.transport.stream( - "/chat/completions", headers=self.transport.bearer(key), json=body - ) + return self.transport.stream("/chat/completions", headers=self.transport.bearer(key), json=body) def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]: return self.transport.post( @@ -132,6 +132,14 @@ class Gateway: response_type=EmbedResponse, ) + def ocr(self, key: str, body: OcrBody) -> Result[OcrResponse]: + return self.transport.post( + "/v1/ocr", + headers=self.transport.bearer(key), + json=body, + response_type=OcrResponse, + ) + # ---- spend read-back ------------------------------------------------ def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]: @@ -150,9 +158,7 @@ class Gateway: def poll_logs_for_key( self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None ) -> list[SpendLogRow]: - return self._poll( - lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate - ) + return self._poll(lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate) def poll_logs_for_request_id( self, diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py deleted file mode 100644 index 6ce59b2b5ac..00000000000 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ /dev/null @@ -1,148 +0,0 @@ -""" -Gateway E2E smoke for Rust-backed OCR. - -Start the proxy with: - -LITELLM_USE_RUST_OCR=1 litellm --config tests/e2e/gateway/litellm-config.yml --port 4000 -""" - -from __future__ import annotations - -import os -from dataclasses import dataclass -from pathlib import Path -from typing import Any - -import httpx -import pytest -import yaml - -TEST_PDF_URL = ( - "https://cdn.jsdelivr.net/gh/BerriAI/litellm" - "@d769e81c90d453240c61fc572cdb27fae06a89d0" - "/tests/llm_translation/fixtures/dummy.pdf" -) -TEST_IMAGE_URL = ( - "https://cdn.jsdelivr.net/gh/BerriAI/litellm" - "@d769e81c90d453240c61fc572cdb27fae06a89d0" - "/tests/image_gen_tests/test_image.png" -) - -RUST_OCR_GATEWAY_CASES = [ - pytest.param( - "rust-ocr-mistral", - {"type": "document_url", "document_url": TEST_PDF_URL}, - id="mistral", - ), - pytest.param( - "rust-ocr-azure-ai", - {"type": "document_url", "document_url": TEST_PDF_URL}, - id="azure_ai", - ), - pytest.param( - "rust-ocr-azure-document-intelligence", - {"type": "document_url", "document_url": TEST_PDF_URL}, - id="azure_document_intelligence", - ), - pytest.param( - "rust-ocr-vertex-mistral", - {"type": "document_url", "document_url": TEST_PDF_URL}, - id="vertex_mistral", - ), - pytest.param( - "rust-ocr-vertex-deepseek", - { - "type": "image_url", - "image_url": os.getenv("RUST_OCR_IMAGE_URL", TEST_IMAGE_URL), - }, - id="vertex_deepseek", - ), -] - -CONFIG_PATH = Path(__file__).with_name("litellm-config.yml") - - -@dataclass(frozen=True) -class OcrGateway: - base_url: str - master_key: str - - def model_names(self) -> set[str]: - with httpx.Client( - timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120")) - ) as client: - response = client.get( - f"{self.base_url.rstrip('/')}/model/info", - headers={"Authorization": f"Bearer {self.master_key}"}, - ) - assert response.status_code == 200, response.text - return { - model["model_name"] - for model in response.json().get("data", []) - if "model_name" in model - } - - def ocr(self, model: str, document: dict[str, str]) -> httpx.Response: - with httpx.Client( - timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120")) - ) as client: - return client.post( - f"{self.base_url.rstrip('/')}/v1/ocr", - headers={"Authorization": f"Bearer {self.master_key}"}, - json={"model": model, "document": document}, - ) - - -@dataclass(frozen=True) -class OcrResources: - gateway: OcrGateway - - -@pytest.fixture -def resources() -> OcrResources: - proxy_url = os.getenv("LITELLM_PROXY_URL") - if not proxy_url: - pytest.skip( - "Start a Rust OCR proxy and set LITELLM_PROXY_URL, e.g. http://localhost:4000" - ) - return OcrResources( - gateway=OcrGateway( - base_url=proxy_url, - master_key=os.getenv("LITELLM_MASTER_KEY", "sk-1234"), - ) - ) - - -def _assert_ocr_response_shape(response_json: dict[str, Any]) -> None: - assert response_json["object"] == "ocr" - assert response_json["model"] - assert isinstance(response_json["pages"], list) - assert len(response_json["pages"]) > 0 - assert "index" in response_json["pages"][0] - assert "markdown" in response_json["pages"][0] - - -class TestRustOcrGateway: - def test_rust_ocr_models_are_on_gateway_config(self) -> None: - config = yaml.safe_load(CONFIG_PATH.read_text()) - configured_models = { - model_config["model_name"] for model_config in config["model_list"] - } - - expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES} - assert expected_models.issubset(configured_models) - - def test_running_gateway_loaded_rust_ocr_models( - self, resources: OcrResources - ) -> None: - expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES} - assert expected_models.issubset(resources.gateway.model_names()) - - @pytest.mark.parametrize(("model", "document"), RUST_OCR_GATEWAY_CASES) - def test_rust_ocr_model_gateway_response( - self, resources: OcrResources, model: str, document: dict[str, str] - ) -> None: - response = resources.gateway.ocr(model, document) - - assert response.status_code == 200, response.text - _assert_ocr_response_shape(response.json()) diff --git a/tests/e2e/llm_translation/test_ocr_rust_e2e.py b/tests/e2e/llm_translation/test_ocr_rust_e2e.py new file mode 100644 index 00000000000..44908da2729 --- /dev/null +++ b/tests/e2e/llm_translation/test_ocr_rust_e2e.py @@ -0,0 +1,117 @@ +"""Live e2e: Rust-backed OCR is reachable through the gateway across providers. + +The gateway config declares one rust-ocr deployment per provider (mistral, +azure_ai, azure document intelligence, vertex mistral, vertex deepseek). Start the +proxy with the Rust OCR path enabled: + + LITELLM_USE_RUST_OCR=1 litellm --config tests/e2e/gateway/litellm-config.yml + +Three behaviors are checked: the config declares every provider's deployment (a +pure config read, no proxy needed); the running proxy loaded them onto /model/info; +and each one returns a well-formed OCR document over /v1/ocr. Per the e2e +"skip on environment, fail on behavior" rule, the proxy-backed cases skip when no +proxy answers but fail (never skip) once a request reaches it, so a provider whose +credentials are missing surfaces as a hard failure rather than silent green. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +import pytest +import yaml +from pydantic import BaseModel + +from e2e_http import unwrap +from models import OcrBody, OcrDocument, OcrResponse +from passthrough_client import PassthroughClient + +# Tiny in-repo fixtures served via jsdelivr (sha-pinned, immutable) so the request +# bodies stay stable across runs. +TEST_PDF_URL = ( + "https://cdn.jsdelivr.net/gh/BerriAI/litellm" + "@d769e81c90d453240c61fc572cdb27fae06a89d0" + "/tests/llm_translation/fixtures/dummy.pdf" +) +TEST_IMAGE_URL = ( + "https://cdn.jsdelivr.net/gh/BerriAI/litellm" + "@d769e81c90d453240c61fc572cdb27fae06a89d0" + "/tests/image_gen_tests/test_image.png" +) + +CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml" + + +@dataclass(frozen=True, slots=True) +class _OcrCase: + model: str + document: OcrDocument + + +RUST_OCR_CASES: tuple[_OcrCase, ...] = ( + _OcrCase( + "rust-ocr-mistral", + OcrDocument(type="document_url", document_url=TEST_PDF_URL), + ), + _OcrCase( + "rust-ocr-azure-ai", + OcrDocument(type="document_url", document_url=TEST_PDF_URL), + ), + _OcrCase( + "rust-ocr-azure-document-intelligence", + OcrDocument(type="document_url", document_url=TEST_PDF_URL), + ), + _OcrCase( + "rust-ocr-vertex-mistral", + OcrDocument(type="document_url", document_url=TEST_PDF_URL), + ), + _OcrCase( + "rust-ocr-vertex-deepseek", + OcrDocument(type="image_url", image_url=TEST_IMAGE_URL), + ), +) + +_EXPECTED_MODELS = frozenset(case.model for case in RUST_OCR_CASES) +_CASE_IDS = tuple(case.model.removeprefix("rust-ocr-") for case in RUST_OCR_CASES) + + +class _ConfiguredModel(BaseModel): + model_name: str + + +class _GatewayConfig(BaseModel): + model_list: list[_ConfiguredModel] + + +def _configured_model_names() -> frozenset[str]: + config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text())) + return frozenset(entry.model_name for entry in config.model_list) + + +def _assert_ocr_document(response: OcrResponse) -> None: + assert response.object == "ocr", f"expected object='ocr', got {response.object!r}" + assert response.model, "response missing the resolved model name" + assert response.pages, "OCR returned no pages" + assert response.pages[0].markdown is not None, "first page has no markdown" + + +def test_rust_ocr_models_declared_in_gateway_config() -> None: + """Pure config read (no proxy): every provider's rust-ocr deployment the suite + exercises is declared in the gateway config the proxy runs with. A case added + here without a matching deployment fails before any live call is attempted.""" + missing = _EXPECTED_MODELS - _configured_model_names() + assert not missing, f"rust-ocr models absent from {CONFIG_PATH.name}: {missing}" + + +@pytest.mark.e2e +class TestRustOcrGateway: + def test_gateway_loaded_rust_ocr_models(self, client: PassthroughClient) -> None: + loaded = frozenset(entry.model_name for entry in client.gateway.model_info()) + missing = _EXPECTED_MODELS - loaded + assert not missing, f"proxy did not load rust-ocr models: {missing}" + + @pytest.mark.parametrize("case", RUST_OCR_CASES, ids=_CASE_IDS) + def test_rust_ocr_response(self, client: PassthroughClient, scoped_key: str, case: _OcrCase) -> None: + response = unwrap(client.gateway.ocr(scoped_key, OcrBody(model=case.model, document=case.document))) + _assert_ocr_document(response) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 83b7d148957..0d0352df9af 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -125,6 +125,34 @@ class EmbedResponse(BaseModel): model: str | None = None +# ---------- ocr ---------- + + +class OcrDocument(BaseModel): + """A document for /v1/ocr in Mistral OCR format: a document_url for PDFs/docs + or an image_url for images. exclude_none on serialize drops the unset one.""" + + type: str + document_url: str | None = None + image_url: str | None = None + + +class OcrBody(BaseModel): + model: str + document: OcrDocument + + +class OcrPage(BaseModel): + index: int + markdown: str + + +class OcrResponse(BaseModel): + object: str | None = None + model: str | None = None + pages: list[OcrPage] = [] + + # ---------- spend logs ---------- @@ -216,14 +244,10 @@ class CustomPricing(BaseModel): def token_cost(self, prompt_tokens: int, completion_tokens: int) -> float: """Spend for a fresh (uncached) call under these rates: the proxy's custom-pricing formula (prompt * input + completion * output).""" - assert ( - self.input_cost_per_token is not None - and self.output_cost_per_token is not None - ), "custom pricing has no per-token rates" - return ( - prompt_tokens * self.input_cost_per_token - + completion_tokens * self.output_cost_per_token + assert self.input_cost_per_token is not None and self.output_cost_per_token is not None, ( + "custom pricing has no per-token rates" ) + return prompt_tokens * self.input_cost_per_token + completion_tokens * self.output_cost_per_token class ModelInfoEntry(BaseModel): From f796547d8061ac228c5f84a6c0e3e65420de6172 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 27 Jun 2026 17:54:40 -0700 Subject: [PATCH 02/37] docs(e2e): add CLAUDE.md harness conventions and coverage registry tests/e2e/CLAUDE.md captures the harness code-style rules (suite-as-a-class, shared transport, typed pydantic models, Result/unwrap, markers, typing) and the coverage-registry naming grammar; CONTRIBUTING.md gets the Contributors Guide intro and a Setup section --- tests/e2e/CLAUDE.md | 132 ++++++++++++++++++++++++++++++++++++++ tests/e2e/CONTRIBUTING.md | 12 +++- 2 files changed, 142 insertions(+), 2 deletions(-) create mode 100644 tests/e2e/CLAUDE.md diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md new file mode 100644 index 00000000000..2c9dadb0ff2 --- /dev/null +++ b/tests/e2e/CLAUDE.md @@ -0,0 +1,132 @@ +# e2e harness conventions + +Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `CLAUDE.md` + +## Lay the pattern down in a class + +Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being + +```python +# pseudo-code to convey intent, not the real API +class TestPromptCompression: + def test_prompt_compression_add_to_virtual_key(self): + new_key = self.resources.create_key(user_id, compression=True) # turn the feature on + resources._defer(new_key) # queue key deletion + + def test_prompt_compression_accumulate_spend(self, key_id, user_id): + for _ in range(10): + response = self.resources.gateway.post("gemini-2.5-flash", key_id, user_id) + compressed_value = ... + assert response.cost == compressed_value # the cost was actually reduced +``` + +That snippet only conveys intent. What you actually write uses the real harness: the `client` fixture for your suite, the `scoped_key` fixture for an auto-deleted key, typed pydantic bodies from `models.py`, and `unwrap(...)` on the tagged-union result. `tests/e2e/llm_translation/test_custom_pricing_e2e.py` is the reference to copy from; it creates a scoped key, drives a real gemini call, polls `/spend/logs` to a deadline for the cost-breakdown row, then asserts the input and output costs match the configured custom rates and that a sibling deployment kept its own price. Read it before writing yours + +## Use the shared transport; never touch requests directly + +Every HTTP call goes through the shared transport, never through `requests.*` in a test. `e2e_http.py` is the only module permitted to call `requests.*`, and that is enforced in CI by `tests/code_coverage_tests/check_e2e_no_raw_requests.py`. A test that imports requests will fail the check + +The shape is layered so tests stay declarative + +`transport.py` exposes a `Transport` Protocol with `post`, `get`, `delete`, `send`, `stream`, `probe`, plus `bearer(key)` and the `master` header. `HttpTransport` fulfils it, and `SplitTransport` routes each call by path to the data plane or the control plane so a split control-plane/data-plane deployment works without any change in the test + +`e2e_gateway.py` holds `Gateway`, a frozen dataclass that wraps a `Transport` and adds the operations tests reuse: `generate_key` / `delete_key` / `key_info`, `model_info`, the LLM calls `chat` / `chat_stream` / `embed` / `ocr`, the spend read-back `spend_logs`, and the poll helpers `poll_logs_for_key` / `poll_logs_for_request_id` that loop to `poll_timeout` instead of sleeping once. Add a new route as a method here so other suites get it for free + +Each suite provides its own `client` fixture (see `llm_translation/passthrough_client.py`), a frozen dataclass that holds the shared `Gateway` and adds suite-specific routes. Cleanup runs through that same `Gateway`, so whatever keys or customers your test creates get torn down by the `resources` fixture + +Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip + +Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache + +## Typing + +The harness is fully typed and new code must not add `Any` or widen the basedpyright budgets. When a response field is untyped, model it in `models.py` (just the fields you read) and let pydantic validate it, rather than threading a `dict` or `Any` through the test + +## Coverage registry + +The set of tests we want is a registry checked into this repo, one row per behavior; that file is the definition of done and the denominator. Each e2e test declares what it covers with `@pytest.mark.covers("...")`, and a small collector diffs the registry against the tests and ships coverage to the existing Grafana. No Allure, no new dependencies + +Coverage is organized as module > feature > test. There are six modules: LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works` + +The metric is coverage: the share of registry rows that have a passing covering test, reported to Grafana per module so a gap surfaces as an uncovered row rather than a silent absence + +### Naming grammar per module + +LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix + +``` +llm..... + endpoint : chat_completions | messages | responses | embeddings | batches | files + | rerank | images_generations | audio_speech | audio_transcriptions | moderations + route : openai | azure_openai | anthropic | bedrock_invoke | bedrock_converse | vertex | azure_foundry + (vocab varies per endpoint; messages is anthropic-format only) + capability : basic | tool_use | prompt_cache_5m | prompt_cache_1h | vision | thinking + | thinking_tool_use | pdf_input | web_search | structured_output | count_tokens + | tool_search | long_context_1m + streaming : stream | nonstream (omit where n/a) + assertion : works | cost_logged + label (not in id): model = haiku-4.5 | sonnet-4.6 | opus-4.7 | gpt-* + e.g. llm.chat_completions.bedrock_converse.tool_use.stream.works + llm.messages.anthropic.prompt_cache_1h.nonstream.cache_hit +``` + +Management / UI - endpoint features (surface tag: api | ui) + +``` +mgmt.. + endpoint : key.generate | key.update | key.delete | team.new | user.new + | budget.new | model.add | ... (one per management route) + assertion : persists | member_forbidden | admin_only | happy_path + e.g. mgmt.key.generate.persists (surface=api) + mgmt.key.generate.happy_path (surface=ui) +``` + +MCPs - endpoint features with the protocol op as the variant + +``` +mcp... + operation : list_tools | call_tool | list_resources | read_resource | list_prompts | get_prompt + auth_family : none | api_key | bearer | oauth + assertion : succeeds | denied_without_permission + e.g. mcp.call_tool.oauth.succeeds +``` + +Reliability & Performance - behavior features (no route; endpoint is exercised_on) + +``` +reliability... + behavior : fallback | retry | cooldown | timeout | ratelimit | routing | cache | circuit_breaker | perf + variant : 5xx | context_window | content_policy | 429 | timeout + simple_shuffle | usage_based | latency_based | cost_based | least_busy + latency | throughput (perf only; SLO/threshold assertion, not binary) + assertion : routes_to_fallback | succeeds_within_retries | picks_under_tpm | returns_cached + | trips_then_recovers | under_slo + e.g. reliability.fallback.context_window.routes_to_fallback exercised_on=[chat_completions] + reliability.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages] +``` + +Logging & Guardrails - behavior features (config-driven; endpoint is exercised_on) + +``` +logging... + integration : langfuse | s3 | otel | prometheus | datadog | ... + event : success | failure | stream + assertion : logs_spend | writes_object | exports_metric + e.g. logging.langfuse.success.logs_spend exercised_on=[chat_completions] + +guardrail... + provider : presidio | lakera | bedrock | aporia | ... + hook_point : pre_call | post_call | during | logging_only + assertion : blocks | masks | allows + e.g. guardrail.presidio.pre_call.masks exercised_on=[chat_completions] +``` + +Other - holding pen (endpoint or behavior) + +``` +other... + area : auth | lifecycle | config | ... + rule : audited periodically; a cluster here promotes to a new component + e.g. other.auth.jwt.valid_token_allows + other.lifecycle.readiness.reports_db +``` diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index c0c08453b25..494ba315a1e 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -1,8 +1,16 @@ -# Contributing e2e tests +# Contributors Guide This directory holds the live end-to-end suites that prove product correctness against a real running proxy and real provider APIs. The goal of this guide is simple: when you ship a feature, you add e2e coverage that walks that feature the way production does, across every route and edge case it touches, so a later change that breaks it fails here first -Read this before adding a test. The harness already encodes most of the rules; your job is to fill in the feature-specific behavior, not to reinvent plumbing +Read this before adding a test and i recommend reading through CLAUDE.md + +When contributing to this directory, please first discuss the change you wish to make via issue or pull request. We require screenshots and proof of your tests working on a live proxy. + + +## Setup + +1. Use litellm-config.yml in gateway/ if your tests require adding new models +2. Deploy a docker image in docker-compose.yml to run your e2e tests ## What a complete test looks like From 1b45ee629a8ebd93e0f30cb01d7d464873069b6a Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 27 Jun 2026 19:30:55 -0700 Subject: [PATCH 03/37] fix: make changes to contributing --- tests/e2e/CONTRIBUTING.md | 33 ++++++++++++++++++++++++++++++--- 1 file changed, 30 insertions(+), 3 deletions(-) diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 494ba315a1e..12aea6bbc25 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -9,8 +9,35 @@ When contributing to this directory, please first discuss the change you wish to ## Setup -1. Use litellm-config.yml in gateway/ if your tests require adding new models -2. Deploy a docker image in docker-compose.yml to run your e2e tests +The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with its Postgres and Redis, serving `gateway/litellm-config.yml`; add any model, pricing override, or guardrail your test needs to that file and read it back in the test rather than hardcoding values. `gateway/` holds proxy configuration only, so never put tests there + +## Running the tests locally + +1. Create a .env file and add provider keys: + ```bash + OPENAI_API_KEY="sk-..." + ANTHROPIC_API_KEY="sk-..." + +2. Bring the stack up from this directory: + + ```bash + docker compose up -d + curl -fs http://localhost:4000/health/liveliness + ``` + +3. Run a suite against it; the harness reads `LITELLM_PROXY_URL` (default `http://localhost:4000`): + + ```bash + uv run pytest tests/e2e/llm_translation/ -v + ``` + +4. Tear it down when you're done: + + ```bash + docker compose down -v + ``` + +Tests marked `@pytest.mark.e2e` skip when no proxy answers `/health/liveliness`, so a run that reports everything skipped means the stack isn't up, not that anything passed ## What a complete test looks like @@ -61,7 +88,7 @@ If a step is missing, the test is not done. That is the whole pattern Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being ```python -# pseudo-code to convey intent, not the real API +# pseudo-code to convey intent class TestPromptCompression: def test_prompt_compression_add_to_virtual_key(self): new_key = self.resources.create_key(user_id, compression=True) # turn the feature on From d5f757fc735c66642fa007d876d453243a93d790 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 27 Jun 2026 19:38:57 -0700 Subject: [PATCH 04/37] docs(e2e): document suite-folder layout and the add-a-folder rule in CLAUDE.md --- tests/e2e/CLAUDE.md | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 2c9dadb0ff2..b155a1a7024 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -2,6 +2,22 @@ Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `CLAUDE.md` +## Suite folders + +Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family or behavior area. If you add a new folder, you must add a line here describing what kind of tests belong in it, so the layout stays self-describing. `gateway/` is the exception: it holds proxy configuration only and never tests + +- `llm_translation/` - LLM endpoint and provider-translation behavior: passthrough, custom pricing, OCR +- `embeddings/` - the `/embeddings` endpoint across providers +- `batches/` - the `/batches` endpoint (placeholder until the first test lands) +- `realtime/` - realtime websocket sessions, including the pipecat audio path +- `budgets/` - budget definition, enforcement, and reset windows (key, team, tag, soft, multi-window) +- `spend_tracking/` - spend logging and cost attribution on `/spend/*` +- `models_mgmt/` - model-management routes (add/update, tpm persistence) +- `logging/` - logging-integration delivery (datadog and friends) +- `security/` - secret handling and log-leak protection +- `router/` - routing and reliability behavior (rate limits, fallbacks, cooldowns) +- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests + ## Lay the pattern down in a class Keep the cases for one feature inside a class so the file reads as a spec for how that feature behaves in production. The class name says what is under test; each method is one behavior. Think of it as documenting the contract, with the rough intent being From 2a5790fe55d1f73846b43ec4bce4c0fd84261f17 Mon Sep 17 00:00:00 2001 From: yucheng-berriai Date: Fri, 26 Jun 2026 12:53:07 -0700 Subject: [PATCH 05/37] fix(proxy): reject team-scoped object_permission on personal keys for non-admins Non-admin callers could create or update a personal key (no team_id) with arbitrary access_group_ids, mcp_toolsets, vector_stores, or search_tools in object_permission. The server persisted the values without ownership validation; runtime authorization then trusted the IDs because they were stored on the key, allowing cross-tenant access to other teams' restricted models, MCP toolsets, and vector stores. The personal-key gate now mirrors the team-key path. enforce_member_can_assign_access_groups raises 403 for non-admin teamless callers. validate_key_mcp_servers_against_team rejects non-empty mcp_toolsets on personal non-admin keys. A new validate_key_vector_stores_against_team enforces the same rule for vector_stores. validate_key_search_tools_against_team gains the same gate for search_tools. The four validators are wired into /key/generate, /key/update, and /key/regenerate. Proxy admins keep their existing carve-out across all fields; team keys are unaffected. Endpoint-level regression coverage lives in tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py (six new parametrised cases through generate_key_fn and _validate_update_key_data) and helper-level coverage in tests/test_litellm/proxy/management_helpers/. Deleting any of the validator calls in _common_key_generation_helper or unmoving the enforce gate in _validate_update_key_data breaks the suite. --- .../key_management_endpoints.py | 74 +++++-- .../object_permission_utils.py | 123 +++++++++-- .../team_member_permission_checks.py | 10 +- .../test_key_management_endpoints.py | 191 ++++++++++++++++++ .../test_object_permission_utils.py | 120 +++++++++++ .../test_team_member_permission_checks.py | 39 +++- 6 files changed, 512 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index eed7d869a5d..f19ea6da529 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -80,6 +80,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, + validate_key_vector_stores_against_team, ) from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, @@ -349,6 +350,12 @@ def _personal_key_membership_check( def _personal_key_generation_check(user_api_key_dict: UserAPIKeyAuth, data: GenerateKeyRequest): + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=user_api_key_dict, + team_table=None, + access_group_ids=data.access_group_ids, + ) + if ( litellm.key_generation_settings is None or litellm.key_generation_settings.get("personal_key_generation") is None @@ -845,17 +852,26 @@ async def _common_key_generation_helper( data_json.pop("tags") # Validate MCP servers in object_permission are within team scope + _is_proxy_admin_caller = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) normalized_object_permission = await validate_key_mcp_servers_against_team( object_permission=data_json.get("object_permission"), team_obj=team_table, prisma_client=prisma_client, - is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value, + is_proxy_admin=_is_proxy_admin_caller, ) if normalized_object_permission is not None: data_json["object_permission"] = normalized_object_permission await validate_key_search_tools_against_team( object_permission=data_json.get("object_permission"), team_obj=team_table, + is_proxy_admin=_is_proxy_admin_caller, + ) + await validate_key_vector_stores_against_team( + object_permission=data_json.get("object_permission"), + team_obj=team_table, + is_proxy_admin=_is_proxy_admin_caller, ) data_json = await _set_object_permission( @@ -2069,13 +2085,7 @@ async def _validate_mcp_servers_for_key_update( user_api_key_cache=user_api_key_cache, check_db_only=True, ) - object_permission_dict: Optional[dict] = None - if data.object_permission is not None: - object_permission_dict = ( - data.object_permission.model_dump(exclude_unset=True) - if hasattr(data.object_permission, "model_dump") - else dict(data.object_permission) # type: ignore[arg-type] - ) + object_permission_dict = _object_permission_to_dict(data.object_permission) normalized_object_permission = await validate_key_mcp_servers_against_team( object_permission=object_permission_dict, team_obj=effective_team_obj, @@ -2085,6 +2095,12 @@ async def _validate_mcp_servers_for_key_update( await validate_key_search_tools_against_team( object_permission=object_permission_dict, team_obj=effective_team_obj, + is_proxy_admin=is_proxy_admin, + ) + await validate_key_vector_stores_against_team( + object_permission=object_permission_dict, + team_obj=effective_team_obj, + is_proxy_admin=is_proxy_admin, ) return normalized_object_permission @@ -2216,14 +2232,6 @@ async def _validate_update_key_data( detail=f"Team not found for team_id={data.team_id}. Non-admin users cannot set keys to non-existent teams.", ) - # Field-level opt-in: non-admin members may only assign access groups when - # the team has enabled KEY_ACCESS_GROUP_ASSIGNMENT. - TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( - user_api_key_dict=user_api_key_dict, - team_table=team_obj, - access_group_ids=data.access_group_ids, - ) - if team_obj is not None: await _check_team_key_limits( team_table=team_obj, @@ -2231,6 +2239,12 @@ async def _validate_update_key_data( prisma_client=prisma_client, ) + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=user_api_key_dict, + team_table=team_obj, + access_group_ids=data.access_group_ids, + ) + # Validate key against project limits if project_id is being set _project_id_to_check = getattr(data, "project_id", None) or getattr(existing_key_row, "project_id", None) if _project_id_to_check is not None and (data.models is not None or data.max_budget is not None): @@ -4458,9 +4472,9 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) - # Gate access_group_ids on regenerate, same as /key/generate and - # /key/update. Use the existing key's team since the body may omit it. - if data is not None and data.access_group_ids: + if data is not None and ( + data.access_group_ids or data.object_permission is not None + ): regenerate_team_table: Optional[LiteLLM_TeamTableCachedObj] = None if _key_in_db.team_id is not None: regenerate_team_table = await get_team_object( @@ -4469,11 +4483,33 @@ async def regenerate_key_fn( user_api_key_cache=user_api_key_cache, check_db_only=True, ) + _regen_is_proxy_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=user_api_key_dict, team_table=regenerate_team_table, access_group_ids=data.access_group_ids, ) + _regen_object_permission_dict = _object_permission_to_dict( + data.object_permission + ) + await validate_key_mcp_servers_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + prisma_client=prisma_client, + is_proxy_admin=_regen_is_proxy_admin, + ) + await validate_key_search_tools_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + is_proxy_admin=_regen_is_proxy_admin, + ) + await validate_key_vector_stores_against_team( + object_permission=_regen_object_permission_dict, + team_obj=regenerate_team_table, + is_proxy_admin=_regen_is_proxy_admin, + ) verbose_proxy_logger.info( "Key regeneration requested: key_alias=%s", diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 9d5f716033f..d980a6f8cfd 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -555,31 +555,98 @@ async def validate_key_mcp_servers_against_team( detail={"error": detail}, ) - # Validate requested toolsets against team's allowed toolsets. - # Only enforce the team-based restriction when a team is present — standalone - # keys (no team) can freely be granted any toolset by an admin. - if requested_toolsets and team_obj is not None: - team_op = team_obj.object_permission - team_mcp_toolsets = team_op.mcp_toolsets if team_op is not None else None - # None or [] means the team has no toolset restriction — allow any toolsets. - if team_mcp_toolsets: - disallowed_toolsets = requested_toolsets - set(team_mcp_toolsets) - if disallowed_toolsets: - team_id = team_obj.team_id - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={ - "error": ( - f"Key requests MCP toolsets not allowed by team '{team_id}': " - f"{sorted(disallowed_toolsets)}. " - f"Team allows: {sorted(team_mcp_toolsets)}." - ) - }, - ) + _validate_requested_toolsets( + requested_toolsets=requested_toolsets, + team_obj=team_obj, + is_proxy_admin=is_proxy_admin, + ) return object_permission +def _validate_requested_toolsets( + requested_toolsets: set[str], + team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool, +) -> None: + """ + Validate mcp_toolsets requested on a key. + + Non-admin callers cannot assign toolsets to a personal (no team) key. Team + keys must request a subset of the team's own toolset allowlist. + """ + if not requested_toolsets: + return + if team_obj is None: + if is_proxy_admin: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. MCP toolsets cannot be assigned to " + "personal keys by non-admin callers. Disallowed toolsets: " + f"{sorted(requested_toolsets)}." + ) + }, + ) + team_op = team_obj.object_permission + team_mcp_toolsets = team_op.mcp_toolsets if team_op is not None else None + if not team_mcp_toolsets: + return + disallowed_toolsets = requested_toolsets - set(team_mcp_toolsets) + if not disallowed_toolsets: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + f"Key requests MCP toolsets not allowed by team '{team_obj.team_id}': " + f"{sorted(disallowed_toolsets)}. " + f"Team allows: {sorted(team_mcp_toolsets)}." + ) + }, + ) + + +def _extract_requested_vector_stores(object_permission: Optional[dict]) -> set[str]: + """Return vector_store IDs from a key's object_permission dict.""" + if not object_permission or not isinstance(object_permission, dict): + return set() + raw = object_permission.get("vector_stores") + if isinstance(raw, list): + return {str(x) for x in raw if x} + return set() + + +async def validate_key_vector_stores_against_team( + object_permission: Optional[dict], + team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool = False, +) -> None: + """ + Reject vector_stores requested on a personal (no team) key by a non-admin + caller. Vector store access is granted at use-time from the key's + object_permission.vector_stores list, so the assignment is the authorization + boundary. Team keys and proxy admins are unaffected. + """ + requested = _extract_requested_vector_stores(object_permission) + if not requested: + return + if team_obj is not None or is_proxy_admin: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. Vector stores cannot be assigned to " + "personal keys by non-admin callers. Disallowed vector stores: " + f"{sorted(requested)}." + ) + }, + ) + + def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[str]: """Return search_tool_name values from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -593,16 +660,30 @@ def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[s async def validate_key_search_tools_against_team( object_permission: Optional[dict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], + is_proxy_admin: bool = False, ) -> None: """ Validate key object_permission.search_tools is a subset of the team's allowlist. Empty team allowlist means no restriction at team layer (skip). + Non-admin callers cannot assign search_tools to a personal (no team) key. """ requested = _extract_requested_search_tools(object_permission) if not requested: return + if team_obj is None and not is_proxy_admin: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Key is not in a team. search_tools cannot be assigned to " + "personal keys by non-admin callers. Disallowed search tools: " + f"{sorted(requested)}." + ) + }, + ) + team_tools: List[str] = [] if team_obj is not None and team_obj.object_permission is not None: st = team_obj.object_permission.search_tools diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index 1353b9ed651..1532668ed19 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -167,9 +167,15 @@ class TeamMemberPermissionChecks: if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return - # Personal (non-team) keys are out of scope for team-member gating. if team_table is None: - return + raise HTTPException( + status_code=403, + detail=( + "Key is not in a team. Access groups cannot be assigned to " + "personal keys by non-admin callers. Disallowed access groups: " + f"{sorted(access_group_ids)}." + ), + ) team_member_object = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index a92eaa3a5f6..7f78e0de9a6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -526,6 +526,197 @@ async def test_key_generation_with_object_permission(monkeypatch): assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "field,request_kwargs,expected_in_error", + [ + ( + "access_group_ids", + {"access_group_ids": ["acme_private"]}, + "Access groups", + ), + ( + "mcp_toolsets", + {"object_permission": {"mcp_toolsets": ["acme_toolset"]}}, + "MCP toolsets", + ), + ( + "vector_stores", + {"object_permission": {"vector_stores": ["acme_vs"]}}, + "Vector stores", + ), + ( + "search_tools", + {"object_permission": {"search_tools": ["acme_search"]}}, + "search_tools", + ), + ], +) +async def test_generate_key_personal_non_admin_denied_for_team_scoped_fields( + monkeypatch, field, request_kwargs, expected_in_error +): + """generate_key_fn must reject access_group_ids and + object_permission.{mcp_toolsets, vector_stores, search_tools} when the + caller is a non-admin and the request has no team_id. Mutating any of the + three validator calls in _common_key_generation_helper or unmoving the + enforce_member_can_assign_access_groups call in _personal_key_generation_check + must break this test.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( + return_value=MagicMock(object_permission_id="should-not-create") + ) + mock_prisma_client.insert_data = AsyncMock(return_value=MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + from litellm.proxy._types import ( + GenerateKeyRequest, + LiteLLM_ObjectPermissionBase, + LitellmUserRoles, + ) + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_fn, + ) + + if "object_permission" in request_kwargs: + request_kwargs = { + **request_kwargs, + "object_permission": LiteLLM_ObjectPermissionBase( + **request_kwargs["object_permission"] + ), + } + request_data = GenerateKeyRequest(**request_kwargs) + + from litellm.proxy._types import ProxyException + + with pytest.raises((HTTPException, ProxyException)) as exc: + await generate_key_fn( + data=request_data, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + ) + code = getattr(exc.value, "status_code", None) or getattr(exc.value, "code", None) + assert int(code) == 403 + body = str( + getattr(exc.value, "detail", None) or getattr(exc.value, "message", exc.value) + ) + assert expected_in_error in body + mock_prisma_client.db.litellm_objectpermissiontable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_personal_non_admin_denied_vector_stores(monkeypatch): + """/key/update must reject vector_stores on a personal key by a non-admin. + Reverting the enforce_member_can_assign_access_groups move (i.e. putting + it back inside `if _team_id_to_check is not None`) does NOT cover + object_permission fields; this test exercises _validate_update_key_data + which calls _validate_mcp_servers_for_key_update.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + mock_prisma_client.db = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(), + ) + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LitellmUserRoles, + UpdateKeyRequest, + ) + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_update_key_data, + ) + + existing_key_row = MagicMock( + token="hashed_alice_personal_key", + user_id="alice", + team_id=None, + created_by="alice", + max_budget=None, + organization_id=None, + project_id=None, + ) + data = UpdateKeyRequest( + key="sk-alice-personal", + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["acme_vs"]), + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "Vector stores" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_update_key_personal_non_admin_denied_access_groups( + monkeypatch, +): + """/key/update on a personal key must also gate access_group_ids for + non-admins. Reverting the enforce move (putting it back inside + `if _team_id_to_check is not None`) breaks this test.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data # type: ignore + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + from litellm.proxy._types import LitellmUserRoles, UpdateKeyRequest + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_update_key_data, + ) + + existing_key_row = MagicMock( + token="hashed_alice_personal_key", + user_id="alice", + team_id=None, + created_by="alice", + max_budget=None, + organization_id=None, + project_id=None, + ) + data = UpdateKeyRequest( + key="sk-alice-personal", + access_group_ids=["ag-private"], + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "Access groups" in str(exc.value.detail) + + @pytest.mark.asyncio async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch): """Ensure generate_key_helper_fn passes access_group_ids into the key insert payload.""" diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 2b38d732e9d..d81511d2322 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -18,6 +18,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, + validate_key_vector_stores_against_team, ) @@ -890,3 +891,122 @@ async def test_validate_search_tools_raises_when_not_subset(): team_obj=_make_team_obj_search(search_tools=["t1"]), ) assert exc.value.status_code == 403 + + +# ---- Personal-key non-admin gates on toolsets / vector_stores / search_tools ---- + + +@pytest.mark.asyncio +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], +) +async def test_personal_non_admin_cannot_assign_mcp_toolsets( + mock_access_groups, mock_allow_all +): + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission={"mcp_toolsets": ["ts-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "ts-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +@patch( + "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + return_value=set(), +) +@patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], +) +async def test_personal_admin_can_assign_mcp_toolsets( + mock_access_groups, mock_allow_all +): + await validate_key_mcp_servers_against_team( + object_permission={"mcp_toolsets": ["ts-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_personal_non_admin_cannot_assign_vector_stores(): + with pytest.raises(HTTPException) as exc: + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "vs-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_personal_admin_can_assign_vector_stores(): + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_team_key_vector_stores_unrestricted_at_create(): + """Team-scoped keys retain their existing trust model at create time.""" + team_obj = _make_team_obj_search() + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": ["vs-anything"]}, + team_obj=team_obj, + is_proxy_admin=False, + ) + + +@pytest.mark.asyncio +async def test_personal_non_admin_cannot_assign_search_tools(): + with pytest.raises(HTTPException) as exc: + await validate_key_search_tools_against_team( + object_permission={"search_tools": ["st-private"]}, + team_obj=None, + is_proxy_admin=False, + ) + assert exc.value.status_code == 403 + assert "st-private" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_personal_admin_can_assign_search_tools(): + await validate_key_search_tools_against_team( + object_permission={"search_tools": ["st-private"]}, + team_obj=None, + is_proxy_admin=True, + ) + + +@pytest.mark.asyncio +async def test_empty_object_permission_passes_for_personal_non_admin(): + """An empty / absent object_permission must not be blocked.""" + await validate_key_vector_stores_against_team( + object_permission=None, + team_obj=None, + is_proxy_admin=False, + ) + await validate_key_vector_stores_against_team( + object_permission={"vector_stores": []}, + team_obj=None, + is_proxy_admin=False, + ) + await validate_key_search_tools_against_team( + object_permission=None, + team_obj=None, + is_proxy_admin=False, + ) diff --git a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py index 29aa75a0f0a..71999e29f96 100644 --- a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py +++ b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py @@ -314,12 +314,45 @@ class TestEnforceMemberCanAssignAccessGroups: access_group_ids=["ag-1"], ) - def test_personal_key_out_of_scope(self): - """Personal (non-team) keys are not gated by team-member permissions.""" + def test_personal_key_non_admin_denied(self): + """A non-admin cannot self-grant access_group_ids on a personal (no + team) key. The access_group_id grants model access at use-time + without any team-membership cross-check, so the assignment is the + authorization boundary.""" + from fastapi import HTTPException + + with pytest.raises(HTTPException) as exc: + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(), + team_table=None, + access_group_ids=["ag-private"], + ) + assert exc.value.status_code == 403 + assert "ag-private" in str(exc.value.detail) + + def test_personal_key_proxy_admin_can_assign(self): + """Proxy admins bypass the personal-key gate and may assign access + groups on personal keys.""" + from litellm.proxy._types import LitellmUserRoles + + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(role=LitellmUserRoles.PROXY_ADMIN.value), + team_table=None, + access_group_ids=["ag-private"], + ) + + def test_personal_key_empty_access_groups_passes(self): + """An empty / absent access_group_ids list must not be rejected even + on a personal key — the gate only fires when the field is non-empty.""" TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=self._user(), team_table=None, - access_group_ids=["ag-1"], + access_group_ids=None, + ) + TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( + user_api_key_dict=self._user(), + team_table=None, + access_group_ids=[], ) def test_team_admin_bypasses(self, monkeypatch): From f2d7cb152adba50a561fb72b6dde4f7ce9c96913 Mon Sep 17 00:00:00 2001 From: yucheng-berriai Date: Fri, 26 Jun 2026 16:33:52 -0700 Subject: [PATCH 06/37] refactor(proxy): type object_permission dict with ObjectPermissionDict Replace bare Optional[dict] on the object_permission validator surfaces with a typed TypedDict mirror of LiteLLM_ObjectPermissionBase. The TypedDict shape matches the Pydantic model field-for-field and supports .get() and item assignment, so the mutation in _rewrite_object_permission_mcp_identifiers continues to work at runtime (TypedDict is a plain dict). Propagated through the surfaces this PR touches: _object_permission_to_dict, _validate_mcp_servers_for_key_update, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against _team, the five _extract_requested_* helpers, and the two _rewrite_object_permission_mcp_* mutators. attach_object_permission_to_dict, handle_update_object_permission_common, and _set_object_permission keep their wider dict typing because they handle the full key/team data_json, which is a superset of ObjectPermissionDict and pre-dates this PR. No behavior change. 373 tests pass; ruff strict + type discipline gates green. --- litellm/proxy/_types.py | 17 +++++++++++ .../key_management_endpoints.py | 26 ++++++++-------- .../object_permission_utils.py | 30 +++++++++++-------- .../test_object_permission_utils.py | 17 ++++++++++- 4 files changed, 63 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d84588a4c24..7d470fba08b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1005,6 +1005,23 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): search_tools: Optional[List[str]] = None +class ObjectPermissionDict(TypedDict, total=False): + """Plain-dict mirror of LiteLLM_ObjectPermissionBase used by validators + that need to mutate the payload before persistence (e.g. MCP server + identifier normalization in object_permission_utils).""" + + mcp_servers: Optional[list[str]] + mcp_access_groups: Optional[list[str]] + mcp_tool_permissions: Optional[dict[str, list[str]]] + mcp_toolsets: Optional[list[str]] + blocked_tools: Optional[list[str]] + vector_stores: Optional[list[str]] + agents: Optional[list[str]] + agent_access_groups: Optional[list[str]] + models: Optional[list[str]] + search_tools: Optional[list[str]] + + from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index f19ea6da529..aeb664a5d1f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -349,6 +349,14 @@ def _personal_key_membership_check( return True +def _object_permission_to_dict( + object_permission: Optional[LiteLLM_ObjectPermissionBase], +) -> Optional[ObjectPermissionDict]: + if object_permission is None: + return None + return cast(ObjectPermissionDict, object_permission.model_dump(exclude_unset=True)) + + def _personal_key_generation_check(user_api_key_dict: UserAPIKeyAuth, data: GenerateKeyRequest): TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=user_api_key_dict, @@ -852,9 +860,7 @@ async def _common_key_generation_helper( data_json.pop("tags") # Validate MCP servers in object_permission are within team scope - _is_proxy_admin_caller = ( - user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value - ) + _is_proxy_admin_caller = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value normalized_object_permission = await validate_key_mcp_servers_against_team( object_permission=data_json.get("object_permission"), team_obj=team_table, @@ -2074,7 +2080,7 @@ async def _validate_mcp_servers_for_key_update( prisma_client: Any, user_api_key_cache: Any, is_proxy_admin: bool, -) -> Optional[dict]: +) -> Optional[ObjectPermissionDict]: """Validate MCP servers in object_permission against the effective team.""" effective_team_obj = team_obj # If team_id isn't being changed, resolve the existing key's team @@ -4472,9 +4478,7 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) - if data is not None and ( - data.access_group_ids or data.object_permission is not None - ): + if data is not None and (data.access_group_ids or data.object_permission is not None): regenerate_team_table: Optional[LiteLLM_TeamTableCachedObj] = None if _key_in_db.team_id is not None: regenerate_team_table = await get_team_object( @@ -4483,17 +4487,13 @@ async def regenerate_key_fn( user_api_key_cache=user_api_key_cache, check_db_only=True, ) - _regen_is_proxy_admin = ( - user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value - ) + _regen_is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value TeamMemberPermissionChecks.enforce_member_can_assign_access_groups( user_api_key_dict=user_api_key_dict, team_table=regenerate_team_table, access_group_ids=data.access_group_ids, ) - _regen_object_permission_dict = _object_permission_to_dict( - data.object_permission - ) + _regen_object_permission_dict = _object_permission_to_dict(data.object_permission) await validate_key_mcp_servers_against_team( object_permission=_regen_object_permission_dict, team_obj=regenerate_team_table, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index d980a6f8cfd..fe96d9c260a 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -11,7 +11,7 @@ from fastapi import HTTPException, status from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import SpecialMCPServerNames +from litellm.proxy._types import ObjectPermissionDict, SpecialMCPServerNames from litellm.proxy.utils import PrismaClient from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.table_repositories import MCPServerRepository @@ -257,7 +257,7 @@ async def _resolve_mcp_server_identifiers_to_ids( def _rewrite_object_permission_mcp_servers( - object_permission: dict, + object_permission: ObjectPermissionDict, identifier_to_server_ids: Dict[str, Set[str]], ) -> None: mcp_servers = object_permission.get("mcp_servers") @@ -274,7 +274,7 @@ def _rewrite_object_permission_mcp_servers( def _rewrite_object_permission_mcp_tool_permissions( - object_permission: dict, + object_permission: ObjectPermissionDict, identifier_to_server_ids: Dict[str, Set[str]], ) -> None: mcp_tool_permissions = object_permission.get("mcp_tool_permissions") @@ -295,7 +295,7 @@ def _rewrite_object_permission_mcp_tool_permissions( def _rewrite_object_permission_mcp_identifiers( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], identifier_to_server_ids: Dict[str, Set[str]], ) -> None: if not object_permission or not isinstance(object_permission, dict): @@ -383,7 +383,7 @@ async def _get_team_allowed_mcp_servers( def _extract_requested_mcp_server_ids( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """ Extract all MCP server IDs referenced in a key's object_permission dict. @@ -409,7 +409,7 @@ def _extract_requested_mcp_server_ids( def _extract_requested_mcp_access_groups( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """Extract MCP access groups from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -422,7 +422,7 @@ def _extract_requested_mcp_access_groups( def _extract_requested_mcp_toolsets( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], ) -> Set[str]: """Extract MCP toolset IDs from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): @@ -435,11 +435,11 @@ def _extract_requested_mcp_toolsets( async def validate_key_mcp_servers_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], prisma_client: Optional[PrismaClient] = None, is_proxy_admin: bool = False, -) -> Optional[dict]: +) -> Optional[ObjectPermissionDict]: """ Validate that MCP servers requested on a key are within the allowed scope. @@ -609,7 +609,9 @@ def _validate_requested_toolsets( ) -def _extract_requested_vector_stores(object_permission: Optional[dict]) -> set[str]: +def _extract_requested_vector_stores( + object_permission: Optional[ObjectPermissionDict], +) -> set[str]: """Return vector_store IDs from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): return set() @@ -620,7 +622,7 @@ def _extract_requested_vector_stores(object_permission: Optional[dict]) -> set[s async def validate_key_vector_stores_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], is_proxy_admin: bool = False, ) -> None: @@ -647,7 +649,9 @@ async def validate_key_vector_stores_against_team( ) -def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[str]: +def _extract_requested_search_tools( + object_permission: Optional[ObjectPermissionDict], +) -> list[str]: """Return search_tool_name values from a key's object_permission dict.""" if not object_permission or not isinstance(object_permission, dict): return [] @@ -658,7 +662,7 @@ def _extract_requested_search_tools(object_permission: Optional[dict]) -> List[s async def validate_key_search_tools_against_team( - object_permission: Optional[dict], + object_permission: Optional[ObjectPermissionDict], team_obj: Optional["LiteLLM_TeamTableCachedObj"], is_proxy_admin: bool = False, ) -> None: diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index d81511d2322..26c8c774812 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -9,7 +9,7 @@ sys.path.insert(0, os.path.abspath("../../../..")) from unittest.mock import AsyncMock, MagicMock, patch -from litellm.proxy._types import LiteLLM_ObjectPermissionTable +from litellm.proxy._types import LiteLLM_ObjectPermissionBase, LiteLLM_ObjectPermissionTable, ObjectPermissionDict from litellm.proxy.management_helpers.object_permission_utils import ( _extract_requested_mcp_access_groups, _extract_requested_mcp_server_ids, @@ -1010,3 +1010,18 @@ async def test_empty_object_permission_passes_for_personal_non_admin(): team_obj=None, is_proxy_admin=False, ) + + +def test_object_permission_dict_mirrors_pydantic_model(): + """ObjectPermissionDict must stay field-for-field aligned with + LiteLLM_ObjectPermissionBase. If a new field is added to the Pydantic + model, this test fails until the TypedDict is updated to match.""" + from typing import get_type_hints + + pydantic_fields = set(LiteLLM_ObjectPermissionBase.model_fields.keys()) + typeddict_fields = set(get_type_hints(ObjectPermissionDict).keys()) + assert pydantic_fields == typeddict_fields, ( + f"ObjectPermissionDict drifted from LiteLLM_ObjectPermissionBase.\n" + f"Only in Pydantic model: {sorted(pydantic_fields - typeddict_fields)}\n" + f"Only in TypedDict: {sorted(typeddict_fields - pydantic_fields)}" + ) From 453aedef95e6d5cf3ad24379588428182dbd1b26 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Sun, 28 Jun 2026 20:49:28 +0000 Subject: [PATCH 07/37] chore(router): simplify unknown-model error message construction The error string is already produced by the f-string interpolation; the trailing .format() call on it was redundant. Add a regression test that the message renders the model name verbatim. --- litellm/router.py | 6 ++---- tests/test_litellm/test_router.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 30 insertions(+), 4 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index cdb45de66db..8abdd60ccad 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10032,11 +10032,9 @@ class Router: # If still no deployments after checking for fallbacks, raise an error if len(healthy_deployments) == 0: if self.get_model_list(model_name=model) is None: - message = f"You passed in model={model}. There is no 'model_name' with this string".format(model) + message = f"You passed in model={model}. There is no 'model_name' with this string" else: - message = f"You passed in model={model}. There are no healthy deployments for this model".format( - model - ) + message = f"You passed in model={model}. There are no healthy deployments for this model" raise litellm.BadRequestError( message=message, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 5f214506197..9c4d83ff7ea 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3268,6 +3268,34 @@ async def test_router_acompletion_with_unknown_model_and_no_fallback(): assert "no healthy deployments for this model" in str(excinfo.value) +@pytest.mark.asyncio +async def test_router_unknown_model_error_message_renders_model_name_literally(): + """ + The unknown-model error message renders the caller-supplied model name + verbatim. A name containing Python format-field syntax must be treated as + literal text, not re-interpreted as a format template, which would distort + the message and balloon its length. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "azure/gpt-4o-real", "api_key": "fake-key"}, + } + ] + ) + + weird_model = "ghost{:>200}model" + messages = [{"role": "user", "content": "hi"}] + + with pytest.raises(litellm.BadRequestError) as excinfo: + await router.acompletion(model=weird_model, messages=messages) + + message = str(excinfo.value) + assert weird_model in message + assert " " not in message # no padding run from an expanded format field + + def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint(): """ Test that get_deployment_credentials_with_provider correctly copies From 9b4442c6df65ad048097a1c0a87d23978deb2c73 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Sun, 28 Jun 2026 21:37:11 +0000 Subject: [PATCH 08/37] chore(router): drop unreachable unknown-model error branch get_model_list always returns a list, never None, so the is-None branch could not execute. Collapse to the single reachable message. --- litellm/router.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 8abdd60ccad..2d66bc3158d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10031,10 +10031,7 @@ class Router: # If still no deployments after checking for fallbacks, raise an error if len(healthy_deployments) == 0: - if self.get_model_list(model_name=model) is None: - message = f"You passed in model={model}. There is no 'model_name' with this string" - else: - message = f"You passed in model={model}. There are no healthy deployments for this model" + message = f"You passed in model={model}. There are no healthy deployments for this model" raise litellm.BadRequestError( message=message, From a04321d2e119e501678e5efd465bd73b1b4cf4a7 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 29 Jun 2026 09:12:51 +0530 Subject: [PATCH 09/37] test(videos): add 1:1 test file scaffold for videos component paths (#30631) Keep only video test files and CI workflow entries; drop unrelated production code and non-video test changes from this branch. Co-authored-by: Cursor --- .github/workflows/test-unit-misc.yml | 1 + .../workflows/test-unit-proxy-endpoints.yml | 1 + .../proxy/video_endpoints/__init__.py | 0 .../proxy/video_endpoints/test_endpoints.py | 656 ++++++++++++++++++ .../proxy/video_endpoints/test_utils.py | 184 +++++ tests/test_litellm/videos/__init__.py | 0 tests/test_litellm/videos/test_main.py | 458 ++++++++++++ tests/test_litellm/videos/test_utils.py | 196 ++++++ 8 files changed, 1496 insertions(+) create mode 100644 tests/test_litellm/proxy/video_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/video_endpoints/test_endpoints.py create mode 100644 tests/test_litellm/proxy/video_endpoints/test_utils.py create mode 100644 tests/test_litellm/videos/__init__.py create mode 100644 tests/test_litellm/videos/test_main.py create mode 100644 tests/test_litellm/videos/test_utils.py diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index 133c135d97a..d411d996d8e 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -36,6 +36,7 @@ jobs: tests/test_litellm/passthrough tests/test_litellm/sandbox tests/test_litellm/vector_stores + tests/test_litellm/videos tests/test_litellm/test_*.py workers: 2 reruns: 2 diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index d4f00050596..23ca7a8f2f9 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -31,6 +31,7 @@ jobs: tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/openai_files_endpoint + tests/test_litellm/proxy/video_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/vector_store_endpoints diff --git a/tests/test_litellm/proxy/video_endpoints/__init__.py b/tests/test_litellm/proxy/video_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py new file mode 100644 index 00000000000..40a26fad3c3 --- /dev/null +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -0,0 +1,656 @@ +""" +Routing-contract tests for litellm/proxy/video_endpoints/endpoints.py + +Unlike the batches layer, every video endpoint funnels into a single downstream +seam - ProxyBaseLLMRequestProcessing.base_process_llm_request - so there is no +provider-dispatch to assert. All of the video-specific, regression-worthy logic +runs *before* that call, while the endpoint assembles the `data` dict. Each test +therefore locks four things: + + 1. ROUTE_TYPE - the exact route_type each endpoint forwards + (avideo_generation/status/content/edit). Swapping two would + silently route requests to the wrong handler. + 2. DATA SHAPE - the entire `data` dict the processor is constructed with: + provider-precedence resolution, video_id passthrough/extraction, + model resolution from the decoded model_id, and file attachment. + 3. RESULT - base_process_llm_request's return value is propagated untouched + (except where the endpoint transforms it). + 4. OUTPUT SHAPE - video_content wraps raw bytes in a Response (video/mp4 + + Content-Disposition). + +Only true I/O boundaries are mocked (the downstream processor call, request body +parsing, file->bytes conversion, the provider-from-request readers, and the +router's model-id resolver). The id decode helpers and get_custom_provider_from_data +run for real, so the data assertions reflect production exactly. base_process is +patched with autospec so the real __init__ still stores self.data (captured via the +mock's call args), and a brand-new kwarg added to this layer surfaces as a failure. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import orjson +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm.proxy.proxy_server as proxy_server +import litellm.proxy.video_endpoints.endpoints as endpoints +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.utils import ProxyLogging +from litellm.router import Router +from litellm.types.videos.utils import ( + encode_character_id_with_provider, + encode_video_id_with_provider, +) + +from fastapi import Response + +# --------------------------------------------------------------------------- # +# A real model-encoded video id: decodes (for real) to provider "azure", +# model_id VIDEO_MODEL_ID, original video id "video_orig123". The router's +# resolver maps that model_id to a model name; an unknown id resolves to None, +# so a wrong/hardcoded model_id cannot produce a plausible-looking result. +# --------------------------------------------------------------------------- # + +VIDEO_MODEL_ID = "deployment-123" +AZURE_VIDEO_ID = encode_video_id_with_provider("video_orig123", "azure", VIDEO_MODEL_ID) +# A real model-encoded character id: decodes to provider "azure", VIDEO_MODEL_ID, +# original character id "char_orig". Distinct from the video id so a test cannot +# pass by reusing the wrong constant. +AZURE_CHARACTER_ID = encode_character_id_with_provider( + "char_orig", "azure", VIDEO_MODEL_ID +) +RESOLVED_MODELS: Dict[str, str] = {VIDEO_MODEL_ID: "azure-sora"} + +# Sentinel propagated by base_process for the passthrough endpoints. +SENTINEL = object() + + +class FakeRequest: + """Minimal stand-in. headers/query_params are read by the provider readers + (mocked) and on the edit path the raw body is parsed for real via orjson.""" + + def __init__( + self, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + raw_body: bytes = b"{}", + ): + self.headers = headers or {} + self.query_params = query or {} + self._raw_body = raw_body + + async def body(self) -> bytes: + return self._raw_body + + +@dataclass +class Harness: + read_body: AsyncMock + batch_to_bytesio: AsyncMock + base_process: MagicMock + handle_exc: AsyncMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + provider_from_body: AsyncMock + router: MagicMock + resolve_model: MagicMock + + def processor_data(self) -> Dict[str, Any]: + """The exact `data` dict the processor was constructed with.""" + assert self.base_process.call_count == 1 + return dict(self.base_process.call_args.args[0].data) + + def route_type(self) -> str: + return self.base_process.call_args.kwargs["route_type"] + + +@pytest.fixture +def harness(): + logging = MagicMock(spec=ProxyLogging) + + router = MagicMock(spec=Router) + resolve_model = MagicMock( + side_effect=lambda model_id: RESOLVED_MODELS.get(model_id) + ) + router.resolve_model_name_from_model_id = resolve_model + + read_body = AsyncMock(return_value={}) + batch_to_bytesio = AsyncMock(return_value=[b"filebytes"]) + handle_exc = AsyncMock(return_value=RuntimeError("handled")) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + provider_from_body = AsyncMock(return_value=None) + + with ExitStack() as stack: + base_process = stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + autospec=True, + ) + ) + base_process.return_value = SENTINEL + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "_handle_llm_api_exception", + handle_exc, + ) + ) + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object(endpoints, "batch_to_bytesio", batch_to_bytesio) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_body", + provider_from_body, + ) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context( + patch.object(proxy_server, "select_data_generator", MagicMock()) + ) + stack.enter_context(patch.object(proxy_server, "user_model", None)) + stack.enter_context(patch.object(proxy_server, "user_temperature", None)) + stack.enter_context(patch.object(proxy_server, "user_request_timeout", None)) + stack.enter_context(patch.object(proxy_server, "user_max_tokens", None)) + stack.enter_context(patch.object(proxy_server, "user_api_base", None)) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + yield Harness( + read_body=read_body, + batch_to_bytesio=batch_to_bytesio, + base_process=base_process, + handle_exc=handle_exc, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + provider_from_body=provider_from_body, + router=router, + resolve_model=resolve_model, + ) + + +def _user() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test") + + +# =========================================================================== # +# POST /v1/videos - video_generation # +# =========================================================================== # + + +async def call_generation( + harness: Harness, *, body: Dict[str, Any], input_reference=None +): + harness.read_body.return_value = body + return await endpoints.video_generation( + request=FakeRequest(), + fastapi_response=Response(), + input_reference=input_reference, + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_generation__route_type_data_and_no_provider_default(harness): + body = {"model": "sora-2", "prompt": "a sunset"} + + resp = await call_generation(harness, body=body) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_generation" + # generation does NOT resolve a provider; data is the body, untouched. A + # future default custom_llm_provider injection would break this row. + assert harness.processor_data() == {"model": "sora-2", "prompt": "a sunset"} + harness.batch_to_bytesio.assert_not_called() + + +@pytest.mark.asyncio +async def test_generation__input_reference_attached(harness): + body = {"model": "sora-2", "prompt": "a sunset"} + upload = MagicMock(name="upload_file") + + await call_generation(harness, body=body, input_reference=upload) + + harness.batch_to_bytesio.assert_called_once_with([upload]) + assert harness.processor_data() == { + "model": "sora-2", + "prompt": "a sunset", + "input_reference": b"filebytes", + } + + +@pytest.mark.asyncio +async def test_generation__exception_routed_through_handler(harness): + harness.base_process.side_effect = ValueError("provider boom") + + with pytest.raises(RuntimeError, match="handled"): + await call_generation(harness, body={"model": "sora-2"}) + + harness.handle_exc.assert_called_once() + assert harness.handle_exc.call_args.kwargs["e"].args[0] == "provider boom" + + +# =========================================================================== # +# GET /v1/videos/{video_id} - video_status # +# =========================================================================== # + + +async def call_status(harness: Harness, video_id: str, *, headers=None, query=None): + return await endpoints.video_status( + video_id=video_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_status__model_encoded_id_full_contract(harness): + resp = await call_status(harness, AZURE_VIDEO_ID) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_status" + # provider comes from the decoded id; model_id resolved to a model name. + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_status__plain_id_defaults_to_openai(harness): + await call_status(harness, "video_plain") + + # plain id -> nothing decoded, no header/query/body provider -> "openai". + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "video_id": "video_plain", + "custom_llm_provider": "openai", + } + + +@pytest.mark.asyncio +async def test_status__header_provider_beats_decoded_id(harness): + harness.provider_from_headers.return_value = "bedrock" + + await call_status(harness, AZURE_VIDEO_ID) + + data = harness.processor_data() + # header wins over the provider decoded from the id ... + assert data["custom_llm_provider"] == "bedrock" + # ... but the model is still resolved from the decoded model_id. + assert data["model"] == "azure-sora" + + +# =========================================================================== # +# GET /v1/videos/{video_id}/content - video_content # +# =========================================================================== # + + +async def call_content(harness: Harness, video_id: str, *, headers=None, query=None): + return await endpoints.video_content( + video_id=video_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_content__wraps_raw_bytes_in_response(harness): + harness.base_process.return_value = b"VIDEOBYTES" + + resp = await call_content(harness, "video_plain") + + assert harness.route_type() == "avideo_content" + assert isinstance(resp, Response) + assert resp.body == b"VIDEOBYTES" + assert resp.media_type == "video/mp4" + assert ( + resp.headers["content-disposition"] + == "attachment; filename=video_video_plain.mp4" + ) + + +@pytest.mark.asyncio +async def test_content__plain_id_has_no_openai_default(harness): + """The high-value asymmetry vs video_status: content stops at the decoded + provider and never injects an 'openai' default, so a plain id leaves + custom_llm_provider unset. A copy-paste of status' fallback breaks this.""" + harness.base_process.return_value = b"x" + + await call_content(harness, "video_plain") + + assert harness.processor_data() == {"video_id": "video_plain"} + + +@pytest.mark.asyncio +async def test_content__model_encoded_id(harness): + harness.base_process.return_value = b"x" + + await call_content(harness, AZURE_VIDEO_ID) + + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +# =========================================================================== # +# POST /v1/videos/edits - video_edit # +# =========================================================================== # + + +async def call_edit( + harness: Harness, *, body: Dict[str, Any], headers=None, query=None +): + return await endpoints.video_edit( + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_edit__extracts_nested_video_id_full_contract(harness): + resp = await call_edit( + harness, body={"prompt": "brighter", "video": {"id": AZURE_VIDEO_ID}} + ) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_edit" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + # nested video object is popped; its id becomes video_id; provider/model + # derived from the encoded id. + assert harness.processor_data() == { + "prompt": "brighter", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_edit__provider_from_body_data_for_plain_id(harness): + """For a plain id, get_custom_provider_from_data (run for real) pulls the + provider out of the request body before the 'openai' default.""" + await call_edit( + harness, + body={ + "prompt": "x", + "video": {"id": "video_plain"}, + "custom_llm_provider": "vertex_ai", + }, + ) + + data = harness.processor_data() + assert data["video_id"] == "video_plain" + assert data["custom_llm_provider"] == "vertex_ai" + harness.resolve_model.assert_not_called() + + +@pytest.mark.asyncio +async def test_edit__missing_video_object_defaults_to_openai(harness): + await call_edit(harness, body={"prompt": "x"}) + + data = harness.processor_data() + # no video object -> empty video_id; plain -> default provider. + assert data["video_id"] == "" + assert data["custom_llm_provider"] == "openai" + assert "video" not in data + + +# =========================================================================== # +# GET /v1/videos - video_list # +# =========================================================================== # + + +async def call_list(harness: Harness, *, headers=None, query=None): + return await endpoints.video_list( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_list__query_params_and_no_provider(harness): + resp = await call_list(harness, query={"limit": "5"}) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_list" + # no provider anywhere -> custom_llm_provider stays absent (only set if truthy). + assert harness.processor_data() == {"query_params": {"limit": "5"}} + + +@pytest.mark.asyncio +async def test_list__provider_from_header(harness): + harness.provider_from_headers.return_value = "bedrock" + + await call_list(harness) + + assert harness.processor_data() == { + "query_params": {}, + "custom_llm_provider": "bedrock", + } + + +# =========================================================================== # +# POST /v1/videos/{video_id}/remix - video_remix # +# =========================================================================== # + + +async def call_remix( + harness: Harness, video_id: str, *, body, headers=None, query=None +): + return await endpoints.video_remix( + video_id=video_id, + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_remix__model_encoded_id_full_contract(harness): + resp = await call_remix(harness, AZURE_VIDEO_ID, body={"prompt": "new colors"}) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_remix" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "prompt": "new colors", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_remix__provider_from_body_data_not_request_body_reader(harness): + """remix resolves the provider from data.get('custom_llm_provider'), never + from the async request-body reader (unlike status/get_character). Setting + that reader to a sentinel and asserting it is untouched locks the difference.""" + harness.provider_from_body.return_value = "must-not-win" + + await call_remix( + harness, + "video_plain", + body={"prompt": "x", "custom_llm_provider": "vertex_ai"}, + ) + + harness.provider_from_body.assert_not_called() + data = harness.processor_data() + assert data["video_id"] == "video_plain" + assert data["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_remix__plain_id_has_no_openai_default(harness): + await call_remix(harness, "video_plain", body={"prompt": "x"}) + + # like video_content, remix stops at provider_from_id with no 'openai' default. + assert harness.processor_data() == {"prompt": "x", "video_id": "video_plain"} + + +# =========================================================================== # +# POST /v1/videos/characters - video_create_character # +# =========================================================================== # + + +async def call_create_character(harness: Harness, *, body, video=None, name="my_char"): + harness.read_body.return_value = body + return await endpoints.video_create_character( + request=FakeRequest(), + fastapi_response=Response(), + video=video if video is not None else MagicMock(name="video_upload"), + name=name, + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_create_character__video_attached_default_provider_no_encode(harness): + upload = MagicMock(name="video_upload") + + resp = await call_create_character(harness, body={"prompt": "x"}, video=upload) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_create_character" + harness.batch_to_bytesio.assert_called_once_with([upload]) + # no target_model_names -> no model injected, no id re-encoding. + assert harness.processor_data() == { + "prompt": "x", + "video": b"filebytes", + "custom_llm_provider": "openai", + } + + +@pytest.mark.asyncio +async def test_create_character__target_model_sets_model_and_encodes_id(harness): + harness.base_process.return_value = {"id": "char_raw"} + + resp = await call_create_character( + harness, + body={"target_model_names": "azure-sora-model", "custom_llm_provider": "azure"}, + ) + + data = harness.processor_data() + assert data["model"] == "azure-sora-model" + assert data["custom_llm_provider"] == "azure" + # response id re-encoded with the resolved provider + model for the round-trip. + assert resp["id"] == encode_character_id_with_provider( + "char_raw", "azure", "azure-sora-model" + ) + + +# =========================================================================== # +# GET /v1/videos/characters/{character_id} - video_get_character # +# =========================================================================== # + + +async def call_get_character( + harness: Harness, character_id: str, *, headers=None, query=None +): + return await endpoints.video_get_character( + character_id=character_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_get_character__encoded_id_full_contract(harness): + harness.base_process.return_value = {"id": "char_raw2"} + + resp = await call_get_character(harness, AZURE_CHARACTER_ID) + + assert harness.route_type() == "avideo_get_character" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + # character_id decoded to its inner value; provider/model from the encoded id. + assert harness.processor_data() == { + "character_id": "char_orig", + "custom_llm_provider": "azure", + "model": "azure-sora", + } + # response id re-encoded for the client round-trip. + assert resp["id"] == encode_character_id_with_provider( + "char_raw2", "azure", VIDEO_MODEL_ID + ) + + +@pytest.mark.asyncio +async def test_get_character__plain_id_defaults_openai_no_encode(harness): + harness.base_process.return_value = {"id": "char_raw3"} + + resp = await call_get_character(harness, "char_plain") + + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "character_id": "char_plain", + "custom_llm_provider": "openai", + } + # id does not start with 'character_' -> returned untouched. + assert resp["id"] == "char_raw3" + + +# =========================================================================== # +# POST /v1/videos/extensions - video_extension # +# =========================================================================== # + + +async def call_extension(harness: Harness, *, body, headers=None, query=None): + return await endpoints.video_extension( + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_extension__extracts_nested_video_id_full_contract(harness): + resp = await call_extension( + harness, body={"prompt": "continue", "video": {"id": AZURE_VIDEO_ID}} + ) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_extension" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "prompt": "continue", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py new file mode 100644 index 00000000000..ae22ae233b5 --- /dev/null +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -0,0 +1,184 @@ +""" +Pure-logic contract tests for litellm/proxy/video_endpoints/utils.py + +Three helpers the video proxy endpoints lean on: + - extract_model_from_target_model_names: first model from a comma string / list + - get_custom_provider_from_data: provider precedence (top-level > extra_body) + - encode_character_id_in_response: re-encode a response id in place + +Every test asserts the exact result (or identity), so a mutation that flips a +branch, drops a strip/filter, or changes precedence fails. The only collaborator +is encode_character_id_with_provider, which runs for real; encoding assertions +are checked by the genuine decode round-trip. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.proxy.video_endpoints.utils import ( + encode_character_id_in_response, + extract_model_from_target_model_names, + get_custom_provider_from_data, +) +from litellm.types.videos.utils import ( + decode_character_id_with_provider, + encode_character_id_with_provider, +) + +# =========================================================================== # +# extract_model_from_target_model_names +# =========================================================================== # + + +@pytest.mark.parametrize( + "value,expected", + [ + ("m1,m2,m3", "m1"), + (" a , b ", "a"), # leading/trailing whitespace stripped + (",, m1 ,,", "m1"), # empty tokens filtered out + ("solo", "solo"), # single token, no comma + ("", None), # empty string -> no tokens + (" , , ", None), # only separators/whitespace -> no tokens + (["x", "y"], "x"), # list -> first element + ([], None), # empty list + ], +) +def test_extract_model__str_and_list(value, expected): + assert extract_model_from_target_model_names(value) == expected + + +@pytest.mark.parametrize("value", [None, 123, {"a": 1}, 4.5]) +def test_extract_model__non_str_non_list_is_none(value): + assert extract_model_from_target_model_names(value) is None + + +# =========================================================================== # +# get_custom_provider_from_data +# =========================================================================== # + + +def test_provider__top_level_wins_over_extra_body(): + data = { + "custom_llm_provider": "azure", + "extra_body": {"custom_llm_provider": "openai"}, + } + assert get_custom_provider_from_data(data) == "azure" + + +@pytest.mark.parametrize("falsy", ["", None]) +def test_provider__falsy_top_level_falls_through_to_extra_body(falsy): + data = { + "custom_llm_provider": falsy, + "extra_body": {"custom_llm_provider": "vertex_ai"}, + } + assert get_custom_provider_from_data(data) == "vertex_ai" + + +def test_provider__from_extra_body_dict(): + assert ( + get_custom_provider_from_data( + {"extra_body": {"custom_llm_provider": "bedrock"}} + ) + == "bedrock" + ) + + +def test_provider__from_extra_body_json_string(): + data = {"extra_body": '{"custom_llm_provider": "gemini"}'} + assert get_custom_provider_from_data(data) == "gemini" + + +def test_provider__invalid_json_string_is_none(): + assert get_custom_provider_from_data({"extra_body": "not-json{"}) is None + + +def test_provider__json_string_parsing_to_non_dict_is_none(): + # parses to a list, not a dict -> no provider extracted. + assert get_custom_provider_from_data({"extra_body": "[1, 2]"}) is None + + +def test_provider__extra_body_provider_not_a_string_is_none(): + assert ( + get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) + is None + ) + + +@pytest.mark.parametrize( + "data", + [ + {}, + {"extra_body": {}}, + {"extra_body": 5}, # non-dict, non-str + {"extra_body": {"other": "x"}}, # dict without provider key + ], +) +def test_provider__no_provider_anywhere_is_none(data): + assert get_custom_provider_from_data(data) is None + + +# =========================================================================== # +# encode_character_id_in_response +# =========================================================================== # + + +class _Resp: + """Minimal response object exposing an `id` attribute.""" + + +def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys(): + response = {"id": "char_raw", "object": "character", "name": "hero"} + + out = encode_character_id_in_response(response, "azure", "model-1") + + assert out is response # same dict, mutated in place + assert out["object"] == "character" and out["name"] == "hero" + assert out["id"] == encode_character_id_with_provider( + "char_raw", "azure", "model-1" + ) + decoded = decode_character_id_with_provider(out["id"]) + assert decoded["custom_llm_provider"] == "azure" + assert decoded["model_id"] == "model-1" + assert decoded["character_id"] == "char_raw" + + +@pytest.mark.parametrize("response", [{}, {"id": ""}, {"id": None}]) +def test_encode__dict_without_usable_id_unchanged(response): + snapshot = dict(response) + out = encode_character_id_in_response(response, "azure", "model-1") + assert out == snapshot + + +def test_encode__object_with_str_id(): + resp = _Resp() + resp.id = "char_raw" + + out = encode_character_id_in_response(resp, "openai", None) + + assert out is resp + assert resp.id == encode_character_id_with_provider("char_raw", "openai", None) + decoded = decode_character_id_with_provider(resp.id) + assert decoded["custom_llm_provider"] == "openai" + assert decoded["character_id"] == "char_raw" + + +@pytest.mark.parametrize("bad_id", [None, 123, ""]) +def test_encode__object_non_str_or_empty_id_unchanged(bad_id): + resp = _Resp() + resp.id = bad_id + + out = encode_character_id_in_response(resp, "azure", "model-1") + + assert out is resp + assert resp.id == bad_id # untouched + + +def test_encode__object_without_id_attr_returned_unchanged(): + resp = _Resp() + out = encode_character_id_in_response(resp, "azure", "model-1") + assert out is resp + assert not hasattr(resp, "id") diff --git a/tests/test_litellm/videos/__init__.py b/tests/test_litellm/videos/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/videos/test_main.py b/tests/test_litellm/videos/test_main.py new file mode 100644 index 00000000000..a04a89ded99 --- /dev/null +++ b/tests/test_litellm/videos/test_main.py @@ -0,0 +1,458 @@ +""" +Dispatch-contract tests for litellm/videos/main.py + +Each public video operation is a pair: a sync `video_*` worker (decorated with +@client) that resolves the provider, fetches the provider config, logs, and then +forwards to exactly one `base_llm_http_handler.video_*_handler`; and an async +`avideo_*` wrapper that delegates to the sync worker in an executor. + +This file locks the contract of that layer so a regression fails loudly: + + 1. DISPATCH - the one correct handler fired and every sibling video handler + asserted NOT called. A copy-paste that calls the wrong handler + (e.g. remix -> edit) flips this. + 2. RESULT - the handler's return value is propagated by identity. + 3. PROVIDER - custom_llm_provider is decoded from an encoded video id when not + passed (status/content/remix/edit/extension), or defaults to + "openai" (list/create_character/get_character). This is the exact + surface of the historical "content defaulted to openai" bug. + 4. PAYLOAD - the provider config object and the operation's identifying args + (video_id/prompt/name/...) reach the handler; _is_async is False + on the sync path. + 5. SHORT-CIRCUIT - mock_response returns a typed object without any handler call. + 6. UNSUPPORTED - a None provider config raises before any handler fires. + 7. DELEGATION - avideo_* returns the sync worker's result untouched, sets + async_call=True, and pre-resolves the provider where it must. + +Seams mocked: the http handler (network), the provider-config registry lookup, +get_llm_provider, and the video-generation optional-param builders. The id decode +helper runs for real against genuinely-encoded ids, so the provider assertions +reflect production. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.types.videos.main import CharacterObject, VideoObject +from litellm.types.videos.utils import encode_video_id_with_provider +from litellm.videos import main as videos_main + +# A real model-encoded video id: decodes (for real) to provider "azure". Used to +# prove the sync workers derive custom_llm_provider from the id, not a hardcode. +AZURE_VIDEO_ID = encode_video_id_with_provider("video_raw", "azure", "deployment-1") + +# The nine sync handlers on base_llm_http_handler. Dispatch tests assert exactly +# one fired and the other eight did not. +SYNC_HANDLERS = ( + "video_generation_handler", + "video_content_handler", + "video_remix_handler", + "video_create_character_handler", + "video_get_character_handler", + "video_edit_handler", + "video_extension_handler", + "video_list_handler", + "video_status_handler", +) + +GEN_OPTIONAL_PARAMS = {"seconds": "8", "size": "720x1280"} + + +@dataclass +class Seams: + handler: MagicMock + get_config: MagicMock + config: MagicMock + + def kwargs_of(self, handler_name: str) -> Dict[str, Any]: + method = getattr(self.handler, handler_name) + assert method.call_count == 1 + return dict(method.call_args.kwargs) + + def assert_only(self, handler_name: str) -> None: + for name in SYNC_HANDLERS: + method = getattr(self.handler, name) + if name == handler_name: + method.assert_called_once() + else: + method.assert_not_called() + + +@pytest.fixture +def seams(): + handler = MagicMock(spec=BaseLLMHTTPHandler) + config = MagicMock(name="provider_video_config") + get_config = MagicMock(return_value=config) + + with ExitStack() as stack: + stack.enter_context(patch.object(videos_main, "base_llm_http_handler", handler)) + stack.enter_context( + patch.object( + videos_main.ProviderConfigManager, + "get_provider_video_config", + get_config, + ) + ) + # video_generation resolves model+provider through get_llm_provider and + # builds optional params; mock those so the dispatch payload is deterministic. + stack.enter_context( + patch.object( + videos_main, + "get_llm_provider", + MagicMock(return_value=("sora-2", "openai", None, None)), + ) + ) + stack.enter_context( + patch.object( + videos_main.VideoGenerationRequestUtils, + "get_requested_video_generation_optional_param", + MagicMock(return_value={"seconds": "8"}), + ) + ) + stack.enter_context( + patch.object( + videos_main.VideoGenerationRequestUtils, + "get_optional_params_video_generation", + MagicMock(return_value=dict(GEN_OPTIONAL_PARAMS)), + ) + ) + yield Seams(handler=handler, get_config=get_config, config=config) + + +# =========================================================================== # +# Dispatch contract - one rich test per sync worker. +# =========================================================================== # + + +def test_video_generation__dispatch(seams): + result = videos_main.video_generation(prompt="a sunset", model="sora-2") + + seams.assert_only("video_generation_handler") + assert result is seams.handler.video_generation_handler.return_value + kw = seams.kwargs_of("video_generation_handler") + assert kw["model"] == "sora-2" + assert kw["prompt"] == "a sunset" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_generation_provider_config"] is seams.config + assert kw["video_generation_optional_request_params"] == GEN_OPTIONAL_PARAMS + assert kw["_is_async"] is False + + +def test_video_status__dispatch_and_provider_from_id(seams): + result = videos_main.video_status(video_id=AZURE_VIDEO_ID) + + seams.assert_only("video_status_handler") + assert result is seams.handler.video_status_handler.return_value + kw = seams.kwargs_of("video_status_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["custom_llm_provider"] == "azure" # decoded from the id, not openai + assert kw["video_status_provider_config"] is seams.config + assert kw["_is_async"] is False + # provider config requested for the decoded provider, not a hardcode. + assert seams.get_config.call_args.kwargs["provider"] == litellm.LlmProviders.AZURE + + +def test_video_content__dispatch_and_provider_from_id(seams): + result = videos_main.video_content(video_id=AZURE_VIDEO_ID, variant="thumbnail") + + seams.assert_only("video_content_handler") + assert result is seams.handler.video_content_handler.return_value + kw = seams.kwargs_of("video_content_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["custom_llm_provider"] == "azure" + assert kw["variant"] == "thumbnail" + assert kw["video_content_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_content__plain_id_defaults_to_openai(seams): + videos_main.video_content(video_id="video_plain") + + assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai" + + +def test_video_remix__dispatch_and_provider_from_id(seams): + result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors") + + seams.assert_only("video_remix_handler") + assert result is seams.handler.video_remix_handler.return_value + kw = seams.kwargs_of("video_remix_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "new colors" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_remix_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_edit__dispatch_and_provider_from_id(seams): + result = videos_main.video_edit(video_id=AZURE_VIDEO_ID, prompt="brighter") + + seams.assert_only("video_edit_handler") + assert result is seams.handler.video_edit_handler.return_value + kw = seams.kwargs_of("video_edit_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "brighter" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_extension__dispatch_and_provider_from_id(seams): + result = videos_main.video_extension( + video_id=AZURE_VIDEO_ID, prompt="continue", seconds="5" + ) + + seams.assert_only("video_extension_handler") + assert result is seams.handler.video_extension_handler.return_value + kw = seams.kwargs_of("video_extension_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "continue" + assert kw["seconds"] == "5" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_list__dispatch_defaults_to_openai(seams): + result = videos_main.video_list(after="cur", limit=5, order="desc") + + seams.assert_only("video_list_handler") + assert result is seams.handler.video_list_handler.return_value + kw = seams.kwargs_of("video_list_handler") + assert kw["after"] == "cur" + assert kw["limit"] == 5 + assert kw["order"] == "desc" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_list_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_create_character__dispatch_defaults_to_openai(seams): + video = MagicMock(name="video_upload") + result = videos_main.video_create_character(name="hero", video=video) + + seams.assert_only("video_create_character_handler") + assert result is seams.handler.video_create_character_handler.return_value + kw = seams.kwargs_of("video_create_character_handler") + assert kw["name"] == "hero" + assert kw["video"] is video + assert kw["custom_llm_provider"] == "openai" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_get_character__dispatch_defaults_to_openai(seams): + result = videos_main.video_get_character(character_id="char_1") + + seams.assert_only("video_get_character_handler") + assert result is seams.handler.video_get_character_handler.return_value + kw = seams.kwargs_of("video_get_character_handler") + assert kw["character_id"] == "char_1" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_explicit_provider_beats_decoded_id(seams): + """An explicit custom_llm_provider wins over the one encoded in the id.""" + videos_main.video_status(video_id=AZURE_VIDEO_ID, custom_llm_provider="vertex_ai") + + assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "vertex_ai" + + +# =========================================================================== # +# mock_response short-circuit - returns a typed object, no handler call. +# =========================================================================== # + + +def test_generation__mock_response_short_circuits(seams): + resp = videos_main.video_generation( + prompt="x", + model="sora-2", + mock_response={"id": "v1", "object": "video", "status": "queued"}, + ) + + assert isinstance(resp, VideoObject) + assert resp.id == "v1" + seams.handler.video_generation_handler.assert_not_called() + + +def test_list__mock_response_short_circuits(seams): + resp = videos_main.video_list( + mock_response=[{"id": "v1", "object": "video", "status": "completed"}] + ) + + assert isinstance(resp, list) + assert resp[0].id == "v1" + seams.handler.video_list_handler.assert_not_called() + + +def test_get_character__mock_response_short_circuits(seams): + resp = videos_main.video_get_character( + character_id="char_1", + mock_response={ + "id": "char_1", + "object": "character", + "created_at": 1, + "name": "hero", + }, + ) + + assert isinstance(resp, CharacterObject) + assert resp.id == "char_1" + seams.handler.video_get_character_handler.assert_not_called() + + +# =========================================================================== # +# Unsupported provider - a None provider config raises before any dispatch. +# =========================================================================== # + + +def test_unsupported_provider_raises_without_dispatch(seams): + seams.get_config.return_value = None + + with pytest.raises(Exception): + videos_main.video_status(video_id=AZURE_VIDEO_ID) + + seams.handler.video_status_handler.assert_not_called() + + +# =========================================================================== # +# Async-wrapper delegation - representative coverage. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_avideo_generation__delegates_with_async_flag(): + sentinel = VideoObject(id="v-async", object="video", status="queued") + with ( + patch.object( + videos_main, "video_generation", MagicMock(return_value=sentinel) + ) as sync, + patch.object( + litellm, + "get_llm_provider", + MagicMock(return_value=("sora-2", "openai", None, None)), + ), + ): + result = await videos_main.avideo_generation(prompt="x", model="sora-2") + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_avideo_status__delegates_untouched(): + sentinel = VideoObject(id="v-async", object="video", status="queued") + with patch.object( + videos_main, "video_status", MagicMock(return_value=sentinel) + ) as sync: + result = await videos_main.avideo_status(video_id="video_plain") + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["video_id"] == "video_plain" + + +@pytest.mark.asyncio +async def test_avideo_content__pre_decodes_provider_before_delegating(): + """avideo_content resolves the provider from the encoded id itself before + handing off, so the sync worker receives the decoded provider, not None.""" + sentinel = b"mp4-bytes" + with patch.object( + videos_main, "video_content", MagicMock(return_value=sentinel) + ) as sync: + result = await videos_main.avideo_content(video_id=AZURE_VIDEO_ID) + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["custom_llm_provider"] == "azure" + + +# =========================================================================== # +# Credential passthrough - DB/YAML model-config credentials the router injects +# via kwargs must reach the provider call for EVERY video handler, carried in +# litellm_params. Distinct per-field values catch a cross-wired field. +# =========================================================================== # + +DB_YAML_CREDS = { + "api_key": "sk-db-credential", + "api_base": "https://db-resource.test", + "api_version": "2024-12-31", + "vertex_project": "db-project-xyz", +} + +CREDENTIAL_OPERATIONS = [ + ( + "video_generation_handler", + lambda: videos_main.video_generation( + prompt="p", model="sora-2", **DB_YAML_CREDS + ), + ), + ( + "video_status_handler", + lambda: videos_main.video_status(video_id=AZURE_VIDEO_ID, **DB_YAML_CREDS), + ), + ( + "video_content_handler", + lambda: videos_main.video_content(video_id=AZURE_VIDEO_ID, **DB_YAML_CREDS), + ), + ( + "video_remix_handler", + lambda: videos_main.video_remix( + video_id=AZURE_VIDEO_ID, prompt="p", **DB_YAML_CREDS + ), + ), + ( + "video_edit_handler", + lambda: videos_main.video_edit( + video_id=AZURE_VIDEO_ID, prompt="p", **DB_YAML_CREDS + ), + ), + ( + "video_extension_handler", + lambda: videos_main.video_extension( + video_id=AZURE_VIDEO_ID, prompt="p", seconds="5", **DB_YAML_CREDS + ), + ), + ( + "video_list_handler", + lambda: videos_main.video_list(**DB_YAML_CREDS), + ), + ( + "video_create_character_handler", + lambda: videos_main.video_create_character( + name="hero", video=MagicMock(name="vid"), **DB_YAML_CREDS + ), + ), + ( + "video_get_character_handler", + lambda: videos_main.video_get_character(character_id="char_1", **DB_YAML_CREDS), + ), +] + + +@pytest.mark.parametrize( + "handler_name,invoke", + CREDENTIAL_OPERATIONS, + ids=[op[0] for op in CREDENTIAL_OPERATIONS], +) +def test_db_yaml_credentials_reach_every_handler(seams, handler_name, invoke): + invoke() + + litellm_params = seams.kwargs_of(handler_name)["litellm_params"] + assert litellm_params.get("api_key") == DB_YAML_CREDS["api_key"] + assert litellm_params.get("api_base") == DB_YAML_CREDS["api_base"] + assert litellm_params.get("api_version") == DB_YAML_CREDS["api_version"] + assert litellm_params.get("vertex_project") == DB_YAML_CREDS["vertex_project"] diff --git a/tests/test_litellm/videos/test_utils.py b/tests/test_litellm/videos/test_utils.py new file mode 100644 index 00000000000..09975829531 --- /dev/null +++ b/tests/test_litellm/videos/test_utils.py @@ -0,0 +1,196 @@ +""" +Pure-logic contract tests for litellm/videos/main.py's request utils +(litellm/videos/utils.py: VideoGenerationRequestUtils). + +These lock the exact param-shaping behavior so a mutation that drops a filter, +flips a precedence, or stops removing a key fails loudly. The only seam is the +provider config's map_openai_params (a provider boundary); filter_out_litellm_params +runs for real, so the "litellm-internal params get stripped" assertions reflect +production. Every test asserts the exact resulting dict, never "ran without error". +""" + +import os +import sys +from unittest.mock import MagicMock + + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.videos.utils import VideoGenerationRequestUtils + +get_requested = ( + VideoGenerationRequestUtils.get_requested_video_generation_optional_param +) +get_optional = VideoGenerationRequestUtils.get_optional_params_video_generation + + +# =========================================================================== # +# get_requested_video_generation_optional_param +# +# Receives the caller's full local_vars; must return only the API-bound optional +# params. filter_out_litellm_params strips known internal keys for real; the +# values used below were chosen against the live set: seconds/size/user/foo_param/ +# vertex_project/extra/a/b survive, api_key/metadata/litellm_* are stripped. +# =========================================================================== # + + +def test_requested__drops_none_and_excluded_keys(): + result = get_requested( + { + "seconds": "8", + "size": None, # None -> dropped + "prompt": "a sunset", # excluded + "model": "sora-2", # excluded + "user": "u1", + } + ) + assert result == {"seconds": "8", "user": "u1"} + + +def test_requested__strips_litellm_internal_params(): + result = get_requested( + { + "seconds": "8", + "api_key": "sk-secret", + "metadata": {"x": 1}, + "litellm_call_id": "id-123", + } + ) + assert result == {"seconds": "8"} + + +def test_requested__timeout_always_removed(): + # timeout is NOT a litellm-internal param, so only the explicit pop removes it. + result = get_requested({"seconds": "8", "timeout": 30}) + assert result == {"seconds": "8"} + + +def test_requested__nested_kwargs_merge_and_override_base(): + result = get_requested( + {"seconds": "8", "kwargs": {"size": "720x1280", "seconds": "override"}} + ) + # nested kwargs win over the top-level base params on collision. + assert result == {"seconds": "override", "size": "720x1280"} + + +def test_requested__non_dict_kwargs_treated_as_empty(): + result = get_requested({"seconds": "8", "kwargs": "not-a-dict"}) + assert result == {"seconds": "8"} + + +def test_requested__none_input_returns_empty(): + assert get_requested(None) == {} + + +def test_requested__top_level_extra_body_spread_and_preserved(): + result = get_requested( + {"seconds": "8", "extra_body": {"vertex_project": "proj", "foo_param": "bar"}} + ) + # extra_body keys are both spread at top level AND kept under "extra_body". + assert result == { + "seconds": "8", + "vertex_project": "proj", + "foo_param": "bar", + "extra_body": {"vertex_project": "proj", "foo_param": "bar"}, + } + + +def test_requested__extra_body_kwargs_overrides_top_level(): + result = get_requested( + { + "extra_body": {"a": "top", "b": "top_b"}, + "kwargs": {"extra_body": {"a": "kw"}}, + } + ) + # kwargs' extra_body wins over the top-level extra_body on collision; the + # non-colliding top-level key survives. + assert result == { + "a": "kw", + "b": "top_b", + "extra_body": {"a": "kw", "b": "top_b"}, + } + + +def test_requested__extra_body_strips_litellm_internal_params(): + result = get_requested({"extra_body": {"api_key": "sk", "foo_param": "bar"}}) + # api_key filtered out of extra_body; only foo_param remains (and is spread). + assert result == {"foo_param": "bar", "extra_body": {"foo_param": "bar"}} + + +def test_requested__empty_extra_body_not_added(): + result = get_requested({"seconds": "8", "extra_body": {}}) + assert result == {"seconds": "8"} + assert "extra_body" not in result + + +# =========================================================================== # +# get_optional_params_video_generation +# +# Delegates mapping to the provider config (the seam) then folds extra_body in. +# =========================================================================== # + + +def _config(map_return): + config = MagicMock() + config.map_openai_params.return_value = map_return + return config + + +def test_optional__delegates_to_map_openai_params_with_drop_params(): + config = _config({"seconds": "8"}) + optional_params = {"seconds": "8"} + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params=optional_params, + ) + + assert result == {"seconds": "8"} + config.map_openai_params.assert_called_once_with( + video_create_optional_params=optional_params, + model="sora-2", + drop_params=litellm.drop_params, + ) + + +def test_optional__extra_body_overrides_mapped_and_is_removed(): + # mapped output carries a leftover extra_body that must be popped; the input + # extra_body overrides a colliding mapped key and is spread in. + config = _config({"seconds": "8", "size": "mapped", "extra_body": {"leftover": 1}}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={ + "extra_body": {"size": "override", "extra": "x"} + }, + ) + + assert result == {"seconds": "8", "size": "override", "extra": "x"} + assert "extra_body" not in result + + +def test_optional__no_extra_body_returns_mapped_unchanged(): + config = _config({"seconds": "8"}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={"seconds": "8"}, + ) + + assert result == {"seconds": "8"} + + +def test_optional__non_dict_extra_body_ignored(): + config = _config({"seconds": "8"}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={"seconds": "8", "extra_body": None}, + ) + + assert result == {"seconds": "8"} From 2cf565ae282584a74d64d92bd2d59fe9a12a5484 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 29 Jun 2026 09:22:58 +0530 Subject: [PATCH 10/37] test(batches): add 1:1 test file scaffold for batches component paths (#30529) * test(batches): add 1:1 test file scaffold for batches component paths Co-authored-by: Cursor * Add harness test for create batch endpoint * Add retrieve endpoint harness tests * Add list endpoint harness tests * Add cancel endpoint harness tests * Add cancel endpoint harness tests * Add test for litellm/batches/main.py * Add test for litellm/tests/test_litellm/batches/test_batch_utils.py * Add handler and transformation tests for all providers * Fix: run batches tests in cicd * fix(tests): remove azure/__init__.py that shadowed azure namespace package Adding __init__.py to tests/test_litellm/llms/azure/ caused pytest to insert tests/test_litellm/llms/ into sys.path[0], making our empty azure/ dir shadow the real azure-identity namespace package. Any test that patched azure.identity.* would then fail with AttributeError. * style(tests): apply ruff format to test_batch_utils.py Base migrated the formatter from black to ruff format (#31317); reformat the batches scaffold test file to match. --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .github/workflows/test-unit-misc.yml | 1 + .../workflows/test-unit-proxy-endpoints.yml | 1 + tests/test_litellm/batches/__init__.py | 0 .../test_litellm/batches/test_batch_utils.py | 735 ++++++ tests/test_litellm/batches/test_main.py | 744 ++++++ tests/test_litellm/llms/anthropic/__init__.py | 0 .../llms/anthropic/batches/__init__.py | 0 .../llms/anthropic/batches/test_handler.py | 286 +++ .../anthropic/batches/test_transformation.py | 650 ++++++ .../llms/azure/batches/__init__.py | 0 .../llms/azure/batches/test_handler.py | 491 ++++ tests/test_litellm/llms/base_llm/__init__.py | 0 .../llms/base_llm/batches/__init__.py | 0 .../batches/base_batches_config_test.py | 128 ++ .../base_llm/batches/test_transformation.py | 231 ++ tests/test_litellm/llms/bedrock/__init__.py | 0 .../llms/bedrock/batches/__init__.py | 0 .../bedrock/batches/test_transformation.py | 684 ++++++ .../llms/vertex_ai/batches/__init__.py | 0 .../llms/vertex_ai/batches/test_handler.py | 805 +++++++ .../vertex_ai/batches/test_transformation.py | 396 ++++ .../proxy/batches_endpoints/__init__.py | 0 .../proxy/batches_endpoints/test_endpoints.py | 2026 +++++++++++++++++ 23 files changed, 7178 insertions(+) create mode 100644 tests/test_litellm/batches/__init__.py create mode 100644 tests/test_litellm/batches/test_batch_utils.py create mode 100644 tests/test_litellm/batches/test_main.py create mode 100644 tests/test_litellm/llms/anthropic/__init__.py create mode 100644 tests/test_litellm/llms/anthropic/batches/__init__.py create mode 100644 tests/test_litellm/llms/anthropic/batches/test_handler.py create mode 100644 tests/test_litellm/llms/anthropic/batches/test_transformation.py create mode 100644 tests/test_litellm/llms/azure/batches/__init__.py create mode 100644 tests/test_litellm/llms/azure/batches/test_handler.py create mode 100644 tests/test_litellm/llms/base_llm/__init__.py create mode 100644 tests/test_litellm/llms/base_llm/batches/__init__.py create mode 100644 tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py create mode 100644 tests/test_litellm/llms/base_llm/batches/test_transformation.py create mode 100644 tests/test_litellm/llms/bedrock/__init__.py create mode 100644 tests/test_litellm/llms/bedrock/batches/__init__.py create mode 100644 tests/test_litellm/llms/bedrock/batches/test_transformation.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/__init__.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/test_handler.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/test_transformation.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/test_endpoints.py diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index d411d996d8e..7c3b195f0ad 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -22,6 +22,7 @@ jobs: uses: ./.github/workflows/_test-unit-base.yml with: test-path: >- + tests/test_litellm/batches tests/test_litellm/secret_managers tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index 23ca7a8f2f9..cbb36eebdb9 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -31,6 +31,7 @@ jobs: tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/openai_files_endpoint + tests/test_litellm/proxy/batches_endpoints tests/test_litellm/proxy/video_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/image_endpoints diff --git a/tests/test_litellm/batches/__init__.py b/tests/test_litellm/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py new file mode 100644 index 00000000000..3aebfcb911e --- /dev/null +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -0,0 +1,735 @@ +""" +Unit tests for litellm/batches/batch_utils.py + +batch_utils.py is the batch cost/usage/parsing layer: it turns a batch output +JSONL into spend (cost), token usage, and the list of models seen, and counts +tokens in batch *input* files for rate limiting. A silent bug here mis-bills +real money or lets callers slip past TPM limits, so these tests assert exact +numeric results rather than "ran without error". + +Pure functions (parsing, token math, credential extraction, success checks) run +for real with exact-value assertions. The few true external seams - the cost +maps (litellm.completion_cost, batch_cost_calculator), the tokenizer +(token_counter), and remote file fetch (afile_content) - are mocked with +deterministic stand-ins so the arithmetic under test is the only variable. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.batches.batch_utils as bu +from litellm.types.utils import Usage + +# --------------------------------------------------------------------------- # +# Builders for batch OUTPUT file rows. +# Shape: {"response": {"status_code": 200, "body": {... "usage": {...}}}} +# --------------------------------------------------------------------------- # + + +def _usage(p, c, t=None): + return { + "prompt_tokens": p, + "completion_tokens": c, + "total_tokens": t if t is not None else p + c, + } + + +def _success_row(model="gpt-4o", usage=None, **body_extra): + body = {"model": model, **body_extra} + if usage is not None: + body["usage"] = usage + return {"response": {"status_code": 200, "body": body}} + + +def _failed_row(status_code=500, model="gpt-4o"): + return {"response": {"status_code": status_code, "body": {"model": model}}} + + +# =========================================================================== # +# _batch_response_was_successful +# =========================================================================== # + + +@pytest.mark.parametrize( + "row,expected", + [ + ({"response": {"status_code": 200}}, True), + ({"response": {"status_code": 500}}, False), + ({"response": {"status_code": 429}}, False), + ({"response": {}}, False), # no status_code + ({}, False), # no response + ({"response": None}, False), # null response + ], +) +def test_batch_response_was_successful(row, expected): + assert bu._batch_response_was_successful(row) is expected + + +# =========================================================================== # +# _get_response_from_batch_job_output_file +# =========================================================================== # + + +def test_get_response_body_present(): + row = {"response": {"body": {"model": "gpt-4o", "usage": {"x": 1}}}} + assert bu._get_response_from_batch_job_output_file(row) == { + "model": "gpt-4o", + "usage": {"x": 1}, + } + + +@pytest.mark.parametrize( + "row", + [ + {}, # no response + {"response": {}}, # no body + {"response": None}, # null response + {"response": {"body": None}}, # null body + ], +) +def test_get_response_body_missing_returns_empty(row): + assert bu._get_response_from_batch_job_output_file(row) == {} + + +# =========================================================================== # +# _get_batch_job_usage_from_response_body +# =========================================================================== # + + +def test_get_usage_from_response_body(): + usage = bu._get_batch_job_usage_from_response_body( + {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}} + ) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 10, + 5, + 15, + ) + + +def test_get_usage_from_response_body_missing_is_zero(): + usage = bu._get_batch_job_usage_from_response_body({}) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 0, + 0, + 0, + ) + + +# =========================================================================== # +# _get_file_content_as_dictionary (JSONL parsing) +# =========================================================================== # + + +def test_parse_jsonl_multiple_lines(): + content = b'{"a": 1}\n{"b": 2}\n{"c": 3}' + assert bu._get_file_content_as_dictionary(content) == [ + {"a": 1}, + {"b": 2}, + {"c": 3}, + ] + + +def test_parse_jsonl_trailing_newline_skipped(): + # outer content is stripped; the trailing-newline empty line is dropped. + content = b'{"a": 1}\n{"b": 2}\n' + assert bu._get_file_content_as_dictionary(content) == [{"a": 1}, {"b": 2}] + + +def test_parse_jsonl_empty_content_is_empty_list(): + assert bu._get_file_content_as_dictionary(b"") == [] + + +def test_parse_jsonl_malformed_raises(): + with pytest.raises(Exception): + bu._get_file_content_as_dictionary(b"not valid json") + + +# =========================================================================== # +# _iter_batch_input_lines / _iter_batch_input_entries (JSONL parsing) +# =========================================================================== # + + +def test_iter_input_lines_skips_blank_and_strips(): + content = b'{"a":1}\n\n \n{"b":2}\n' + assert list(bu._iter_batch_input_lines(content)) == [b'{"a":1}', b'{"b":2}'] + + +def test_iter_input_lines_handles_missing_trailing_newline(): + assert list(bu._iter_batch_input_lines(b'{"a":1}')) == [b'{"a":1}'] + + +def test_iter_input_lines_empty(): + assert list(bu._iter_batch_input_lines(b"")) == [] + + +def test_iter_input_entries_parses_each_row(): + content = b'{"body": {"model": "gpt-4o"}}\n{"body": {"model": "claude-3"}}\n' + assert list(bu._iter_batch_input_entries(content)) == [ + {"body": {"model": "gpt-4o"}}, + {"body": {"model": "claude-3"}}, + ] + + +def test_iter_input_entries_raises_on_malformed_line(): + # _iter_batch_input_entries raises on a bad row; callers that must survive + # bad rows iterate _iter_batch_input_lines and parse per-row instead. + with pytest.raises(Exception): + list(bu._iter_batch_input_entries(b'{"ok":1}\nnot-json\n')) + + +# =========================================================================== # +# _estimate_batch_entry_tokens (regression: an uncountable/malformed row must +# never contribute zero tokens, or a crafted batch could evade the TPM limit) +# =========================================================================== # + + +def test_estimate_tokens_scales_with_size(): + # 4 bytes per token, floored, with a minimum of 1. + assert bu._estimate_batch_entry_tokens(b"a" * 40) == 10 + + +def test_estimate_tokens_never_zero_for_short_rows(): + assert bu._estimate_batch_entry_tokens(b"") == 1 + assert bu._estimate_batch_entry_tokens(b"abc") == 1 + + +# =========================================================================== # +# _get_batch_models_from_file_content (output file) +# =========================================================================== # + + +def test_output_models_uses_model_name_override(): + # model_name short-circuits: content is ignored entirely. + assert bu._get_batch_models_from_file_content([_success_row(model="ignored")], model_name="forced-model") == [ + "forced-model" + ] + + +def test_output_models_collects_from_successful_only(): + rows = [ + _success_row(model="gpt-4o"), + _failed_row(model="should-be-skipped"), + _success_row(model="claude-3"), + ] + assert bu._get_batch_models_from_file_content(rows) == ["gpt-4o", "claude-3"] + + +def test_output_models_skips_successful_without_model(): + rows = [{"response": {"status_code": 200, "body": {}}}] + assert bu._get_batch_models_from_file_content(rows) == [] + + +# =========================================================================== # +# _extract_file_access_credentials +# =========================================================================== # + + +def test_extract_credentials_only_known_keys(): + params = { + "api_key": "sk-1", + "api_base": "https://b", + "vertex_project": "proj", + "model": "gpt-4o", # not a credential key + "unrelated": "x", + } + assert bu._extract_file_access_credentials(params) == { + "api_key": "sk-1", + "api_base": "https://b", + "vertex_project": "proj", + } + + +@pytest.mark.parametrize("params", [None, {}]) +def test_extract_credentials_empty(params): + assert bu._extract_file_access_credentials(params) == {} + + +def test_extract_credentials_all_supported_keys(): + keys = { + "api_key", + "api_base", + "api_version", + "organization", + "azure_ad_token", + "azure_ad_token_provider", + "vertex_project", + "vertex_location", + "vertex_credentials", + "timeout", + "max_retries", + } + params = {k: f"val-{k}" for k in keys} + assert bu._extract_file_access_credentials(params) == params + + +# =========================================================================== # +# _count_prompt_or_input_tokens (regression-critical: list[list[int]] used to +# count as zero and let callers slip past TPM limits). token_counter stubbed to +# len(text) so every shape has an exact expected value. +# =========================================================================== # + + +@pytest.fixture +def fake_token_counter(monkeypatch): + def _tc(model=None, text=None, messages=None, **kw): + if messages is not None: + return len(messages) + if text is not None: + return len(text) + return 0 + + monkeypatch.setattr(bu, "token_counter", _tc) + return _tc + + +def test_count_tokens_str(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", "hello") == 5 # len("hello") + + +def test_count_tokens_list_of_str(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", ["ab", "cde"]) == 5 # 2 + 3 + + +def test_count_tokens_list_of_int(fake_token_counter): + # pre-tokenized prompt: each int counts as one token. + assert bu._count_prompt_or_input_tokens("m", [1, 2, 3, 4]) == 4 + + +def test_count_tokens_list_of_list_of_int(fake_token_counter): + # the bug-fix shape: nested pre-tokenized prompts, each int = 1 token. + assert bu._count_prompt_or_input_tokens("m", [[1, 2, 3], [4, 5]]) == 5 + + +def test_count_tokens_mixed_nested(fake_token_counter): + # nested list with ints + a string: 2 ints (=2) + len("xyz")=3 -> 5 + assert bu._count_prompt_or_input_tokens("m", [[1, 2, "xyz"]]) == 5 + + +def test_count_tokens_unsupported_shape_is_zero(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", 12345) == 0 + assert bu._count_prompt_or_input_tokens("m", {"a": 1}) == 0 + + +# =========================================================================== # +# _count_entry_tokens (per-entry rate-limit token counting). The individual +# prompt/input/embedding shapes are covered in test_batch_file_validation.py; +# here we pin the body-field precedence and the empty/fallback behavior. +# =========================================================================== # + + +def test_count_entry_messages_path(fake_token_counter): + entry = {"body": {"model": "gpt-4o", "messages": [{"role": "user"}, {"role": "x"}]}} + assert bu._count_entry_tokens(entry) == 2 # len(messages) + + +def test_count_entry_prompt_path(fake_token_counter): + assert bu._count_entry_tokens({"body": {"model": "gpt-4o", "prompt": "abcd"}}) == 4 + + +def test_count_entry_input_path(fake_token_counter): + assert bu._count_entry_tokens({"body": {"model": "gpt-4o", "input": "ab"}}) == 2 + + +def test_count_entry_messages_beats_prompt(fake_token_counter): + # messages present -> prompt/input are ignored (messages is checked first). + entry = { + "body": { + "model": "gpt-4o", + "messages": [{"role": "user"}], + "prompt": "this-should-be-ignored", + } + } + assert bu._count_entry_tokens(entry) == 1 + + +def test_count_entry_prompt_beats_input(fake_token_counter): + entry = {"body": {"model": "gpt-4o", "prompt": "abc", "input": "this-is-longer"}} + assert bu._count_entry_tokens(entry) == 3 + + +def test_count_entry_empty_body_is_zero(fake_token_counter): + assert bu._count_entry_tokens({"body": {}}) == 0 + assert bu._count_entry_tokens({}) == 0 + + +def test_count_entry_uses_model_name_fallback(monkeypatch): + # No body.model -> the model_name argument is forwarded to the token counter. + captured = {} + + def _tc(model=None, text=None, messages=None, **kw): + captured["model"] = model + return len(text or "") + + monkeypatch.setattr(bu, "token_counter", _tc) + bu._count_entry_tokens({"body": {"prompt": "ab"}}, model_name="fallback-model") + assert captured["model"] == "fallback-model" + + +# =========================================================================== # +# _get_batch_job_total_usage_from_file_content (output usage aggregation) +# =========================================================================== # + + +def test_total_usage_sums_successful_only(): + rows = [ + _success_row(usage=_usage(10, 5)), # 15 + _failed_row(), # excluded + _success_row(usage=_usage(20, 10)), # 30 + ] + usage = bu._get_batch_job_total_usage_from_file_content(rows) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 30, + 15, + 45, + ) + + +def test_total_usage_empty_is_zero(): + usage = bu._get_batch_job_total_usage_from_file_content([]) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 0, + 0, + 0, + ) + + +# =========================================================================== # +# _get_batch_job_cost_from_file_content (cost maps mocked) +# =========================================================================== # + + +def test_cost_from_content_completion_cost_path(monkeypatch): + # model_info is None -> litellm.completion_cost per successful row. + calls = [] + + def _completion_cost(**kw): + calls.append(kw) + return 0.5 + + monkeypatch.setattr(litellm, "completion_cost", _completion_cost) + rows = [ + _success_row(usage=_usage(10, 5)), + _failed_row(), # excluded -> not costed + _success_row(usage=_usage(20, 10)), + ] + + total = bu._get_batch_job_cost_from_file_content(rows, custom_llm_provider="openai") + + assert total == 1.0 # 2 successful * 0.5 + assert len(calls) == 2 # failed row not costed + + +def test_cost_from_content_model_info_path(monkeypatch): + # model_info set -> batch_cost_calculator(prompt_cost, completion_cost). + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.1, 0.2)) + rows = [ + _success_row(usage=_usage(10, 5)), + _success_row(usage=_usage(20, 10)), + ] + + total = bu._get_batch_job_cost_from_file_content( + rows, + custom_llm_provider="openai", + model_info={"input_cost_per_token": 0.0}, # type: ignore[arg-type] # truthy -> model_info path + ) + + assert total == pytest.approx(0.6) # 2 * (0.1 + 0.2) + + +# =========================================================================== # +# _batch_cost_calculator (dispatch: vertex-disable-transform vs generic) +# =========================================================================== # + + +def test_batch_cost_calculator_generic_path(monkeypatch): + monkeypatch.setattr(bu, "_get_batch_job_cost_from_file_content", lambda **kw: 4.2) + assert bu._batch_cost_calculator([], custom_llm_provider="openai", model_name="gpt-4o") == 4.2 + + +def test_batch_cost_calculator_vertex_disable_transform_path(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, + "calculate_vertex_ai_batch_cost_and_usage", + lambda content, model: (9.9, Usage()), + ) + # generic path must NOT be taken + monkeypatch.setattr( + bu, + "_get_batch_job_cost_from_file_content", + lambda **kw: pytest.fail("generic path should not run"), + ) + + cost = bu._batch_cost_calculator([], custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001") + assert cost == 9.9 + + +# =========================================================================== # +# calculate_vertex_ai_batch_cost_and_usage (usageMetadata aggregation) +# =========================================================================== # + + +def test_vertex_cost_and_usage_aggregation(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.1, 0.2)) + responses = [ + { + "response": { + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 5, + "totalTokenCount": 15, + } + } + }, + { + "response": { + "usageMetadata": { + "promptTokenCount": 20, + "candidatesTokenCount": 10, + "totalTokenCount": 30, + } + } + }, + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + + assert cost == pytest.approx(0.6) # 2 * (0.1 + 0.2) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 30, + 15, + 45, + ) + + +def test_vertex_cost_skips_none_response_body(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (1.0, 0.0)) + responses = [ + {"response": None}, # skipped + { + "response": { + "usageMetadata": { + "promptTokenCount": 7, + "candidatesTokenCount": 3, + "totalTokenCount": 10, + } + } + }, + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + + assert cost == pytest.approx(1.0) # only one line costed + assert usage.total_tokens == 10 + + +def test_vertex_usage_total_token_fallback(monkeypatch): + # no totalTokenCount -> falls back to prompt + completion. + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.0, 0.0)) + responses = [{"response": {"usageMetadata": {"promptTokenCount": 8, "candidatesTokenCount": 4}}}] + + _, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + assert usage.total_tokens == 12 + + +def test_vertex_cost_error_in_line_is_swallowed(monkeypatch): + # a cost error on one line must not abort aggregation; usage still tallies. + import litellm.cost_calculator as cc + + def _boom(**kw): + raise RuntimeError("price map miss") + + monkeypatch.setattr(cc, "batch_cost_calculator", _boom) + responses = [ + { + "response": { + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 5, + "totalTokenCount": 10, + } + } + } + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + assert cost == 0.0 + assert usage.total_tokens == 10 + + +# =========================================================================== # +# calculate_batch_cost_and_usage (async orchestrator) +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_calculate_batch_cost_and_usage_orchestration(monkeypatch): + rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] + monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 2.5) + monkeypatch.setattr( + bu, + "_get_batch_job_total_usage_from_file_content", + lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + + cost, usage, models = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=rows, custom_llm_provider="openai" + ) + + assert cost == 2.5 + assert usage.total_tokens == 15 + assert models == ["gpt-4o"] # real _get_batch_models_from_file_content + + +# =========================================================================== # +# _get_batch_output_file_content_as_dictionary (file fetch + credential merge) +# =========================================================================== # + + +def _batch(output_file_id): + from litellm.types.llms.openai import Batch + + return Batch( + id="b", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="f", + object="batch", + status="completed", + output_file_id=output_file_id, + ) + + +@pytest.mark.asyncio +async def test_output_file_content_vertex_raises(): + with pytest.raises(ValueError, match="Vertex AI does not support"): + await bu._get_batch_output_file_content_as_dictionary(_batch("of"), custom_llm_provider="vertex_ai") + + +@pytest.mark.asyncio +async def test_output_file_content_no_output_file_id_raises(): + with pytest.raises(ValueError, match="Output file id is None"): + await bu._get_batch_output_file_content_as_dictionary(_batch(None), custom_llm_provider="openai") + + +@pytest.mark.asyncio +async def test_output_file_content_fetches_and_parses(monkeypatch): + import litellm.files.main as files_main + import litellm.proxy.openai_files_endpoints.common_utils as cu + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b'{"a": 1}\n{"b": 2}'})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + monkeypatch.setattr(cu, "_is_base64_encoded_unified_file_id", lambda fid: False) + + result = await bu._get_batch_output_file_content_as_dictionary( + _batch("file-out"), + custom_llm_provider="azure", + litellm_params={"api_key": "sk-az", "api_base": "https://az", "model": "x"}, + ) + + assert result == [{"a": 1}, {"b": 2}] + # afile_content received the file id + extracted credentials (not "model"). + assert captured["file_id"] == "file-out" + assert captured["custom_llm_provider"] == "azure" + assert captured["api_key"] == "sk-az" + assert captured["api_base"] == "https://az" + assert "model" not in captured + + +@pytest.mark.asyncio +async def test_output_file_content_unified_file_id_extraction(monkeypatch): + # a base64 unified id carries the real provider file id inside + # "llm_output_file_id,;" - it must be unwrapped before the fetch. + import litellm.files.main as files_main + import litellm.proxy.openai_files_endpoints.common_utils as cu + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b'{"a": 1}'})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + monkeypatch.setattr( + cu, + "_is_base64_encoded_unified_file_id", + lambda fid: "litellm_proxy;llm_output_file_id,real-file-99;rest", + ) + + await bu._get_batch_output_file_content_as_dictionary(_batch("encoded-blob"), custom_llm_provider="openai") + + assert captured["file_id"] == "real-file-99" + + +# =========================================================================== # +# _handle_completed_batch (async orchestrator: fetch -> cost/usage/models) +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_handle_completed_batch_orchestration(monkeypatch): + rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] + + async def fake_get_content(batch, custom_llm_provider, litellm_params=None): + return rows + + monkeypatch.setattr(bu, "_get_batch_output_file_content_as_dictionary", fake_get_content) + monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 3.3) + monkeypatch.setattr( + bu, + "_get_batch_job_total_usage_from_file_content", + lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + + cost, usage, models = await bu._handle_completed_batch(_batch("of"), custom_llm_provider="openai") + + assert cost == 3.3 + assert usage.total_tokens == 15 + assert models == ["gpt-4o"] + + +# =========================================================================== # +# Remaining branch: vertex usage disable-transform path. +# +# NOTE: the error path of _get_batch_job_cost_from_file_content (its `raise e`) +# is intentionally NOT tested: the preceding line logs via +# `verbose_logger.error("...", e)`, which passes the exception as a logging +# format-arg with no placeholder and itself raises TypeError under +# logging.raiseExceptions, masking the original error. Asserting that masked +# behavior would lock a source bug; left uncovered on purpose. +# =========================================================================== # + + +def test_total_usage_vertex_disable_transform_path(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, + "calculate_vertex_ai_batch_cost_and_usage", + lambda content, model: ( + 0.0, + Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3), + ), + ) + + usage = bu._get_batch_job_total_usage_from_file_content([], custom_llm_provider="vertex_ai", model_name="gemini-x") + assert usage.total_tokens == 3 diff --git a/tests/test_litellm/batches/test_main.py b/tests/test_litellm/batches/test_main.py new file mode 100644 index 00000000000..1f7a91a5511 --- /dev/null +++ b/tests/test_litellm/batches/test_main.py @@ -0,0 +1,744 @@ +""" +Provider-dispatch contract tests for litellm/batches/main.py + +main.py is the SDK layer beneath the proxy batch endpoints: each of +create/retrieve/list/cancel_batch is a switch on `custom_llm_provider` (and, for +create/retrieve, on whether a provider-config + model is present) that hands off +to exactly one provider handler. These tests lock that dispatch: + + 1. DISPATCH - exactly which provider seam fired (openai_batches_instance vs + azure vs vertex vs anthropic vs base_llm_http_handler vs the + Bedrock ARN handlers), with every sibling seam asserted NOT + called. A reordered/negated branch flips this. + 2. PAYLOAD - the request object (CreateBatchRequest/RetrieveBatchRequest/...) + and the _is_async flag forwarded to the handler. + 3. RESULT - the handler's return value is what the function returns. + 4. DELEGATION - the async wrappers (a*_batch) forward to the sync function in an + executor with the right "_is_async" flag, and pass the result + back untouched. + +Only the provider handler instances are mocked (true network boundaries). The +real public functions run (including the @client decorator) so dispatch reflects +production. Provider env vars are not required: missing creds resolve to None and +flow through harmlessly because the handler is mocked. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.batches.main as bm + + +# --------------------------------------------------------------------------- # +# Seam harness - one mock per provider handler instance + the Bedrock ARN +# handler. Each handler method auto-returns a unique sentinel (its +# return_value), so "result is seam..return_value" verifies dispatch. +# --------------------------------------------------------------------------- # + + +@dataclass +class Seams: + openai: MagicMock + azure: MagicMock + vertex: MagicMock + anthropic: MagicMock + base_http: MagicMock + bedrock_arn: MagicMock + + +@pytest.fixture +def seams(): + openai_i = MagicMock(name="openai_batches_instance") + azure_i = MagicMock(name="azure_batches_instance") + vertex_i = MagicMock(name="vertex_ai_batches_instance") + anthropic_i = MagicMock(name="anthropic_batches_instance") + base_http = MagicMock(name="base_llm_http_handler") + bedrock_arn = MagicMock(name="BedrockBatchesHandler") + + with ExitStack() as stack: + stack.enter_context(patch.object(bm, "openai_batches_instance", openai_i)) + stack.enter_context(patch.object(bm, "azure_batches_instance", azure_i)) + stack.enter_context(patch.object(bm, "vertex_ai_batches_instance", vertex_i)) + stack.enter_context( + patch.object(bm, "anthropic_batches_instance", anthropic_i) + ) + stack.enter_context(patch.object(bm, "base_llm_http_handler", base_http)) + stack.enter_context(patch.object(bm, "BedrockBatchesHandler", bedrock_arn)) + yield Seams( + openai=openai_i, + azure=azure_i, + vertex=vertex_i, + anthropic=anthropic_i, + base_http=base_http, + bedrock_arn=bedrock_arn, + ) + + +# Every handler method across all provider instances - used to assert +# "no sibling seam fired" exhaustively. +def _all_seam_methods(seams: Seams, op: str): + return [ + getattr(seams.openai, op), + getattr(seams.azure, op), + getattr(seams.vertex, op), + getattr(seams.anthropic, op), + getattr(seams.base_http, op), + ] + + +def _assert_only(fired, seams: Seams, op: str): + """Assert `fired` was called exactly once and every other op seam was not.""" + assert fired.call_count == 1 + for m in _all_seam_methods(seams, op): + if m is not fired: + m.assert_not_called() + + +CREATE_KW: Dict[str, Any] = dict( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id="file-abc", +) + + +# =========================================================================== # +# create_batch +# =========================================================================== # + + +def test_create__openai_dispatch_and_payload(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="openai") + + # DISPATCH + RESULT + assert result is seams.openai.create_batch.return_value + _assert_only(seams.openai.create_batch, seams, "create_batch") + seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + + # PAYLOAD - request object built from the call, sync flag off. + kw = seams.openai.create_batch.call_args.kwargs + assert kw["create_batch_data"] == { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", + "metadata": None, + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + assert kw["timeout"] == 600.0 + + +def test_create__hosted_vllm_routes_to_openai_instance(seams): + """hosted_vllm is in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, so it shares + the openai handler. Locks that set membership.""" + result = bm.create_batch(**CREATE_KW, custom_llm_provider="hosted_vllm") + + assert result is seams.openai.create_batch.return_value + _assert_only(seams.openai.create_batch, seams, "create_batch") + + +def test_create__azure_dispatch(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="azure") + + assert result is seams.azure.create_batch.return_value + _assert_only(seams.azure.create_batch, seams, "create_batch") + + +def test_create__vertex_ai_dispatch(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="vertex_ai") + + assert result is seams.vertex.create_batch.return_value + _assert_only(seams.vertex.create_batch, seams, "create_batch") + + +def test_create__provider_config_routes_to_base_http_handler(seams): + """model + a provider batches config (bedrock-style) routes to the generic + base_llm_http_handler, NOT the per-provider instance.""" + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + result = bm.create_batch( + **CREATE_KW, custom_llm_provider="bedrock", model="bedrock/my-batch-model" + ) + + assert result is seams.base_http.create_batch.return_value + _assert_only(seams.base_http.create_batch, seams, "create_batch") + + +def test_create__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.create_batch(**CREATE_KW, custom_llm_provider="cohere") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "create_batch"): + m.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__async_path_propagates_is_async(seams): + """Through the real async wrapper, the handler is invoked with _is_async=True. + (Calling the @client sync create_batch with acreate_batch=True directly is not + a real code path - logging-obj setup only happens on the async wrapper path.)""" + await bm.acreate_batch(**CREATE_KW, custom_llm_provider="openai") + + assert seams.openai.create_batch.call_args.kwargs["_is_async"] is True + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve__openai_dispatch_and_payload(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="openai") + + assert result is seams.openai.retrieve_batch.return_value + _assert_only(seams.openai.retrieve_batch, seams, "retrieve_batch") + + kw = seams.openai.retrieve_batch.call_args.kwargs + assert kw["retrieve_batch_data"] == { + "batch_id": "batch-1", + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + + +def test_retrieve__hosted_vllm_routes_to_openai_instance(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="hosted_vllm") + + assert result is seams.openai.retrieve_batch.return_value + _assert_only(seams.openai.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__azure_dispatch(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="azure") + + assert result is seams.azure.retrieve_batch.return_value + _assert_only(seams.azure.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__vertex_ai_dispatch(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="vertex_ai") + + assert result is seams.vertex.retrieve_batch.return_value + _assert_only(seams.vertex.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__anthropic_dispatch(seams): + """anthropic is retrieve-capable (not in create's provider set).""" + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="anthropic") + + assert result is seams.anthropic.retrieve_batch.return_value + _assert_only(seams.anthropic.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__provider_config_routes_to_base_http_handler(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + result = bm.retrieve_batch( + batch_id="batch-1", + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + ) + + assert result is seams.base_http.retrieve_batch.return_value + _assert_only(seams.base_http.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__bedrock_async_invoke_arn(seams): + arn = "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123" + result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") + + seams.bedrock_arn._handle_async_invoke_status.assert_called_once() + assert result is seams.bedrock_arn._handle_async_invoke_status.return_value + # provider instances untouched. + for m in _all_seam_methods(seams, "retrieve_batch"): + m.assert_not_called() + + +def test_retrieve__bedrock_model_invocation_job_arn(seams): + arn = "arn:aws:bedrock:us-east-1:123456789012:model-invocation-job/xyz789" + result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") + + seams.bedrock_arn._handle_model_invocation_job_status.assert_called_once() + assert ( + result is seams.bedrock_arn._handle_model_invocation_job_status.return_value + ) + seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + + +def test_retrieve__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="cohere") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "retrieve_batch"): + m.assert_not_called() + + +# =========================================================================== # +# list_batches (supported: openai, hosted_vllm, azure, vertex_ai) +# =========================================================================== # + + +def test_list__openai_dispatch_and_payload(seams): + result = bm.list_batches(custom_llm_provider="openai", after="cur", limit=5) + + assert result is seams.openai.list_batches.return_value + _assert_only(seams.openai.list_batches, seams, "list_batches") + + kw = seams.openai.list_batches.call_args.kwargs + assert kw["after"] == "cur" + assert kw["limit"] == 5 + assert kw["_is_async"] is False + + +def test_list__hosted_vllm_routes_to_openai_instance(seams): + result = bm.list_batches(custom_llm_provider="hosted_vllm") + + assert result is seams.openai.list_batches.return_value + _assert_only(seams.openai.list_batches, seams, "list_batches") + + +def test_list__azure_dispatch(seams): + result = bm.list_batches(custom_llm_provider="azure") + + assert result is seams.azure.list_batches.return_value + _assert_only(seams.azure.list_batches, seams, "list_batches") + + +def test_list__vertex_ai_dispatch(seams): + result = bm.list_batches(custom_llm_provider="vertex_ai") + + assert result is seams.vertex.list_batches.return_value + _assert_only(seams.vertex.list_batches, seams, "list_batches") + + +def test_list__unsupported_provider_raises_badrequest(seams): + # anthropic supports retrieve but NOT list - good negative case. + with pytest.raises(litellm.exceptions.BadRequestError): + bm.list_batches(custom_llm_provider="anthropic") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "list_batches"): + m.assert_not_called() + + +# =========================================================================== # +# cancel_batch (supported: openai, hosted_vllm, azure, vertex_ai; no @client) +# =========================================================================== # + + +def test_cancel__openai_dispatch_and_payload(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="openai") + + assert result is seams.openai.cancel_batch.return_value + _assert_only(seams.openai.cancel_batch, seams, "cancel_batch") + + kw = seams.openai.cancel_batch.call_args.kwargs + assert kw["cancel_batch_data"] == { + "batch_id": "batch-1", + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + + +def test_cancel__azure_dispatch(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="azure") + + assert result is seams.azure.cancel_batch.return_value + _assert_only(seams.azure.cancel_batch, seams, "cancel_batch") + + +def test_cancel__vertex_ai_dispatch(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="vertex_ai") + + assert result is seams.vertex.cancel_batch.return_value + _assert_only(seams.vertex.cancel_batch, seams, "cancel_batch") + + +def test_cancel__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.cancel_batch(batch_id="batch-1", custom_llm_provider="cohere") + + for m in _all_seam_methods(seams, "cancel_batch"): + m.assert_not_called() + + +def test_cancel__async_flag_propagates_is_async(seams): + bm.cancel_batch( + batch_id="batch-1", custom_llm_provider="openai", acancel_batch=True + ) + + assert seams.openai.cancel_batch.call_args.kwargs["_is_async"] is True + + +# =========================================================================== # +# Async wrappers - delegate to the sync function in an executor, set the right +# "_is_async" flag, and return the result untouched. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_acreate_batch_delegates_to_create_batch(): + with patch.object(bm, "create_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.acreate_batch(**CREATE_KW, custom_llm_provider="openai") + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("acreate_batch") is True + # positional handoff: (completion_window, endpoint, input_file_id, provider, ...) + assert m.call_args.args[0] == "24h" + assert m.call_args.args[2] == "file-abc" + assert m.call_args.args[3] == "openai" + + +@pytest.mark.asyncio +async def test_aretrieve_batch_delegates_to_retrieve_batch(): + with patch.object(bm, "retrieve_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.aretrieve_batch( + batch_id="batch-1", custom_llm_provider="azure" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("aretrieve_batch") is True + assert m.call_args.args[0] == "batch-1" + assert m.call_args.args[1] == "azure" + + +@pytest.mark.asyncio +async def test_alist_batches_delegates_to_list_batches(): + with patch.object(bm, "list_batches", MagicMock(return_value="SENTINEL")) as m: + result = await bm.alist_batches( + after="cur", limit=3, custom_llm_provider="vertex_ai" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("alist_batches") is True + assert m.call_args.args[0] == "cur" + assert m.call_args.args[1] == 3 + assert m.call_args.args[2] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_acancel_batch_delegates_to_cancel_batch(): + with patch.object(bm, "cancel_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.acancel_batch( + batch_id="batch-1", custom_llm_provider="openai" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("acancel_batch") is True + assert m.call_args.args[0] == "batch-1" + + +# =========================================================================== # +# Credential passthrough - when the caller supplies credentials in kwargs, they +# must reach the provider handler. Explicit kwargs win over litellm.* globals and +# env vars (they are first in each `optional_params.x or litellm.x or env` chain), +# so these assertions are deterministic regardless of the test environment. +# +# The credential-resolution blocks are copy-pasted per provider in EACH of +# create/retrieve/list/cancel, so a regression can land in any one independently; +# every function is checked. +# =========================================================================== # + + +# Distinct values so a cross-wired field (e.g. api_key forwarded as api_base) is +# impossible to miss. +OPENAI_CREDS: Dict[str, Any] = dict( + api_key="sk-user-openai", + api_base="https://openai.user.test", + organization="org-user-123", + max_retries=7, +) +AZURE_CREDS: Dict[str, Any] = dict( + api_key="sk-user-azure", + api_base="https://azure.user.test", + api_version="2024-12-99", +) +VERTEX_CREDS: Dict[str, Any] = dict( + vertex_project="proj-user", + vertex_location="loc-user", + vertex_credentials="cred-user", + api_base="https://vertex.user.test", +) + + +def _sent(mock_method, *keys): + """Subset of the call kwargs limited to `keys`, for exact comparison.""" + kw = mock_method.call_args.kwargs + return {k: kw.get(k) for k in keys} + + +# ---- create_batch ---------------------------------------------------------- # + + +def test_create__openai_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.create_batch, "api_key", "api_base", "organization", "max_retries" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + "max_retries": 7, + } + + +def test_create__azure_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.create_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_create__vertex_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.create_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + "api_base", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + "api_base": "https://vertex.user.test", + } + + +def test_create__provider_config_credentials_passthrough(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + bm.create_batch( + **CREATE_KW, + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + api_key="sk-user-bedrock", + api_base="https://bedrock.user.test", + ) + + assert _sent(seams.base_http.create_batch, "api_key", "api_base") == { + "api_key": "sk-user-bedrock", + "api_base": "https://bedrock.user.test", + } + + +# ---- retrieve_batch -------------------------------------------------------- # + + +def test_retrieve__openai_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.retrieve_batch, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_retrieve__azure_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.retrieve_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_retrieve__vertex_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.retrieve_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +def test_retrieve__anthropic_credentials_passthrough(seams): + bm.retrieve_batch( + batch_id="b1", + custom_llm_provider="anthropic", + api_key="sk-user-anthropic", + api_base="https://anthropic.user.test", + ) + + assert _sent(seams.anthropic.retrieve_batch, "api_key", "api_base") == { + "api_key": "sk-user-anthropic", + "api_base": "https://anthropic.user.test", + } + + +def test_retrieve__provider_config_credentials_passthrough(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + bm.retrieve_batch( + batch_id="b1", + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + api_key="sk-user-bedrock", + api_base="https://bedrock.user.test", + ) + + assert _sent(seams.base_http.retrieve_batch, "api_key", "api_base") == { + "api_key": "sk-user-bedrock", + "api_base": "https://bedrock.user.test", + } + + +# ---- list_batches ---------------------------------------------------------- # + + +def test_list__openai_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.list_batches, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_list__azure_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.list_batches, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_list__vertex_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.list_batches, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +# ---- cancel_batch ---------------------------------------------------------- # + + +def test_cancel__openai_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.cancel_batch, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_cancel__azure_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.cancel_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_cancel__vertex_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.cancel_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +# =========================================================================== # +# _resolve_timeout - pure helper (used by create_batch). +# =========================================================================== # + + +def _params(**kw): + from litellm.types.router import GenericLiteLLMParams + + return GenericLiteLLMParams(**kw) + + +def test_resolve_timeout__explicit_numeric(): + assert bm._resolve_timeout(_params(timeout=30), {}, "openai") == 30.0 + + +def test_resolve_timeout__default_when_unset(): + assert bm._resolve_timeout(_params(), {}, "openai") == 600.0 + + +def test_resolve_timeout__request_timeout_kwarg_fallback(): + assert bm._resolve_timeout(_params(), {"request_timeout": 45}, "openai") == 45.0 + + +def test_resolve_timeout__httpx_timeout_returns_float_read(): + import httpx + + t = httpx.Timeout(99.0, connect=5.0) + resolved = bm._resolve_timeout(_params(timeout=t), {}, "openai") + assert isinstance(resolved, float) + assert resolved == 99.0 diff --git a/tests/test_litellm/llms/anthropic/__init__.py b/tests/test_litellm/llms/anthropic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/anthropic/batches/__init__.py b/tests/test_litellm/llms/anthropic/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/anthropic/batches/test_handler.py b/tests/test_litellm/llms/anthropic/batches/test_handler.py new file mode 100644 index 00000000000..0a472d86257 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/batches/test_handler.py @@ -0,0 +1,286 @@ +""" +Unit tests for litellm/llms/anthropic/batches/handler.py + +AnthropicBatchesHandler is the HTTP/auth glue for retrieving Anthropic Message +Batches. It resolves credentials, builds the retrieve URL + auth headers via the +provider config, fires a single GET against the async httpx client, and hands the +response to the config's transform. These tests mock ONLY the genuine I/O seams - +the async httpx client (network) and credential resolution (secret managers / +env) - and assert exactly which seam fired, with what URL/headers, and that the +parsed result is the LiteLLMBatch the transform produced. + +The sync ``retrieve_batch`` dispatch (``_is_async`` true -> coroutine, false -> +asyncio.run) is exercised directly, mirroring the dispatch-contract discipline in +tests/test_litellm/batches/test_main.py. +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler +from litellm.types.utils import LiteLLMBatch + + +def _ok_batch_response(): + """A real httpx.Response shaped like an Anthropic MessageBatch retrieval.""" + return httpx.Response( + status_code=200, + json={ + "id": "msgbatch_abc", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "request_counts": {"succeeded": 2, "errored": 0}, + }, + request=httpx.Request( + "GET", "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ), + ) + + +@pytest.fixture +def handler(): + return AnthropicBatchesHandler() + + +@pytest.fixture +def patched_client(): + """Patch the async httpx client seam; yield the (fake_client, factory).""" + fake_client = MagicMock() + fake_client.get = AsyncMock(return_value=_ok_batch_response()) + with patch( + "litellm.llms.anthropic.batches.handler.get_async_httpx_client", + return_value=fake_client, + ) as factory: + yield fake_client, factory + + +@pytest.mark.asyncio +async def test_aretrieve_batch_fires_get_with_correct_url_and_headers( + handler, patched_client +): + fake_client, factory = patched_client + + batch = await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + + # The single network seam fired exactly once. + fake_client.get.assert_awaited_once() + _, call_kwargs = fake_client.get.call_args + # Exact URL built by get_retrieve_batch_url. + assert call_kwargs["url"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + # Auth + version + beta headers built by validate_environment. + headers = call_kwargs["headers"] + assert headers["x-api-key"] == "sk-ant-test" + assert headers["anthropic-version"] == "2023-06-01" + assert headers["anthropic-beta"] == "message-batches-2024-09-24" + + # Response parsed through the config transform. + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.status == "completed" + assert batch.request_counts.completed == 2 + + +@pytest.mark.asyncio +async def test_aretrieve_batch_uses_anthropic_provider_for_client( + handler, patched_client +): + from litellm.types.utils import LlmProviders + + _, factory = patched_client + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + _, kwargs = factory.call_args + assert kwargs["llm_provider"] == LlmProviders.ANTHROPIC + + +@pytest.mark.asyncio +async def test_aretrieve_batch_resolves_api_key_from_model_info( + handler, patched_client +): + fake_client, _ = patched_client + # api_key=None -> handler falls back to AnthropicModelInfo.get_api_key(). + with patch.object( + handler.anthropic_model_info, "get_api_key", return_value="sk-from-env" + ): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key=None, + timeout=60.0, + max_retries=0, + ) + _, call_kwargs = fake_client.get.call_args + assert call_kwargs["headers"]["x-api-key"] == "sk-from-env" + + +@pytest.mark.asyncio +async def test_aretrieve_batch_missing_api_key_raises(handler, patched_client): + fake_client, _ = patched_client + # No api_key and resolver yields None -> hard error before any network call. + with patch.object( + handler.anthropic_model_info, "get_api_key", return_value=None + ): + with pytest.raises(ValueError, match="Missing Anthropic API Key"): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key=None, + timeout=60.0, + max_retries=0, + ) + fake_client.get.assert_not_called() + + +@pytest.mark.asyncio +async def test_aretrieve_batch_resolves_default_api_base(handler, patched_client): + fake_client, _ = patched_client + # api_base=None -> resolved via get_api_base() default before URL build. + with patch.object( + handler.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base=None, + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + _, call_kwargs = fake_client.get.call_args + assert call_kwargs["url"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_raises_for_status(handler): + # A non-2xx response must surface via raise_for_status (no silent parse). + error_response = httpx.Response( + status_code=404, + json={"error": "not found"}, + request=httpx.Request( + "GET", "https://api.anthropic.com/v1/messages/batches/missing" + ), + ) + fake_client = MagicMock() + fake_client.get = AsyncMock(return_value=error_response) + with patch( + "litellm.llms.anthropic.batches.handler.get_async_httpx_client", + return_value=fake_client, + ): + with pytest.raises(httpx.HTTPStatusError): + await handler.aretrieve_batch( + batch_id="missing", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_invokes_pre_call_logging(handler, patched_client): + fake_client, _ = patched_client + logging_obj = MagicMock() + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + logging_obj=logging_obj, + ) + logging_obj.pre_call.assert_called_once() + pre_kwargs = logging_obj.pre_call.call_args.kwargs + assert pre_kwargs["input"] == "msgbatch_abc" + assert pre_kwargs["api_key"] == "sk-ant-test" + # The logged api_base is the full retrieve URL, not the bare base. + assert pre_kwargs["additional_args"]["api_base"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_builds_default_logging_obj_when_absent( + handler, patched_client +): + # logging_obj=None -> handler constructs a real Logging object; the call + # must still complete (no AttributeError on a missing logger). + _, _ = patched_client + with patch( + "litellm.litellm_core_utils.litellm_logging.Logging" + ) as logging_cls: + logging_cls.return_value = MagicMock() + batch = await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + logging_obj=None, + ) + logging_cls.assert_called_once() + # call_type wires through to the constructed logging object. + assert logging_cls.call_args.kwargs["call_type"] == "batch_retrieve" + assert batch.id == "msgbatch_abc" + + +# =========================================================================== # +# retrieve_batch dispatch (sync wrapper) +# =========================================================================== # + + +async def test_retrieve_batch_async_returns_coroutine(handler, patched_client): + # _is_async=True -> returns the un-awaited coroutine (caller awaits it). + import asyncio + + coro = handler.retrieve_batch( + _is_async=True, + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + assert asyncio.iscoroutine(coro) + # Await directly - robust under asyncio_mode=auto's session-scoped loop + # (manually driving get_event_loop().run_until_complete() breaks when prior + # async tests in the suite have already used/closed that loop). + batch = await coro + assert batch.id == "msgbatch_abc" + + +def test_retrieve_batch_sync_runs_to_result(handler, patched_client): + # _is_async=False -> asyncio.run(...) returns the resolved LiteLLMBatch. + batch = handler.retrieve_batch( + _is_async=False, + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.status == "completed" diff --git a/tests/test_litellm/llms/anthropic/batches/test_transformation.py b/tests/test_litellm/llms/anthropic/batches/test_transformation.py new file mode 100644 index 00000000000..4a2adb01ea5 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/batches/test_transformation.py @@ -0,0 +1,650 @@ +""" +Unit tests for litellm/llms/anthropic/batches/transformation.py + +AnthropicBatchesConfig is the pure request/response mapping layer for Anthropic +Message Batches. It builds auth headers, constructs the batch create/retrieve +URLs, and (most importantly) maps an Anthropic MessageBatch JSON response into a +LiteLLM/OpenAI ``LiteLLMBatch`` (status mapping, timestamp parsing, request +counts). A silent bug here mis-reports batch status or counts to the caller, so +these tests assert EXACT output values rather than "ran without error". + +Pure transform code runs for real. The only mocked boundaries are the credential +resolvers on AnthropicModelInfo (get_api_base / get_auth_header), which would +otherwise read process env / secret managers - mocking them keeps the URL/header +assertions deterministic without touching production transform logic. +""" + +import os +import sys +import time +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig +from litellm.types.utils import LiteLLMBatch, LlmProviders + + +@pytest.fixture +def config(): + return AnthropicBatchesConfig() + + +def _response(payload): + """A real httpx.Response whose .json() yields ``payload``.""" + return httpx.Response( + status_code=200, + json=payload, + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + +# =========================================================================== # +# custom_llm_provider +# =========================================================================== # + + +def test_custom_llm_provider_is_anthropic(config): + assert config.custom_llm_provider == LlmProviders.ANTHROPIC + + +# =========================================================================== # +# validate_environment (auth + fixed headers + beta header) +# =========================================================================== # + + +def test_validate_environment_builds_headers_with_api_key(config): + headers = config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-test", + ) + assert headers["accept"] == "application/json" + assert headers["anthropic-version"] == "2023-06-01" + assert headers["content-type"] == "application/json" + # Plain api key -> x-api-key auth header. + assert headers["x-api-key"] == "sk-ant-test" + # Beta header is injected when not already present. + assert headers["anthropic-beta"] == "message-batches-2024-09-24" + + +def test_validate_environment_preserves_existing_beta_header(config): + headers = config.validate_environment( + headers={"anthropic-beta": "custom-beta-value"}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-test", + ) + # Existing beta header must NOT be overwritten. + assert headers["anthropic-beta"] == "custom-beta-value" + + +def test_validate_environment_oauth_key_uses_bearer(config): + headers = config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-oat-abc123", + ) + # OAuth tokens map to Authorization: Bearer, not x-api-key. + assert headers["authorization"] == "Bearer sk-ant-oat-abc123" + assert "x-api-key" not in headers + + +def test_validate_environment_missing_key_raises(config): + # No api_key passed and no env credentials -> get_auth_header returns None. + with patch.object( + config.anthropic_model_info, "get_auth_header", return_value=None + ): + with pytest.raises(ValueError, match="Missing Anthropic API Key"): + config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + +# =========================================================================== # +# get_complete_batch_url (batch creation URL) +# =========================================================================== # + + +def test_get_complete_batch_url_appends_path(config): + url = config.get_complete_batch_url( + api_base="https://api.anthropic.com", + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +def test_get_complete_batch_url_strips_trailing_slash(config): + url = config.get_complete_batch_url( + api_base="https://api.anthropic.com/", + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +def test_get_complete_batch_url_already_complete_is_unchanged(config): + complete = "https://proxy.internal/v1/messages/batches" + url = config.get_complete_batch_url( + api_base=complete, + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == complete + + +def test_get_complete_batch_url_uses_default_api_base(config): + # api_base=None -> falls back to get_api_base() default. + with patch.object( + config.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + url = config.get_complete_batch_url( + api_base=None, + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +# =========================================================================== # +# get_retrieve_batch_url (batch retrieval URL + path encoding) +# =========================================================================== # + + +def test_get_retrieve_batch_url_happy_path(config): + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +def test_get_retrieve_batch_url_strips_trailing_slash(config): + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com/", + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +def test_get_retrieve_batch_url_encodes_batch_id(config): + # batch_id is user-controlled; a path-traversal attempt must be percent-encoded. + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="a/b id", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/a%2Fb%20id" + + +def test_get_retrieve_batch_url_rejects_dot_segment(config): + with pytest.raises(ValueError, match="dot path segment"): + config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="..", + optional_params={}, + litellm_params={}, + ) + + +def test_get_retrieve_batch_url_uses_default_api_base(config): + with patch.object( + config.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + url = config.get_retrieve_batch_url( + api_base=None, + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +# =========================================================================== # +# transform_retrieve_batch_request (no-op for Anthropic) +# =========================================================================== # + + +def test_transform_retrieve_batch_request_returns_empty_dict(config): + assert ( + config.transform_retrieve_batch_request( + batch_id="msgbatch_123", optional_params={}, litellm_params={} + ) + == {} + ) + + +# =========================================================================== # +# Unimplemented create-batch methods raise NotImplementedError +# =========================================================================== # + + +def test_transform_create_batch_request_not_implemented(config): + with pytest.raises(NotImplementedError, match="not yet implemented"): + config.transform_create_batch_request( + model="claude-3", + create_batch_data={}, # type: ignore[arg-type] + optional_params={}, + litellm_params={}, + ) + + +def test_transform_create_batch_response_not_implemented(config): + with pytest.raises(NotImplementedError, match="not yet implemented"): + config.transform_create_batch_response( + model="claude-3", + raw_response=_response({}), + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# =========================================================================== # +# transform_retrieve_batch_response (the core mapping - exact values) +# =========================================================================== # + + +def test_transform_retrieve_response_in_progress(config): + raw = _response( + { + "id": "msgbatch_abc", + "processing_status": "in_progress", + "created_at": "2024-09-24T10:00:00Z", + "expires_at": "2024-09-25T10:00:00Z", + "request_counts": { + "processing": 3, + "succeeded": 2, + "errored": 1, + "canceled": 0, + "expired": 0, + }, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.object == "batch" + assert batch.endpoint == "/v1/messages" + assert batch.status == "in_progress" + # output_file_id mirrors the batch id for Anthropic. + assert batch.output_file_id == "msgbatch_abc" + assert batch.input_file_id == "None" + assert batch.completion_window == "24h" + # created_at parsed from ISO8601 (UTC). + assert batch.created_at == 1727172000 + assert batch.expires_at == 1727258400 + # in_progress -> in_progress_at is set to created_at. + assert batch.in_progress_at == 1727172000 + assert batch.completed_at is None + assert batch.cancelling_at is None + assert batch.cancelled_at is None + # request_counts: total = processing+succeeded+errored+canceled+expired. + assert batch.request_counts.total == 6 + assert batch.request_counts.completed == 2 + assert batch.request_counts.failed == 1 + assert batch.metadata == {} + + +def test_transform_retrieve_response_ended_maps_to_completed(config): + raw = _response( + { + "id": "msgbatch_done", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "request_counts": {"succeeded": 5, "errored": 0}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # "ended" -> OpenAI "completed". + assert batch.status == "completed" + # completed_at populated only because processing_status == "ended". + assert batch.completed_at == 1727175600 + # not in_progress -> in_progress_at stays None. + assert batch.in_progress_at is None + assert batch.request_counts.total == 5 + assert batch.request_counts.completed == 5 + + +def test_transform_retrieve_response_canceling_maps_to_cancelling(config): + raw = _response( + { + "id": "msgbatch_cancel", + "processing_status": "canceling", + "created_at": "2024-09-24T10:00:00Z", + "cancel_initiated_at": "2024-09-24T10:30:00Z", + "ended_at": "2024-09-24T10:45:00Z", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # "canceling" -> OpenAI "cancelling". + assert batch.status == "cancelling" + assert batch.cancelling_at == 1727173800 + # cancelled_at = ended_at when canceling and ended_at present. + assert batch.cancelled_at == 1727174700 + assert batch.completed_at is None + + +def test_transform_retrieve_response_unknown_status_defaults_in_progress(config): + raw = _response( + { + "id": "msgbatch_x", + "processing_status": "some_future_status", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # Unmapped status falls back to in_progress (don't 500 on new enum values). + assert batch.status == "in_progress" + + +def test_transform_retrieve_response_missing_id_and_status_defaults(config): + # Empty body: id defaults to "", status defaults to "in_progress". + raw = _response({}) + before = int(time.time()) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + after = int(time.time()) + assert batch.id == "" + assert batch.status == "in_progress" + # No created_at -> created_at falls back to int(time.time()). + assert before <= batch.created_at <= after + # No created_at -> in_progress_at (which mirrors created_at) is None. + assert batch.in_progress_at is None + assert batch.request_counts.total == 0 + + +def test_transform_retrieve_response_archived_sets_expired_at(config): + raw = _response( + { + "id": "msgbatch_arch", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "archived_at": "2024-09-26T10:00:00Z", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # archived_at present -> expired_at populated. + assert batch.expired_at == 1727344800 + + +def test_transform_retrieve_response_bad_timestamp_is_none(config): + raw = _response( + { + "id": "msgbatch_bad", + "processing_status": "in_progress", + "created_at": "not-a-real-timestamp", + "request_counts": {}, + } + ) + before = int(time.time()) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + after = int(time.time()) + # Unparseable created_at -> parse_timestamp returns None, created_at falls + # back to time.time(). + assert before <= batch.created_at <= after + + +def test_transform_retrieve_response_unparseable_json_raises(config): + bad = httpx.Response( + status_code=200, + content=b"not json", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + with pytest.raises(ValueError, match="Failed to parse Anthropic batch response"): + config.transform_retrieve_batch_response( + model=None, raw_response=bad, logging_obj=MagicMock(), litellm_params={} + ) + + +# =========================================================================== # +# get_error_class +# =========================================================================== # + + +def test_get_error_class_with_dict_headers(config): + err = config.get_error_class( + error_message="rate limited", status_code=429, headers={"x-ratelimit": "0"} + ) + from litellm.llms.anthropic.common_utils import AnthropicError + + assert isinstance(err, AnthropicError) + assert err.status_code == 429 + assert err.message == "rate limited" + + +def test_get_error_class_with_httpx_headers(config): + hdrs = httpx.Headers({"retry-after": "5"}) + err = config.get_error_class( + error_message="server error", status_code=500, headers=hdrs + ) + assert err.status_code == 500 + assert err.message == "server error" + + +# =========================================================================== # +# transform_response (batch results JSONL -> summed usage on ModelResponse) +# =========================================================================== # + + +def test_transform_response_sums_usage_across_lines(config): + from litellm.types.utils import ModelResponse, Usage + + # Two result lines; transform_parsed_response is stubbed to attach a fixed + # Usage per line so we can assert the SUM is what lands on model_response. + line1 = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + line2 = '{"result": {"message": {"content": [{"type": "text", "text": "b"}]}}}' + raw = httpx.Response( + status_code=200, + text=f"{line1}\n{line2}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + model_response = ModelResponse() + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr( + mr, + "usage", + Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + return mr + + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ): + out = config.transform_response( + model="claude-3", + raw_response=raw, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert out is model_response + usage = getattr(out, "usage") + # Two lines * (10 prompt, 5 completion) summed. + assert usage.prompt_tokens == 20 + assert usage.completion_tokens == 10 + assert usage.total_tokens == 30 + + +def test_transform_response_skips_malformed_lines(config): + from litellm.types.utils import ModelResponse, Usage + + valid = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + # Interior blank line (survives the outer strip) exercises the empty-line + # `continue`; leading not-json exercises the JSONDecodeError `continue`. + raw = httpx.Response( + status_code=200, + text=f"not-json\n\n{valid}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + model_response = ModelResponse() + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr( + mr, "usage", Usage(prompt_tokens=7, completion_tokens=3, total_tokens=10) + ) + return mr + + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ): + out = config.transform_response( + model="claude-3", + raw_response=raw, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + # Only the single valid line contributed usage; malformed/empty skipped. + usage = getattr(out, "usage") + assert usage.prompt_tokens == 7 + assert usage.completion_tokens == 3 + + +def test_transform_response_reraises_unexpected_error(config): + from litellm.types.utils import ModelResponse, Usage + + valid = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + raw = httpx.Response( + status_code=200, + text=f"{valid}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr(mr, "usage", Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)) + return mr + + # A non-JSONDecodeError raised during usage aggregation must propagate + # (the outer `except Exception: raise e`), not be swallowed. + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ), patch( + "litellm.cost_calculator.BaseTokenUsageProcessor.combine_usage_objects", + side_effect=RuntimeError("boom"), + ): + with pytest.raises(RuntimeError, match="boom"): + config.transform_response( + model="claude-3", + raw_response=raw, + model_response=ModelResponse(), + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +# --------------------------------------------------------------------------- # +# Shared BaseBatchesConfig contract suite (consistency net across providers). +# This subclass supplies anthropic fixtures; the inherited contract tests run +# automatically. See base_batches_config_test.py. +# --------------------------------------------------------------------------- # + +from litellm.types.utils import LlmProviders # noqa: E402 +from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 + BatchesConfigContractTests, +) + + +class TestAnthropicBatchesContract(BatchesConfigContractTests): + def make_config(self): + from litellm.llms.anthropic.batches.transformation import ( + AnthropicBatchesConfig, + ) + + return AnthropicBatchesConfig() + + expected_provider = LlmProviders.ANTHROPIC + supports_create = False # anthropic raises NotImplementedError on create + supports_retrieve_response = True + + def sample_retrieve_response_body(self) -> dict: + return { + "id": "msgbatch_123", + "processing_status": "ended", + "created_at": "2024-01-01T00:00:00Z", + "ended_at": "2024-01-02T00:00:00Z", + "request_counts": {"succeeded": 2, "errored": 1}, + } + + expected_retrieve_batch_id = "msgbatch_123" + expected_retrieve_status = "completed" # "ended" -> "completed" diff --git a/tests/test_litellm/llms/azure/batches/__init__.py b/tests/test_litellm/llms/azure/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/azure/batches/test_handler.py b/tests/test_litellm/llms/azure/batches/test_handler.py new file mode 100644 index 00000000000..f2332a7de7c --- /dev/null +++ b/tests/test_litellm/llms/azure/batches/test_handler.py @@ -0,0 +1,491 @@ +"""Unit tests for ``AzureBatchesAPI`` (litellm/llms/azure/batches/handler.py). + +The Azure batches handler is HTTP/auth glue: each public method +(create/retrieve/cancel/list) resolves an Azure OpenAI client via the inherited +``get_azure_openai_client`` seam, branches on ``_is_async`` (returning the +``a*`` coroutine in the async case, calling the sync client otherwise), validates +the client type, and parses the SDK response into ``LiteLLMBatch``. + +We mock only true boundaries: + * ``get_azure_openai_client`` - the credential/client-construction seam. We + assert the EXACT auth args (api_key / api_base / api_version / client / + _is_async / litellm_params) forwarded to it. + * the returned Azure OpenAI client's ``batches.*`` methods - the network call. + We assert the request data forwarded and that the SDK response is parsed + into ``LiteLLMBatch`` (sibling SDK methods asserted NOT called). + +Pure logic (the _is_async branch, the isinstance guards, the model_dump parse) +runs for real. +""" + +from __future__ import annotations + +import asyncio +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from openai import AsyncOpenAI, OpenAI # noqa: E402 + +from litellm.llms.azure.azure import AsyncAzureOpenAI, AzureOpenAI # noqa: E402 +from litellm.llms.azure.batches.handler import AzureBatchesAPI # noqa: E402 +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +GET_CLIENT = "litellm.llms.azure.batches.handler.AzureBatchesAPI.get_azure_openai_client" + +AUTH_KW = dict( + api_key="sk-azure-test", + api_base="https://my-azure.openai.azure.com", + api_version="2024-12-01", + timeout=600.0, + max_retries=3, +) + +CREATE_DATA = { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", +} +RETRIEVE_DATA = {"batch_id": "batch-123"} +CANCEL_DATA = {"batch_id": "batch-123"} + + +def _batch_dict(batch_id: str = "batch-123", status: str = "completed") -> dict: + """A minimal-but-valid dict for ``LiteLLMBatch(**response.model_dump())``.""" + return { + "id": batch_id, + "completion_window": "24h", + "created_at": 1700000000, + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", + "object": "batch", + "status": status, + "output_file_id": "file-out-xyz", + } + + +def _sdk_response(batch_dict: dict) -> MagicMock: + """An object that mimics the OpenAI SDK Batch: only ``.model_dump()`` is used.""" + resp = MagicMock() + resp.model_dump.return_value = batch_dict + return resp + + +def _sync_client() -> MagicMock: + """A sync Azure client (passes ``isinstance(.., AzureOpenAI)``).""" + return MagicMock(spec=AzureOpenAI) + + +def _async_client() -> MagicMock: + """An async Azure client (passes ``isinstance(.., AsyncAzureOpenAI)``). + + The ``batches.*`` SDK methods are awaited by the handler, so they must be + AsyncMocks. + """ + client = MagicMock(spec=AsyncAzureOpenAI) + client.batches.create = AsyncMock() + client.batches.retrieve = AsyncMock() + client.batches.cancel = AsyncMock() + client.batches.list = AsyncMock() + return client + + +@pytest.fixture +def handler() -> AzureBatchesAPI: + return AzureBatchesAPI() + + +# =========================================================================== # +# create_batch - sync path +# =========================================================================== # + + +def test_create_sync_forwards_auth_to_client_seam(handler): + client = _sync_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + # EXACT auth args forwarded to the client-construction seam. + assert get_client.call_count == 1 + kw = get_client.call_args.kwargs + assert kw["api_key"] == "sk-azure-test" + assert kw["api_base"] == "https://my-azure.openai.azure.com" + assert kw["api_version"] == "2024-12-01" + assert kw["_is_async"] is False + assert kw["client"] is None + # litellm_params defaults to {} (not None) when not supplied. + assert kw["litellm_params"] == {} + + # PAYLOAD: request data forwarded verbatim to the SDK as kwargs. + client.batches.create.assert_called_once_with(**CREATE_DATA) + # sibling SDK seams untouched. + client.batches.retrieve.assert_not_called() + client.batches.cancel.assert_not_called() + + # RESULT: parsed into LiteLLMBatch from the SDK response's model_dump. + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + assert result.status == "completed" + assert result.output_file_id == "file-out-xyz" + + +def test_create_sync_passes_litellm_params_through(handler): + client = _sync_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + lp = {"azure_ad_token": "tok", "tenant_id": "t1"} + + with patch(GET_CLIENT, return_value=client) as get_client: + handler.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + litellm_params=lp, + **AUTH_KW, + ) + + assert get_client.call_args.kwargs["litellm_params"] == lp + + +def test_create_sync_explicit_client_forwarded_to_seam(handler): + sentinel_client = _sync_client() + sentinel_client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=sentinel_client) as get_client: + handler.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + client=sentinel_client, + **AUTH_KW, + ) + + assert get_client.call_args.kwargs["client"] is sentinel_client + + +def test_create_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + +# =========================================================================== # +# create_batch - async path +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create_async_returns_coroutine_and_awaits_async_client(handler): + client = _async_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.create.assert_awaited_once_with(**CREATE_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + + +@pytest.mark.asyncio +async def test_create_async_rejects_sync_client(handler): + """_is_async=True but seam returns a sync client -> ValueError, no network.""" + sync_client = _sync_client() + + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + sync_client.batches.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_acreate_batch_parses_response(handler): + client = _async_client() + client.batches.create.return_value = _sdk_response(_batch_dict(status="validating")) + + result = await handler.acreate_batch( + create_batch_data=CREATE_DATA, azure_client=client + ) + + client.batches.create.assert_awaited_once_with(**CREATE_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "validating" + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve_sync_dispatch_payload_and_result(handler): + client = _sync_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.retrieve_batch( + _is_async=False, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.retrieve.assert_called_once_with(**RETRIEVE_DATA) + client.batches.create.assert_not_called() + client.batches.cancel.assert_not_called() + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + + +def test_retrieve_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.retrieve_batch( + _is_async=False, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + + +@pytest.mark.asyncio +async def test_retrieve_async_returns_coroutine_and_awaits(handler): + client = _async_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.retrieve_batch( + _is_async=True, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.retrieve.assert_awaited_once_with(**RETRIEVE_DATA) + assert isinstance(result, LiteLLMBatch) + + +@pytest.mark.asyncio +async def test_retrieve_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.retrieve_batch( + _is_async=True, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + sync_client.batches.retrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_aretrieve_batch_parses_response(handler): + client = _async_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + result = await handler.aretrieve_batch( + retrieve_batch_data=RETRIEVE_DATA, client=client + ) + + client.batches.retrieve.assert_awaited_once_with(**RETRIEVE_DATA) + assert isinstance(result, LiteLLMBatch) + + +# =========================================================================== # +# cancel_batch (has an EXTRA sync-side isinstance guard the others lack) +# =========================================================================== # + + +def test_cancel_sync_dispatch_payload_and_result(handler): + client = _sync_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.cancel.assert_called_once_with(**CANCEL_DATA) + client.batches.create.assert_not_called() + client.batches.retrieve.assert_not_called() + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +def test_cancel_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + +def test_cancel_sync_rejects_non_sync_client(handler): + """cancel_batch has a unique sync-side guard: if _is_async is False but the + resolved client is async (neither AzureOpenAI nor OpenAI), it must raise + rather than call .cancel().""" + async_client = _async_client() + + with patch(GET_CLIENT, return_value=async_client): + with pytest.raises(ValueError, match="sync client"): + handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + async_client.batches.cancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_cancel_async_returns_coroutine_and_awaits(handler): + client = _async_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.cancel_batch( + _is_async=True, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.cancel.assert_awaited_once_with(**CANCEL_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +@pytest.mark.asyncio +async def test_cancel_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="async client"): + handler.cancel_batch( + _is_async=True, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + sync_client.batches.cancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_acancel_batch_parses_response(handler): + client = _async_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + result = await handler.acancel_batch(cancel_batch_data=CANCEL_DATA, client=client) + + client.batches.cancel.assert_awaited_once_with(**CANCEL_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +# =========================================================================== # +# list_batches (returns the raw SDK response, NOT a LiteLLMBatch) +# =========================================================================== # + + +def test_list_sync_forwards_after_limit_and_returns_raw_response(handler): + client = _sync_client() + raw = MagicMock(name="raw_list_response") + client.batches.list.return_value = raw + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.list_batches( + _is_async=False, after="cur-1", limit=20, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.list.assert_called_once_with(after="cur-1", limit=20) + # list returns the SDK response untouched (no LiteLLMBatch parsing). + assert result is raw + + +def test_list_sync_defaults_after_and_limit_to_none(handler): + client = _sync_client() + client.batches.list.return_value = MagicMock() + + with patch(GET_CLIENT, return_value=client): + handler.list_batches(_is_async=False, **AUTH_KW) + + client.batches.list.assert_called_once_with(after=None, limit=None) + + +def test_list_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.list_batches(_is_async=False, **AUTH_KW) + + +@pytest.mark.asyncio +async def test_list_async_returns_coroutine_and_awaits(handler): + client = _async_client() + raw = MagicMock(name="raw_async_list_response") + client.batches.list.return_value = raw + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.list_batches( + _is_async=True, after="cur-2", limit=7, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.list.assert_awaited_once_with(after="cur-2", limit=7) + assert result is raw + + +@pytest.mark.asyncio +async def test_list_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.list_batches(_is_async=True, **AUTH_KW) + sync_client.batches.list.assert_not_called() + + +@pytest.mark.asyncio +async def test_alist_batches_returns_raw_response(handler): + client = _async_client() + raw = MagicMock(name="raw") + client.batches.list.return_value = raw + + result = await handler.alist_batches(client=client, after="a", limit=2) + + client.batches.list.assert_awaited_once_with(after="a", limit=2) + assert result is raw + + +# =========================================================================== # +# Cross-cutting: an OpenAI (non-Azure) client also satisfies the type guards, +# since the Union allows OpenAI / AsyncOpenAI (Azure-v1 path returns these). +# =========================================================================== # + + +def test_create_sync_accepts_plain_openai_client(handler): + client = MagicMock(spec=OpenAI) + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client): + result = handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + assert isinstance(result, LiteLLMBatch) + + +@pytest.mark.asyncio +async def test_create_async_accepts_plain_async_openai_client(handler): + client = MagicMock(spec=AsyncOpenAI) + client.batches.create = AsyncMock(return_value=_sdk_response(_batch_dict())) + + with patch(GET_CLIENT, return_value=client): + result = await handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + assert isinstance(result, LiteLLMBatch) diff --git a/tests/test_litellm/llms/base_llm/__init__.py b/tests/test_litellm/llms/base_llm/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/base_llm/batches/__init__.py b/tests/test_litellm/llms/base_llm/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py b/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py new file mode 100644 index 00000000000..fd526c55de4 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py @@ -0,0 +1,128 @@ +""" +Reusable contract test suite for BaseBatchesConfig implementations. + +Any provider whose batch transformation subclasses +`litellm.llms.base_llm.batches.transformation.BaseBatchesConfig` gets a shared +consistency net by subclassing `BatchesConfigContractTests` in its own +`test_transformation.py` (as a `Test*`-named class) and overriding the hooks +below. pytest then runs every contract test against that provider, guaranteeing +all provider batch transformations honour the same BaseBatchesConfig contract - +e.g. every `transform_retrieve_batch_response` returns a real `LiteLLMBatch` +with `object == "batch"`, a valid status, and an int `created_at`. + +This module is intentionally NOT named `test_*`: it holds no standalone tests +and must not be collected on its own. It mirrors the established repo pattern in +`tests/llm_translation/base_*_unit_tests.py`. + +Providers that do NOT implement BaseBatchesConfig (e.g. vertex_ai, whose +transformation is a standalone class with a different shape) cannot use this and +keep fully standalone tests. +""" + +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.types.utils import LiteLLMBatch, LlmProviders + +# The OpenAI BatchJobStatus literal set - every provider must map into this. +VALID_BATCH_STATUSES = { + "validating", + "failed", + "in_progress", + "finalizing", + "completed", + "expired", + "cancelling", + "cancelled", +} + + +def make_raw_response(body: dict, status_code: int = 200) -> httpx.Response: + """Build an httpx.Response whose .json() yields `body` - the input shape the + transform_*_response methods consume.""" + return httpx.Response(status_code=status_code, json=body) + + +class BatchesConfigContractTests: + """Contract every BaseBatchesConfig implementation must satisfy. + + Subclass this with a `Test`-prefixed class and override the hooks. Do NOT + add a `Test` prefix here - this base must not be collected directly. + """ + + # ----------------------------------------------------------------------- # + # Hooks - providers MUST override these. + # ----------------------------------------------------------------------- # + + def make_config(self): + """Return a fresh instance of the provider's BaseBatchesConfig.""" + raise NotImplementedError("override make_config()") + + # The LlmProviders value this config reports. + expected_provider: LlmProviders = None # type: ignore[assignment] + + # Does this provider implement batch CREATE via the transformation? + # (anthropic raises NotImplementedError; bedrock/others may support it.) + supports_create: bool = False + + # Does this provider parse retrieve responses in the transformation layer? + supports_retrieve_response: bool = True + + def sample_retrieve_response_body(self) -> dict: + """A representative raw provider retrieve-batch response body.""" + raise NotImplementedError("override sample_retrieve_response_body()") + + # Expected mapped values for the sample above. + expected_retrieve_batch_id: str = None # type: ignore[assignment] + expected_retrieve_status: str = None # type: ignore[assignment] + + # ----------------------------------------------------------------------- # + # Contract tests - run for every provider subclass. + # ----------------------------------------------------------------------- # + + def test_contract__custom_llm_provider(self): + assert self.make_config().custom_llm_provider == self.expected_provider + + def test_contract__get_error_class_is_exception_with_status(self): + err = self.make_config().get_error_class( + error_message="boom", status_code=429, headers={} + ) + assert isinstance(err, Exception) + assert getattr(err, "status_code", None) == 429 + + def test_contract__create_unsupported_raises(self): + if self.supports_create: + pytest.skip("provider supports batch create; see provider-specific tests") + with pytest.raises(NotImplementedError): + self.make_config().transform_create_batch_request( + model="m", + create_batch_data={ + "input_file_id": "f", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + optional_params={}, + litellm_params={}, + ) + + def test_contract__retrieve_response_is_valid_litellm_batch(self): + if not self.supports_retrieve_response: + pytest.skip("provider handles retrieve outside the transformation layer") + out = self.make_config().transform_retrieve_batch_response( + model=None, + raw_response=make_raw_response(self.sample_retrieve_response_body()), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert isinstance(out, LiteLLMBatch) + assert out.object == "batch" + assert out.status in VALID_BATCH_STATUSES + assert isinstance(out.created_at, int) + assert out.id == self.expected_retrieve_batch_id + assert out.status == self.expected_retrieve_status diff --git a/tests/test_litellm/llms/base_llm/batches/test_transformation.py b/tests/test_litellm/llms/base_llm/batches/test_transformation.py new file mode 100644 index 00000000000..cfb9f278f80 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/batches/test_transformation.py @@ -0,0 +1,231 @@ +""" +Unit tests for litellm/llms/base_llm/batches/transformation.py + +BaseBatchesConfig is the abstract base class that every provider-specific +batches config subclasses. It is almost entirely interface (abstractmethods + +one abstract property), so the only concrete behavior to regression-lock is: + + - the abstractness contract: the base class cannot be instantiated, and a + subclass missing any abstract member also cannot be instantiated; a + subclass implementing all of them can. + - get_config(): a classmethod that reflects over ``cls.__dict__`` and returns + the class-level config attributes, filtering out dunders, ``_abc`` internals, + callables (function/builtin/classmethod/staticmethod), and ``None`` values. + +These tests assert the exact dict get_config() produces for hand-built +subclasses, so a change to the filter predicate (e.g. dropping the ``None`` +filter, dropping the staticmethod/classmethod filter, or widening the prefix +filter to all single-underscore names) makes a test fail. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig +from litellm.types.utils import LlmProviders + + +# --------------------------------------------------------------------------- # +# A fully-concrete subclass: implements every abstract member with trivial +# bodies so it can be instantiated and so get_config() has a real cls to +# reflect over. Class-level attributes here are the get_config() fixtures. +# --------------------------------------------------------------------------- # + + +class _ConcreteBatchesConfig(BaseBatchesConfig): + string_attr = "hello" + int_attr = 42 + list_attr = [1, 2, 3] + none_attr = None + _single_underscore = "kept" + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.OPENAI + + def validate_environment( + self, + headers, + model, + messages, + optional_params, + litellm_params, + api_key=None, + api_base=None, + ) -> dict: + return headers + + def get_complete_batch_url( + self, api_base, api_key, model, optional_params, litellm_params, data + ) -> str: + return "https://example.com/batch" + + def transform_create_batch_request( + self, model, create_batch_data, optional_params, litellm_params + ): + return {"created": True} + + def transform_create_batch_response( + self, model, raw_response, logging_obj, litellm_params + ): + return raw_response + + def transform_retrieve_batch_request( + self, batch_id, optional_params, litellm_params + ): + return {"batch_id": batch_id} + + def transform_retrieve_batch_response( + self, model, raw_response, logging_obj, litellm_params + ): + return raw_response + + def get_error_class(self, error_message, status_code, headers): + return Exception(error_message) + + +# =========================================================================== # +# Abstractness contract +# =========================================================================== # + + +def test_base_class_cannot_be_instantiated(): + """The base class has unimplemented abstractmethods, so direct + instantiation must raise TypeError.""" + with pytest.raises(TypeError): + BaseBatchesConfig() + + +def test_fully_concrete_subclass_can_be_instantiated(): + instance = _ConcreteBatchesConfig() + assert isinstance(instance, BaseBatchesConfig) + + +@pytest.mark.parametrize( + "missing_member", + [ + "custom_llm_provider", + "validate_environment", + "get_complete_batch_url", + "transform_create_batch_request", + "transform_create_batch_response", + "transform_retrieve_batch_request", + "transform_retrieve_batch_response", + "get_error_class", + ], +) +def test_subclass_missing_any_abstract_member_cannot_instantiate(missing_member): + """Every abstract member is part of the contract: dropping any one of them + leaves the subclass abstract and uninstantiable.""" + namespace = { + k: v + for k, v in _ConcreteBatchesConfig.__dict__.items() + if not k.startswith("__") + } + namespace.pop(missing_member) + Incomplete = type("Incomplete", (BaseBatchesConfig,), namespace) + with pytest.raises(TypeError): + Incomplete() + + +def test_concrete_instance_methods_run(): + """Sanity: the trivial overrides actually execute through the base contract.""" + instance = _ConcreteBatchesConfig() + assert instance.custom_llm_provider == LlmProviders.OPENAI + assert instance.validate_environment( + headers={"x": "1"}, + model="m", + messages=[], + optional_params={}, + litellm_params={}, + ) == {"x": "1"} + assert instance.transform_retrieve_batch_request( + batch_id="b-1", optional_params={}, litellm_params={} + ) == {"batch_id": "b-1"} + + +# =========================================================================== # +# get_config() +# =========================================================================== # + + +def test_get_config_returns_class_level_non_none_data_attrs(): + """Exact contents: only class-level data attributes that are not None, + not dunders, not callables. Single-underscore names ARE kept (only ``__`` + and ``_abc`` prefixes are filtered). The ``custom_llm_provider`` property + object also survives the filter (a property is neither a function nor None), + matching how real provider subclasses define it.""" + config = _ConcreteBatchesConfig.get_config() + custom_llm_provider = config.pop("custom_llm_provider") + assert isinstance(custom_llm_provider, property) + assert config == { + "string_attr": "hello", + "int_attr": 42, + "list_attr": [1, 2, 3], + "_single_underscore": "kept", + } + + +def test_get_config_excludes_none_valued_attrs(): + assert "none_attr" not in _ConcreteBatchesConfig.get_config() + + +def test_get_config_excludes_methods_and_property(): + config = _ConcreteBatchesConfig.get_config() + for method_name in ( + "validate_environment", + "get_complete_batch_url", + "transform_create_batch_request", + "transform_create_batch_response", + "transform_retrieve_batch_request", + "transform_retrieve_batch_response", + "get_error_class", + "get_config", + ): + assert method_name not in config + + +def test_get_config_excludes_classmethod_and_staticmethod(): + """classmethod and staticmethod objects are filtered even though they are + not plain FunctionType.""" + + class WithCallables(_ConcreteBatchesConfig): + keep_me = "yes" + + @staticmethod + def a_static(): + return 1 + + @classmethod + def a_class(cls): + return 2 + + config = WithCallables.get_config() + assert config == {"keep_me": "yes"} + + +def test_get_config_only_reflects_own_dict_not_inherited(): + """get_config reflects cls.__dict__ only, so attributes defined on a parent + do not leak into a child's config.""" + + class Parent(_ConcreteBatchesConfig): + parent_attr = "parent" + + class Child(Parent): + child_attr = "child" + + assert Parent.get_config() == {"parent_attr": "parent"} + assert Child.get_config() == {"child_attr": "child"} + + +def test_get_config_on_base_class_exposes_only_the_abstract_property(): + """On the base class itself, the only ``__dict__`` member that survives the + filter is the ``custom_llm_provider`` property object (a property is neither + a function nor None and its name has no filtered prefix).""" + config = BaseBatchesConfig.get_config() + assert list(config.keys()) == ["custom_llm_provider"] + assert isinstance(config["custom_llm_provider"], property) diff --git a/tests/test_litellm/llms/bedrock/__init__.py b/tests/test_litellm/llms/bedrock/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/bedrock/batches/__init__.py b/tests/test_litellm/llms/bedrock/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py new file mode 100644 index 00000000000..d1ad5943ae6 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py @@ -0,0 +1,684 @@ +""" +Regression tests for ``BedrockBatchesConfig`` (the BaseBatchesConfig +implementation for Bedrock model-invocation-job batches). + +This file complements (does not duplicate): + - ``test_batch_metadata_sanitization.py`` (covers + ``_get_openai_compatible_batch_metadata`` exhaustively) + - ``test_handler.py`` (covers the boto3-backed handler, not this transform) + +Here we lock the pure transform logic in ``transformation.py``: request +construction (S3 input/output config, model id, job name, role ARN), the +AWS-JobStatus -> OpenAI-status mapping, timestamp parsing, retrieve-request +URL/ARN handling, and the error class. AWS auth/sigv4 is the only external seam +we mock; everything else runs for real. +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig +from litellm.types.utils import LiteLLMBatch, LlmProviders + +# AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py +# (both transform_create_batch_response and transform_retrieve_batch_response). +STATUS_MAP = { + "Submitted": "validating", + "Validating": "validating", + "Scheduled": "in_progress", + "InProgress": "in_progress", + "PartiallyCompleted": "completed", + "Completed": "completed", + "Failed": "failed", + "Stopping": "cancelling", + "Stopped": "cancelled", + "Expired": "expired", +} + +ARN = "arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/abc1234567" + + +@pytest.fixture +def config(): + return BedrockBatchesConfig() + + +def _raw(body: dict, status_code: int = 200) -> httpx.Response: + return httpx.Response(status_code=status_code, json=body) + + +# --------------------------------------------------------------------------- # +# get_complete_batch_url +# --------------------------------------------------------------------------- # + + +def test_get_complete_batch_url_uses_region(config): + url = config.get_complete_batch_url( + api_base=None, + api_key=None, + model="anthropic.claude-3", + optional_params={"aws_region_name": "eu-central-1"}, + litellm_params={}, + data={"input_file_id": "s3://b/k"}, + ) + assert url == "https://bedrock.eu-central-1.amazonaws.com/model-invocation-job" + + +# --------------------------------------------------------------------------- # +# transform_create_batch_request - request construction (sign_aws_request mocked) +# --------------------------------------------------------------------------- # + + +def test_create_request_builds_s3_input_output_and_arn(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-deadbeef", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({"Authorization": "signed"}, b'{"x": 1}') + result = config.transform_create_batch_request( + model="anthropic.claude-3-5-sonnet", + create_batch_data={ + "input_file_id": "s3://in-bucket/path/to/input.jsonl", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + optional_params={"aws_region_name": "us-west-2"}, + litellm_params={ + "s3_output_bucket_name": "out-bucket", + "aws_batch_role_arn": "arn:aws:iam::123:role/my-batch-role", + }, + ) + + bedrock_request = mock_sign.call_args.kwargs["data"] + assert bedrock_request["modelId"] == "anthropic.claude-3-5-sonnet" + assert bedrock_request["jobName"] == "litellm-batch-deadbeef" + assert bedrock_request["roleArn"] == "arn:aws:iam::123:role/my-batch-role" + assert ( + bedrock_request["inputDataConfig"]["s3InputDataConfig"]["s3Uri"] + == "s3://in-bucket/path/to/input.jsonl" + ) + assert ( + bedrock_request["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"] + == "s3://out-bucket/litellm-batch-outputs/litellm-batch-deadbeef/" + ) + # 24h completion window -> 24 hour timeout + assert bedrock_request["timeoutDurationInHours"] == 24 + # signing was over the bedrock endpoint via POST + assert mock_sign.call_args.kwargs["service_name"] == "bedrock" + assert mock_sign.call_args.kwargs["method"] == "POST" + assert mock_sign.call_args.kwargs["endpoint_url"] == ( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job" + ) + # the transform returns the pre-signed envelope + assert result["method"] == "POST" + assert result["url"] == ( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job" + ) + assert result["headers"] == {"Authorization": "signed"} + + +def test_create_request_defaults_output_bucket_to_input_bucket(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-cafef00d", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://same-bucket/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + bedrock_request = mock_sign.call_args.kwargs["data"] + assert ( + bedrock_request["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"] + == "s3://same-bucket/litellm-batch-outputs/litellm-batch-cafef00d/" + ) + + +def test_create_request_adds_kms_encryption_key_when_provided(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={ + "aws_batch_role_arn": "arn:aws:iam::1:role/r", + "s3_encryption_key_id": "kms-key-123", + }, + ) + s3out = mock_sign.call_args.kwargs["data"]["outputDataConfig"][ + "s3OutputDataConfig" + ] + assert s3out["s3EncryptionKeyId"] == "kms-key-123" + + +def test_create_request_omits_kms_key_when_absent(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign, patch( + "litellm.llms.bedrock.batches.transformation.get_secret_str", + return_value=None, + ): + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + s3out = mock_sign.call_args.kwargs["data"]["outputDataConfig"][ + "s3OutputDataConfig" + ] + assert "s3EncryptionKeyId" not in s3out + + +def test_create_request_missing_input_file_id_raises(config): + with pytest.raises(ValueError, match="input_file_id is required"): + config.transform_create_batch_request( + model="m", + create_batch_data={}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + + +def test_create_request_missing_role_arn_raises(config, monkeypatch): + monkeypatch.delenv("AWS_BATCH_ROLE_ARN", raising=False) + with pytest.raises(ValueError, match="IAM role ARN is required"): + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={}, + ) + + +def test_create_request_role_arn_from_env(config, monkeypatch): + monkeypatch.setenv("AWS_BATCH_ROLE_ARN", "arn:aws:iam::9:role/env-role") + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={}, + ) + assert ( + mock_sign.call_args.kwargs["data"]["roleArn"] + == "arn:aws:iam::9:role/env-role" + ) + + +def test_create_request_missing_model_raises(config): + with pytest.raises(ValueError, match="Could not determine Bedrock model ID"): + config.transform_create_batch_request( + model="", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + + +def test_create_request_no_timeout_for_non_24h_window(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={ + "input_file_id": "s3://b/in.jsonl", + "completion_window": "48h", + }, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + assert "timeoutDurationInHours" not in mock_sign.call_args.kwargs["data"] + + +# --------------------------------------------------------------------------- # +# transform_create_batch_response - status mapping + LiteLLMBatch shape +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize("bedrock_status,openai_status", list(STATUS_MAP.items())) +def test_create_response_status_mapping(config, bedrock_status, openai_status): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": bedrock_status}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == openai_status + assert out.id == ARN + assert out.object == "batch" + + +def test_create_response_unknown_status_falls_back_to_validating(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "SomeFutureStatus"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_create_response_default_status_when_missing(config): + # status defaults to "Submitted" -> "validating" + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_create_response_in_progress_sets_in_progress_at(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "InProgress"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "in_progress" + assert isinstance(out.in_progress_at, int) + + +def test_create_response_non_in_progress_leaves_in_progress_at_none(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.in_progress_at is None + + +def test_create_response_uses_original_request_fields(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={ + "original_batch_request": { + "endpoint": "/v1/embeddings", + "input_file_id": "s3://b/in.jsonl", + "completion_window": "24h", + "metadata": {"user": "alice"}, + } + }, + ) + assert out.endpoint == "/v1/embeddings" + assert out.input_file_id == "s3://b/in.jsonl" + assert out.metadata == {"user": "alice"} + + +def test_create_response_default_endpoint_when_no_original_request(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.endpoint == "/v1/chat/completions" + assert out.completion_window == "24h" + + +def test_create_response_raises_on_unparseable_body(config): + bad = httpx.Response(status_code=200, text="not-json") + with pytest.raises(ValueError, match="Failed to parse Bedrock batch response"): + config.transform_create_batch_response( + model=None, + raw_response=bad, + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# --------------------------------------------------------------------------- # +# transform_retrieve_batch_request - ARN validation + URL construction +# --------------------------------------------------------------------------- # + + +def test_retrieve_request_builds_encoded_arn_url(config): + with patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({"Authorization": "signed"}, b"") + result = config.transform_retrieve_batch_request( + batch_id=ARN, optional_params={}, litellm_params={} + ) + # ARN is URL-encoded (colons and slashes escaped) into the path + assert result["method"] == "GET" + assert result["data"] is None + assert result["headers"] == {"Authorization": "signed"} + assert result["url"].startswith( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job/" + ) + assert "%3A" in result["url"] # colon encoded + assert "%2F" in result["url"] # slash encoded + assert mock_sign.call_args.kwargs["method"] == "GET" + assert mock_sign.call_args.kwargs["data"] == {} + + +def test_retrieve_request_rejects_non_arn(config): + with pytest.raises(ValueError, match="Expected ARN"): + config.transform_retrieve_batch_request( + batch_id="abc1234567", optional_params={}, litellm_params={} + ) + + +def test_retrieve_request_rejects_short_arn(config): + with pytest.raises(ValueError, match="Invalid ARN format"): + config.transform_retrieve_batch_request( + batch_id="arn:aws:bedrock:us-west-2", optional_params={}, litellm_params={} + ) + + +def test_retrieve_request_rejects_bad_region(config): + bad = "arn:aws:bedrock:US_WEST:123:model-invocation-job/x" + with pytest.raises(ValueError, match="Invalid region in ARN"): + config.transform_retrieve_batch_request( + batch_id=bad, optional_params={}, litellm_params={} + ) + + +# --------------------------------------------------------------------------- # +# transform_retrieve_batch_response - status, timestamps, files, errors +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize("bedrock_status,openai_status", list(STATUS_MAP.items())) +def test_retrieve_response_status_mapping(config, bedrock_status, openai_status): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": bedrock_status}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == openai_status + + +def test_retrieve_response_unknown_status_falls_back_to_validating(config): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "NewStatus"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_retrieve_response_extracts_file_configs(config): + body = { + "jobArn": ARN, + "status": "Completed", + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": "s3://b/in.jsonl"}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": "s3://b/out/"}}, + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.input_file_id == "s3://b/in.jsonl" + assert out.output_file_id == "s3://b/out/" + + +def test_retrieve_response_parses_timestamps_for_completed(config): + body = { + "jobArn": ARN, + "status": "Completed", + "submitTime": "2026-04-28T12:00:00Z", + "endTime": "2026-04-28T12:30:00Z", + "jobExpirationTime": "2026-05-28T12:00:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + expect_created = int( + datetime.datetime.fromisoformat("2026-04-28T12:00:00+00:00").timestamp() + ) + expect_completed = int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + expect_expires = int( + datetime.datetime.fromisoformat("2026-05-28T12:00:00+00:00").timestamp() + ) + assert out.created_at == expect_created + assert out.completed_at == expect_completed + assert out.expires_at == expect_expires + # completed -> not failed/cancelled, no in_progress timestamp + assert out.failed_at is None + assert out.cancelled_at is None + assert out.in_progress_at is None + + +def test_retrieve_response_failed_sets_failed_at_from_end_time(config): + body = { + "jobArn": ARN, + "status": "Failed", + "endTime": "2026-04-28T12:30:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.failed_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + assert out.completed_at is None + assert out.cancelled_at is None + + +def test_retrieve_response_stopped_sets_cancelled_at(config): + body = {"jobArn": ARN, "status": "Stopped", "endTime": "2026-04-28T12:30:00Z"} + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.cancelled_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + assert out.completed_at is None + assert out.failed_at is None + + +def test_retrieve_response_in_progress_sets_in_progress_at_from_last_modified(config): + body = { + "jobArn": ARN, + "status": "InProgress", + "lastModifiedTime": "2026-04-28T12:15:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.in_progress_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:15:00+00:00").timestamp() + ) + + +def test_retrieve_response_invalid_timestamp_becomes_none(config): + body = { + "jobArn": ARN, + "status": "Completed", + "submitTime": "not-a-timestamp", + "endTime": "also-bad", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + # created_at falls back to int(time.time()) when submitTime unparseable + assert isinstance(out.created_at, int) + assert out.completed_at is None + + +def test_retrieve_response_builds_errors_from_message(config): + body = {"jobArn": ARN, "status": "Failed", "message": "validation failed"} + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body, status_code=400), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.errors is not None + assert out.errors.data[0].message == "validation failed" + assert out.errors.data[0].code == "400" + + +def test_retrieve_response_no_errors_when_no_message(config): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Completed"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.errors is None + + +def test_retrieve_response_enriches_metadata(config): + body = { + "jobArn": ARN, + "status": "Completed", + "jobName": "litellm-batch-1", + "modelId": "anthropic.claude-3", + "roleArn": "arn:aws:iam::1:role/r", + "timeoutDurationInHours": 24, + "vpcConfig": {"subnetIds": ["subnet-1"]}, + "clientRequestToken": None, + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.metadata["jobName"] == "litellm-batch-1" + assert out.metadata["modelId"] == "anthropic.claude-3" + assert out.metadata["roleArn"] == "arn:aws:iam::1:role/r" + # non-string scalar is stringified + assert out.metadata["timeoutDurationInHours"] == "24" + # dict/list serialized to JSON string + assert out.metadata["vpcConfig"] == '{"subnetIds": ["subnet-1"]}' + # None-valued fields dropped + assert "clientRequestToken" not in out.metadata + + +def test_retrieve_response_raises_on_unparseable_body(config): + bad = httpx.Response(status_code=200, text="<<>>") + with pytest.raises(ValueError, match="Failed to parse Bedrock batch response"): + config.transform_retrieve_batch_response( + model=None, + raw_response=bad, + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# --------------------------------------------------------------------------- # +# get_error_class + custom_llm_provider +# --------------------------------------------------------------------------- # + + +def test_get_error_class_returns_bedrock_error(config): + err = config.get_error_class( + error_message="throttled", status_code=429, headers={} + ) + assert isinstance(err, Exception) + assert err.status_code == 429 + assert "throttled" in str(err) + + +def test_custom_llm_provider_is_bedrock(config): + assert config.custom_llm_provider == LlmProviders.BEDROCK + + +def test_validate_environment_passes_headers_through(config): + headers = {"X-Custom": "v"} + out = config.validate_environment( + headers=headers, + model="m", + messages=[], + optional_params={}, + litellm_params={}, + ) + assert out == headers + + +# --------------------------------------------------------------------------- # +# Shared BaseBatchesConfig contract suite. +# --------------------------------------------------------------------------- # + +from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 + BatchesConfigContractTests, +) + + +class TestBedrockBatchesContract(BatchesConfigContractTests): + def make_config(self): + return BedrockBatchesConfig() + + expected_provider = LlmProviders.BEDROCK + # Bedrock builds a real create request (no NotImplementedError); the + # bedrock-specific create tests above cover the request/response shape. + supports_create = True + # Retrieve responses are parsed in this transformation layer + # (transform_retrieve_batch_response). + supports_retrieve_response = True + + def sample_retrieve_response_body(self) -> dict: + return { + "jobArn": ARN, + "status": "Completed", + "submitTime": "2026-04-28T12:00:00Z", + "endTime": "2026-04-28T12:30:00Z", + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": "s3://b/in.jsonl"}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": "s3://b/out/"}}, + } + + expected_retrieve_batch_id = ARN + expected_retrieve_status = "completed" diff --git a/tests/test_litellm/llms/vertex_ai/batches/__init__.py b/tests/test_litellm/llms/vertex_ai/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_handler.py b/tests/test_litellm/llms/vertex_ai/batches/test_handler.py new file mode 100644 index 00000000000..cacea234777 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/batches/test_handler.py @@ -0,0 +1,805 @@ +""" +Unit tests for ``VertexAIBatchPrediction`` (litellm/llms/vertex_ai/batches/handler.py). + +The handler is HTTP/auth glue around the (separately-tested) pure +``VertexAIBatchTransformation``. Each public method (create / retrieve / list / +cancel) resolves a Vertex access token + URL, branches on ``_is_async`` +(returning the coroutine in the async case, doing the sync HTTP call otherwise), +checks the HTTP status, and parses the JSON into ``LiteLLMBatch`` (or the OpenAI +list shape). + +We mock only true I/O / auth seams: + * ``_ensure_access_token`` - the Vertex credential seam. Returns a fixed + (token, project) so we can assert the ``Authorization: Bearer `` + header is forwarded. + * ``_check_custom_proxy`` - returns ``(None, url)``; we let it pass the + computed default url straight through so we can assert the request URL. + * the httpx client factories (``_get_httpx_client`` / + ``get_async_httpx_client``) and the SSRF wrappers (``safe_get`` / + ``async_safe_get``) - the network calls. We assert which seam fired with + what URL/headers/body, and that the response is parsed into the litellm + type. Sibling seams are asserted NOT called where relevant. + +The ``_is_async`` branch, status-code error paths, and the cancel +retrieve-after-cancel sequencing run for real. +""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.vertex_ai.batches.handler import ( # noqa: E402 + VertexAIBatchPrediction, +) +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +HMOD = "litellm.llms.vertex_ai.batches.handler" +TOKEN = "ya29.fake-access-token" +PROJECT = "my-project" +LOCATION = "us-central1" +BATCH_ID = "3814889423749775360" + +CREATE_DATA = { + "input_file_id": ( + "gs://bucket/publishers/google/models/gemini-1.5-flash-001/file-uuid" + ) +} + + +def _vertex_job_response(state: str = "JOB_STATE_SUCCEEDED") -> dict: + return { + "name": f"projects/p/locations/{LOCATION}/batchPredictionJobs/{BATCH_ID}", + "state": state, + "createTime": "2024-12-04T21:53:12.120184Z", + "inputConfig": { + "instancesFormat": "jsonl", + "gcsSource": {"uris": ["gs://bucket/in.jsonl"]}, + }, + "outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}, + } + + +def _http_response(status_code: int = 200, json_body: dict | None = None) -> MagicMock: + resp = MagicMock() + resp.status_code = status_code + resp.text = "error text" + resp.json.return_value = json_body if json_body is not None else _vertex_job_response() + return resp + + +def _make_handler() -> VertexAIBatchPrediction: + """Construct the handler with auth + proxy seams patched at the instance level. + + ``_ensure_access_token`` and ``_check_custom_proxy`` are inherited from + ``VertexLLM``; we patch them on the instance (DI-style) so the URL/auth + plumbing is deterministic and we can assert what got forwarded downstream. + """ + h = VertexAIBatchPrediction(gcs_bucket_name="litellm-testing-bucket") + h._ensure_access_token = MagicMock(return_value=(TOKEN, PROJECT)) # type: ignore[method-assign] + # pass the computed default url straight through (no custom proxy) + h._check_custom_proxy = MagicMock( # type: ignore[method-assign] + side_effect=lambda **kw: (None, kw["url"]) + ) + return h + + +def _run(coro): + return asyncio.run(coro) + + +# =========================================================================== # +# create_vertex_batch_url +# =========================================================================== # + + +def test_create_vertex_batch_url(): + h = _make_handler() + url = h.create_vertex_batch_url(vertex_location=LOCATION, vertex_project=PROJECT) + assert url == ( + f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}" + f"/locations/{LOCATION}/batchPredictionJobs" + ) + + +# =========================================================================== # +# create_batch +# =========================================================================== # + + +def test_create_batch_sync_posts_and_parses(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response() + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + assert out.status == "completed" + + # auth seam fired + h._ensure_access_token.assert_called_once() + # the POST hit the batchPredictionJobs collection url with bearer auth + _, kwargs = client.post.call_args + assert kwargs["url"].endswith(f"/projects/{PROJECT}/locations/{LOCATION}/batchPredictionJobs") + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + # body is the transformed vertex job (json-serialized) + sent = json.loads(kwargs["data"]) + assert sent["model"] == "publishers/google/models/gemini-1.5-flash-001" + assert sent["inputConfig"]["gcsSource"]["uris"] == [CREATE_DATA["input_file_id"]] + + +def test_create_batch_async_returns_coroutine_and_uses_async_client(): + h = _make_handler() + async_client = MagicMock() + async_client.post = AsyncMock(return_value=_http_response()) + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.create_batch( + _is_async=True, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + async_client.post.assert_awaited_once() + # the async branch must NOT use the sync client for the request + sync_client.post.assert_not_called() + + +def test_create_batch_sync_non_200_raises(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response(status_code=500) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + with pytest.raises(Exception, match="Error: 500"): + h.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +def test_create_batch_async_non_200_raises(): + h = _make_handler() + async_client = MagicMock() + async_client.post = AsyncMock(return_value=_http_response(status_code=403)) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.create_batch( + _is_async=True, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(Exception, match="Error: 403"): + _run(coro) + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve_batch_sync_uses_safe_get_with_batch_id_url(): + h = _make_handler() + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.safe_get", return_value=_http_response()) as safe_get, + ): + out = h.retrieve_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + # SSRF-wrapped fetch fired with the batch-id-appended url + bearer header + args, kwargs = safe_get.call_args + assert args[0] is sync_client + assert args[1].endswith(f"/batchPredictionJobs/{BATCH_ID}") + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + # plain client.get must NOT be used (SSRF wrapper is the seam) + sync_client.get.assert_not_called() + + +def test_retrieve_batch_async_returns_coroutine_uses_async_safe_get(): + h = _make_handler() + async_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response()), + ) as async_safe_get, + ): + coro = h.retrieve_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert isinstance(out, LiteLLMBatch) + async_safe_get.assert_awaited_once() + args, _ = async_safe_get.await_args + assert args[1].endswith(f"/batchPredictionJobs/{BATCH_ID}") + + +def test_retrieve_batch_sync_non_200_raises(): + h = _make_handler() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.safe_get", return_value=_http_response(status_code=404)), + ): + with pytest.raises(Exception, match="Error: 404"): + h.retrieve_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +def test_retrieve_batch_sync_invokes_logging_pre_call(): + """When a real ``Logging`` obj is passed, ``pre_call`` is invoked with the + request url + headers (the curl-redaction branch).""" + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + h = _make_handler() + logging_obj = MagicMock(spec=LiteLLMLogging) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.safe_get", return_value=_http_response()), + ): + h.retrieve_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + logging_obj=logging_obj, + ) + + logging_obj.pre_call.assert_called_once() + _, kwargs = logging_obj.pre_call.call_args + assert kwargs["additional_args"]["api_base"].endswith( + f"/batchPredictionJobs/{BATCH_ID}" + ) + + +# =========================================================================== # +# list_batches +# =========================================================================== # + + +def _list_response() -> dict: + return { + "batchPredictionJobs": [_vertex_job_response()], + "nextPageToken": "next-tok", + } + + +def test_list_batches_sync_passes_pagination_params(): + h = _make_handler() + client = MagicMock() + client.get.return_value = _http_response(json_body=_list_response()) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.list_batches( + _is_async=False, + after="cursor-xyz", + limit=7, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert out["object"] == "list" + assert out["data"][0].id == BATCH_ID + assert out["has_more"] is True + assert out["next_page_token"] == "next-tok" + + _, kwargs = client.get.call_args + # limit -> pageSize (stringified), after -> pageToken + assert kwargs["params"] == {"pageSize": "7", "pageToken": "cursor-xyz"} + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + + +def test_list_batches_sync_omits_unset_pagination_params(): + h = _make_handler() + client = MagicMock() + client.get.return_value = _http_response(json_body={"batchPredictionJobs": []}) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.list_batches( + _is_async=False, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + _, kwargs = client.get.call_args + assert kwargs["params"] == {} + assert out["data"] == [] + assert out["has_more"] is False + + +def test_list_batches_async_returns_coroutine(): + h = _make_handler() + async_client = MagicMock() + async_client.get = AsyncMock( + return_value=_http_response(json_body=_list_response()) + ) + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.list_batches( + _is_async=True, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert out["data"][0].id == BATCH_ID + async_client.get.assert_awaited_once() + sync_client.get.assert_not_called() + + +def test_list_batches_sync_non_200_raises(): + h = _make_handler() + client = MagicMock() + client.get.return_value = _http_response(status_code=500) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + with pytest.raises(Exception, match="Error: 500"): + h.list_batches( + _is_async=False, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +# =========================================================================== # +# cancel_batch +# =========================================================================== # + + +def test_cancel_batch_sync_posts_cancel_then_retrieves(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response(json_body={}) + client.get.return_value = _http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.status == "cancelled" + + # POST hit the :cancel url + _, post_kwargs = client.post.call_args + assert post_kwargs["url"].endswith(f"/batchPredictionJobs/{BATCH_ID}:cancel") + assert post_kwargs["data"] == json.dumps({}) + # then GET hit the plain retrieve url (no :cancel suffix) + _, get_kwargs = client.get.call_args + assert get_kwargs["url"].endswith(f"/batchPredictionJobs/{BATCH_ID}") + assert not get_kwargs["url"].endswith(":cancel") + + +def test_cancel_batch_async_returns_coroutine_posts_then_retrieves(): + h = _make_handler() + async_client = MagicMock() + async_client.post = AsyncMock(return_value=_http_response(json_body={})) + async_client.get = AsyncMock( + return_value=_http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + ) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.cancel_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert out.status == "cancelled" + async_client.post.assert_awaited_once() + async_client.get.assert_awaited_once() + _, post_kwargs = async_client.post.await_args + assert post_kwargs["url"].endswith(":cancel") + + +def test_cancel_batch_sync_cancel_post_non_200_raises(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response(status_code=500) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + with pytest.raises(Exception, match="Error: 500"): + h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + # cancel POST failed -> retrieve GET must never fire + client.get.assert_not_called() + + +def test_cancel_batch_sync_retrieve_non_200_raises(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response(json_body={}) + client.get.return_value = _http_response(status_code=404) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + with pytest.raises(Exception, match="Error: 404"): + h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +def test_cancel_batch_sync_proxy_url_without_cancel_suffix_uses_rsplit_branch(): + """If ``_check_custom_proxy`` hands back a url that does NOT end in + ``:cancel`` (e.g. a custom proxy rewrote it), the retrieve url is derived + via the ``rsplit(':cancel')`` else-branch rather than ``removesuffix``.""" + h = _make_handler() + # override the proxy seam to return a non-:cancel-suffixed url + h._check_custom_proxy = MagicMock( # type: ignore[method-assign] + return_value=(None, "https://proxy.internal/vertex/batch") + ) + client = MagicMock() + client.post.return_value = _http_response(json_body={}) + client.get.return_value = _http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base="https://proxy.internal", + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert out.status == "cancelled" + _, get_kwargs = client.get.call_args + # rsplit(":cancel")[0].rstrip("/") of a url with no :cancel -> url unchanged + assert get_kwargs["url"] == "https://proxy.internal/vertex/batch" + + +def test_cancel_batch_sync_httpstatuserror_logged_and_reraised(): + """The cancel POST ``httpx.HTTPStatusError`` except-branch logs + re-raises.""" + h = _make_handler() + client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs/1:cancel") + err_response = httpx.Response(status_code=502, request=request, text="bad gw") + client.post.side_effect = httpx.HTTPStatusError( + "boom", request=request, response=err_response + ) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + with pytest.raises(httpx.HTTPStatusError): + h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + client.get.assert_not_called() + + +def test_create_batch_async_httpstatuserror_logged_and_reraised(): + h = _make_handler() + async_client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs") + err_response = httpx.Response(status_code=500, request=request, text="boom") + async_client.post = AsyncMock( + side_effect=httpx.HTTPStatusError( + "boom", request=request, response=err_response + ) + ) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.create_batch( + _is_async=True, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(httpx.HTTPStatusError): + _run(coro) + + +def test_async_retrieve_batch_non_200_raises(): + h = _make_handler() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response(status_code=500)), + ), + ): + coro = h.retrieve_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(Exception, match="Error: 500"): + _run(coro) + + +def test_async_retrieve_batch_invokes_logging_pre_call(): + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + h = _make_handler() + logging_obj = MagicMock(spec=LiteLLMLogging) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response()), + ), + ): + coro = h.retrieve_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + logging_obj=logging_obj, + ) + _run(coro) + + logging_obj.pre_call.assert_called_once() + + +def test_async_list_batches_non_200_raises(): + h = _make_handler() + async_client = MagicMock() + async_client.get = AsyncMock(return_value=_http_response(status_code=500)) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.list_batches( + _is_async=True, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(Exception, match="Error: 500"): + _run(coro) + + +def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200(): + """Async cancel: POST HTTPStatusError re-raises; and separately the + retrieve-after-cancel non-200 raises.""" + h = _make_handler() + + # (a) POST raises HTTPStatusError + async_client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs/1:cancel") + err_response = httpx.Response(status_code=502, request=request, text="bad") + async_client.post = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=request, response=err_response) + ) + async_client.get = AsyncMock() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.cancel_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(httpx.HTTPStatusError): + _run(coro) + async_client.get.assert_not_awaited() + + # (a2) cancel POST returns a plain non-200 (no exception) -> raises + async_client_post500 = MagicMock() + async_client_post500.post = AsyncMock(return_value=_http_response(status_code=500)) + async_client_post500.get = AsyncMock() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client_post500), + ): + coro = h.cancel_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(Exception, match="Error: 500"): + _run(coro) + async_client_post500.get.assert_not_awaited() + + # (b) retrieve-after-cancel returns non-200 + async_client2 = MagicMock() + async_client2.post = AsyncMock(return_value=_http_response(json_body={})) + async_client2.get = AsyncMock(return_value=_http_response(status_code=404)) + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client2), + ): + coro = h.cancel_batch( + _is_async=True, + batch_id=BATCH_ID, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + with pytest.raises(Exception, match="Error: 404"): + _run(coro) diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py new file mode 100644 index 00000000000..37084d43441 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py @@ -0,0 +1,396 @@ +""" +Unit tests for ``VertexAIBatchTransformation`` +(litellm/llms/vertex_ai/batches/transformation.py). + +This module is pure transformation logic: it maps OpenAI-shaped batch requests +into Vertex AI ``VertexAIBatchPredictionJob`` payloads, and maps Vertex AI batch +responses back into ``LiteLLMBatch`` / OpenAI list shapes. Unlike anthropic / +bedrock, this class does NOT subclass ``BaseBatchesConfig`` - it's a standalone +set of classmethods with a Vertex-specific shape, so these tests are fully +standalone and assert exact values rather than "ran without error". + +There are no real I/O seams here; ``uuid.uuid4`` is the only nondeterministic +dependency and is patched where the displayName is asserted. +""" + +import os +import sys +from unittest.mock import patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402 + VertexAIBatchTransformation, +) +from litellm.llms.vertex_ai.common_utils import ( # noqa: E402 + _convert_vertex_datetime_to_openai_datetime, +) +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +T = VertexAIBatchTransformation + +INPUT_FILE = ( + "gs://litellm-testing-bucket/litellm-vertex-files/publishers/google/" + "models/gemini-1.5-flash-001/e9412502-2c91-42a6-8e61-f5c294cc0fc8" +) + + +# =========================================================================== # +# transform_openai_batch_request_to_vertex_ai_batch_request +# =========================================================================== # + + +def test_transform_openai_request_builds_full_vertex_job(): + with patch( + "litellm.llms.vertex_ai.batches.transformation.uuid.uuid4", + return_value="fixed-uuid", + ): + job = T.transform_openai_batch_request_to_vertex_ai_batch_request( + {"input_file_id": INPUT_FILE} + ) + + assert job["displayName"] == "litellm-vertex-batch-fixed-uuid" + assert job["model"] == "publishers/google/models/gemini-1.5-flash-001" + + assert job["inputConfig"]["instancesFormat"] == "jsonl" + assert job["inputConfig"]["gcsSource"]["uris"] == [INPUT_FILE] + + assert job["outputConfig"]["predictionsFormat"] == "jsonl" + # gcs uri prefix == file path with the filename stripped + assert ( + job["outputConfig"]["gcsDestination"]["outputUriPrefix"] + == "gs://litellm-testing-bucket/litellm-vertex-files/publishers/google/" + "models/gemini-1.5-flash-001" + ) + + +def test_transform_openai_request_missing_input_file_id_raises(): + with pytest.raises(ValueError, match="input_file_id is required"): + T.transform_openai_batch_request_to_vertex_ai_batch_request({}) + + +# =========================================================================== # +# transform_vertex_ai_batch_response_to_openai_batch_response +# =========================================================================== # + + +def test_transform_vertex_response_full_mapping(): + response = { + "name": "projects/510528649030/locations/us-central1/batchPredictionJobs/3814889423749775360", + "state": "JOB_STATE_SUCCEEDED", + "createTime": "2024-12-04T21:53:12.120184Z", + "inputConfig": { + "instancesFormat": "jsonl", + "gcsSource": {"uris": ["gs://bucket/in.jsonl"]}, + }, + "outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}, + } + batch = T.transform_vertex_ai_batch_response_to_openai_batch_response(response) + + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "3814889423749775360" + assert batch.completion_window == "24hrs" + # created_at is parsed via the shared helper (uses local tz); assert the + # transform forwards createTime through that helper rather than a hardcoded + # epoch that would be tz-dependent + assert batch.created_at == _convert_vertex_datetime_to_openai_datetime( + "2024-12-04T21:53:12.120184Z" + ) + assert batch.endpoint == "" + assert batch.object == "batch" + assert batch.input_file_id == "gs://bucket/in.jsonl" + assert batch.status == "completed" + assert batch.error_file_id is None + assert batch.output_file_id == "gs://bucket/out/predictions.jsonl" + + +def test_transform_vertex_response_error_file_id_always_none(): + batch = T.transform_vertex_ai_batch_response_to_openai_batch_response( + { + "name": "x/y/123", + "state": "JOB_STATE_FAILED", + "createTime": "2024-12-04T21:53:12.120184Z", + } + ) + assert batch.error_file_id is None + + +# =========================================================================== # +# _get_batch_job_status_from_vertex_ai_batch_response (test EVERY entry) +# =========================================================================== # + + +@pytest.mark.parametrize( + "vertex_state,expected", + [ + ("JOB_STATE_UNSPECIFIED", "failed"), + ("JOB_STATE_QUEUED", "validating"), + ("JOB_STATE_PENDING", "validating"), + ("JOB_STATE_RUNNING", "in_progress"), + ("JOB_STATE_SUCCEEDED", "completed"), + ("JOB_STATE_FAILED", "failed"), + ("JOB_STATE_CANCELLING", "cancelling"), + ("JOB_STATE_CANCELLED", "cancelled"), + ("JOB_STATE_PAUSED", "in_progress"), + ("JOB_STATE_EXPIRED", "expired"), + ("JOB_STATE_UPDATING", "in_progress"), + ("JOB_STATE_PARTIALLY_SUCCEEDED", "completed"), + ], +) +def test_status_mapping_every_entry(vertex_state, expected): + assert ( + T._get_batch_job_status_from_vertex_ai_batch_response({"state": vertex_state}) + == expected + ) + + +def test_status_mapping_defaults_to_unspecified_when_missing(): + # No "state" key -> defaults to JOB_STATE_UNSPECIFIED -> "failed" + assert T._get_batch_job_status_from_vertex_ai_batch_response({}) == "failed" + + +def test_status_mapping_unknown_state_raises_keyerror(): + with pytest.raises(KeyError): + T._get_batch_job_status_from_vertex_ai_batch_response({"state": "NOPE"}) + + +# =========================================================================== # +# _get_batch_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_batch_id_splits_path(): + assert ( + T._get_batch_id_from_vertex_ai_batch_response( + {"name": "projects/p/locations/l/batchPredictionJobs/999"} + ) + == "999" + ) + + +def test_get_batch_id_no_slash_returns_name(): + assert T._get_batch_id_from_vertex_ai_batch_response({"name": "abc"}) == "abc" + + +def test_get_batch_id_empty_name_returns_empty(): + assert T._get_batch_id_from_vertex_ai_batch_response({"name": ""}) == "" + assert T._get_batch_id_from_vertex_ai_batch_response({}) == "" + + +# =========================================================================== # +# _get_input_file_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_input_file_id_happy_path(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response( + {"inputConfig": {"gcsSource": {"uris": ["gs://b/a.jsonl", "gs://b/c.jsonl"]}}} + ) + == "gs://b/a.jsonl" + ) + + +def test_get_input_file_id_missing_input_config(): + assert T._get_input_file_id_from_vertex_ai_batch_response({}) == "" + + +def test_get_input_file_id_missing_gcs_source(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response({"inputConfig": {}}) == "" + ) + + +def test_get_input_file_id_empty_uris(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response( + {"inputConfig": {"gcsSource": {"uris": []}}} + ) + == "" + ) + + +# =========================================================================== # +# _get_output_file_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_output_file_id_from_output_info(): + # outputInfo branch: rstrip trailing slash, append predictions.jsonl + assert ( + T._get_output_file_id_from_vertex_ai_batch_response( + {"outputInfo": {"gcsOutputDirectory": "gs://bucket/out/"}} + ) + == "gs://bucket/out/predictions.jsonl" + ) + + +def test_get_output_file_id_output_info_no_trailing_slash(): + assert ( + T._get_output_file_id_from_vertex_ai_batch_response( + {"outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}} + ) + == "gs://bucket/out/predictions.jsonl" + ) + + +def test_get_output_file_id_empty_output_info_falls_through_to_output_config(): + # gcsOutputDirectory missing -> "" -> the "/predictions.jsonl" guard skips + # the outputInfo branch, falls through to outputConfig + resp = { + "outputInfo": {}, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}}, + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_no_output_info_and_no_output_config(): + assert T._get_output_file_id_from_vertex_ai_batch_response({}) == "" + + +def test_get_output_file_id_output_config_missing_gcs_destination(): + # outputConfig present but no gcsDestination -> returns the running "" value + assert ( + T._get_output_file_id_from_vertex_ai_batch_response({"outputConfig": {}}) == "" + ) + + +def test_get_output_file_id_output_config_already_has_suffix(): + # outputUriPrefix already ends in /predictions.jsonl -> returned as-is (no double append) + resp = { + "outputConfig": { + "gcsDestination": {"outputUriPrefix": "gs://b/cfg/predictions.jsonl"} + } + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_output_config_strips_trailing_slash(): + resp = { + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/"}} + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_output_info_takes_precedence_over_output_config(): + resp = { + "outputInfo": {"gcsOutputDirectory": "gs://from-info"}, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://from-config"}}, + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://from-info/predictions.jsonl" + ) + + +# =========================================================================== # +# _get_gcs_uri_prefix_from_file +# =========================================================================== # + + +def test_get_gcs_uri_prefix_root(): + assert ( + T._get_gcs_uri_prefix_from_file("gs://litellm-testing-bucket/vtx_batch.jsonl") + == "gs://litellm-testing-bucket" + ) + + +def test_get_gcs_uri_prefix_nested(): + assert ( + T._get_gcs_uri_prefix_from_file( + "gs://litellm-testing-bucket/batches/vtx_batch.jsonl" + ) + == "gs://litellm-testing-bucket/batches" + ) + + +# =========================================================================== # +# _get_model_from_gcs_file +# =========================================================================== # + + +def test_get_model_from_gcs_file_plain(): + assert ( + T._get_model_from_gcs_file(INPUT_FILE) + == "publishers/google/models/gemini-1.5-flash-001" + ) + + +def test_get_model_from_gcs_file_url_encoded(): + # %2F decodes to "/" via urllib.unquote before splitting + encoded = ( + "gs://bucket/publishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2Fuuid" + ) + assert ( + T._get_model_from_gcs_file(encoded) + == "publishers/google/models/gemini-1.5-flash-001" + ) + + +def test_get_model_from_gcs_file_no_publishers_raises(): + with pytest.raises(IndexError): + T._get_model_from_gcs_file("gs://bucket/no-model-here.jsonl") + + +# =========================================================================== # +# transform_vertex_ai_batch_list_response_to_openai_list_response +# =========================================================================== # + + +def _job(batch_id: str) -> dict: + return { + "name": f"projects/p/locations/l/batchPredictionJobs/{batch_id}", + "state": "JOB_STATE_SUCCEEDED", + "createTime": "2024-12-04T21:53:12.120184Z", + } + + +def test_list_response_multiple_jobs(): + response = { + "batchPredictionJobs": [_job("111"), _job("222"), _job("333")], + "nextPageToken": "tok-abc", + } + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response(response) + + assert out["object"] == "list" + assert [b.id for b in out["data"]] == ["111", "222", "333"] + assert out["first_id"] == "111" + assert out["last_id"] == "333" + assert out["has_more"] is True + assert out["next_page_token"] == "tok-abc" + + +def test_list_response_no_next_page_token(): + response = {"batchPredictionJobs": [_job("111")]} + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response(response) + assert out["has_more"] is False + assert out["next_page_token"] is None + assert out["first_id"] == "111" + assert out["last_id"] == "111" + + +def test_list_response_empty(): + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response({}) + assert out["data"] == [] + assert out["first_id"] is None + assert out["last_id"] is None + assert out["has_more"] is False + + +def test_list_response_none_jobs_treated_as_empty(): + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response( + {"batchPredictionJobs": None} + ) + assert out["data"] == [] + assert out["first_id"] is None diff --git a/tests/test_litellm/proxy/batches_endpoints/__init__.py b/tests/test_litellm/proxy/batches_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py new file mode 100644 index 00000000000..26c654cd154 --- /dev/null +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -0,0 +1,2026 @@ +""" +Routing-contract tests for litellm/proxy/batches_endpoints/endpoints.py + +These are not happy-path smoke tests. Each row of the matrix locks the full +contract of a single routing branch so that *any* behavior change in this layer +fails loudly: + + 1. DISPATCH - exactly which downstream seam fired (litellm.acreate_batch + vs llm_router.acreate_batch), and every sibling seam is + asserted NOT called. A reordered/negated branch flips this. + 2. CREDENTIALS - the credential resolver receives the model derived from the + request, not a hardcoded value. The router's + get_deployment_credentials_with_provider is input-locked. + 3. SEAM PAYLOAD - the *entire* kwargs dict forwarded to the provider call is + exact-matched. Because LiteLLMBatchCreateRequest is a + TypedDict (zero runtime filtering), nothing else stops a + newly-added param from silently reaching every provider. + This exact-match is that missing guard: a new key fails the + test and forces a reviewer to ask "does this work for all + providers, or just openai". + 4. OUTPUT SHAPE - the id encode/decode round-trip clients depend on. + +Only true I/O boundaries are mocked (provider call, router, proxy logging, +request parsing, pre-call enrichment). The pure encode/decode/credential-merge +helpers run for real so the payload assertions reflect production exactly. + +The object mocks are spec'd to their real classes, so a brand-new method call +added to this layer raises instead of silently passing - the inventory of seams +cannot drift without a test failure. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.proxy.batches_endpoints.endpoints as endpoints +import litellm.proxy.proxy_server as proxy_server +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.openai_files_endpoints.common_utils import ( + encode_file_id_with_model, +) +from litellm.proxy.utils import ProxyLogging +from litellm.router import Router +from litellm.types.llms.openai import BatchJobStatus +from litellm.types.utils import LiteLLMBatch + +from fastapi import Response + +# --------------------------------------------------------------------------- # +# Fixtures: distinguishable credentials per model so a wrong/hardcoded model_id +# produces wrong creds (or KeyError) and is impossible to hide. +# --------------------------------------------------------------------------- # + +CREDS: Dict[str, Dict[str, str]] = { + "azure/gpt-4o": { + "custom_llm_provider": "azure", + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + }, + "vertex-model": { + "custom_llm_provider": "vertex_ai", + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + }, +} + +# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123". +AZURE_FILE_ID = encode_file_id_with_model( + "file-original123", "azure/gpt-4o", id_type="file" +) + + +def make_batch( + *, + id: str = "batch-provider-id", + output_file_id: Optional[str] = None, + error_file_id: Optional[str] = None, + input_file_id: Optional[str] = None, + status: BatchJobStatus = "validating", +) -> LiteLLMBatch: + batch = LiteLLMBatch( + id=id, + completion_window="24h", + created_at=1234567890, + endpoint="/v1/chat/completions", + input_file_id=input_file_id or "file-provider-input", + object="batch", + status=status, + ) + if output_file_id is not None: + batch.output_file_id = output_file_id + if error_file_id is not None: + batch.error_file_id = error_file_id + batch._hidden_params = {} + return batch + + +class FakeRequest: + """Minimal stand-in. The request is only read via .headers/.query_params on + the model-param fallback path; everything else that touches it is mocked.""" + + def __init__( + self, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + ): + self.headers = headers or {} + self.query_params = query or {} + + +@dataclass +class Harness: + """Holds every mocked seam so a test can configure inputs and assert calls.""" + + body: Dict[str, Any] + read_body: AsyncMock + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + is_known_model: MagicMock + litellm_acreate: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + + @property + def router_acreate(self) -> AsyncMock: + return self.router.acreate_batch + + def acreate_kwargs(self) -> Dict[str, Any]: + """Exact kwargs forwarded to litellm.acreate_batch.""" + assert self.litellm_acreate.call_count == 1 + return dict(self.litellm_acreate.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_acreate.call_count == 1 + return dict(self.router_acreate.call_args.kwargs) + + +def _creds_lookup(*, model_id: str) -> Dict[str, str]: + # KeyError on an unknown/hardcoded model_id - the bug cannot hide. + return dict(CREDS[model_id]) + + +@pytest.fixture +def harness(): + """Seam harness. Patches only true I/O boundaries; pure encode/decode/merge + helpers run for real. Object mocks are spec'd so unknown method calls raise.""" + body_holder: Dict[str, Any] = {} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.acreate_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) + pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + is_known_model = MagicMock(return_value=False) + litellm_acreate = AsyncMock(return_value=make_batch()) + + with ExitStack() as stack: + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context(patch.object(endpoints, "is_known_model", is_known_model)) + stack.enter_context(patch.object(litellm, "acreate_batch", litellm_acreate)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + h = Harness( + body=body_holder, + read_body=read_body, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + is_known_model=is_known_model, + litellm_acreate=litellm_acreate, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + ) + yield h + + +def set_body(harness: Harness, body: Dict[str, Any]) -> None: + harness.body["body"] = body + + +async def call_create( + harness: Harness, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, +): + return await endpoints.create_batch( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + provider=provider, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + ) + + +# =========================================================================== # +# SCENARIO 1 - input_file_id encoded with model. The full showcase: every +# assertion type from the design lives here. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + resp = await call_create(harness) + + # 1. DISPATCH - model-credential path fired via litellm, router did not. + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + + # 2. CREDENTIALS - resolved for the model decoded FROM the file id. + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # 3. SEAM PAYLOAD - exact, whole dict. A new forwarded key breaks this. + assert harness.acreate_kwargs() == { + "custom_llm_provider": "azure", + "input_file_id": "file-original123", # encoding stripped by this layer + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, # sanitize_openai_provider_metadata(None) + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + } + + # 4. OUTPUT SHAPE - ids re-encoded with the model; input_file_id restored. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + assert resp.input_file_id == AZURE_FILE_ID + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id__encodes_output_and_error_ids(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.litellm_acreate.return_value = make_batch( + id="batch-xyz", + output_file_id="file-out-raw", + error_file_id="file-err-raw", + ) + + resp = await call_create(harness) + + assert resp.output_file_id == encode_file_id_with_model( + "file-out-raw", "azure/gpt-4o" + ) + assert resp.error_file_id == encode_file_id_with_model( + "file-err-raw", "azure/gpt-4o" + ) + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id__resolver_gets_decoded_model(harness): + """Regression guard: model_id for credential resolution must be derived from + the file id. A hardcode would call the resolver with the wrong model.""" + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# SCENARIO 2 - model from body / header / query. Locks source precedence. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__model_from_body(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "vertex-model", + }, + ) + + resp = await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + payload = harness.acreate_kwargs() + assert payload["custom_llm_provider"] == "vertex_ai" + assert payload["input_file_id"] == "file-plain" + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "vertex-model", id_type="batch" + ) + + +@pytest.mark.asyncio +async def test_create__model_from_header(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, headers={"x-litellm-model": "vertex-model"}) + + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + harness.router_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_from_query(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, query={"model": "vertex-model"}) + + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + harness.router_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__body_model_beats_header_and_query(harness): + """Precedence row: body > header > query (data.get('model') first).""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "azure/gpt-4o", + }, + ) + + await call_create( + harness, + headers={"x-litellm-model": "vertex-model"}, + query={"model": "vertex-model"}, + ) + + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__fallback_default_openai(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_not_called() # inverse-bug guard + assert harness.acreate_kwargs()["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_create__fallback_provider_path_param(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, provider="anthropic") + + harness.creds_resolver.assert_not_called() + assert harness.acreate_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_create__fallback_body_custom_llm_provider(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "custom_llm_provider": "bedrock", + }, + ) + + await call_create(harness) + + payload = harness.acreate_kwargs() + assert payload["custom_llm_provider"] == "bedrock" + + +# =========================================================================== # +# Unified file id routing (-> llm_router). Helpers mocked only here because a +# real unified id is opaque base64; the routing contract is what we lock. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__unified_file_id_single_model(harness): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"] + ): + resp = await call_create(harness) + + # DISPATCH - router fired, direct litellm did not. + assert harness.router_acreate.call_count == 1 + harness.litellm_acreate.assert_not_called() + # model injected from the unified id, input_file_id restored, hidden param set + assert harness.router_kwargs()["model"] == "gpt-4o-mini" + assert resp.input_file_id == "litellm_proxy_unified_id" + assert resp._hidden_params["unified_file_id"] == "unified-xyz" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("models", [[], ["m1", "m2"]]) +async def test_create__unified_file_id_not_exactly_one_model_400(harness, models): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=models + ): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "400" + harness.router_acreate.assert_not_called() + harness.litellm_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_encoded_beats_unified(harness): + """Precedence row: a file id that is BOTH model-encoded and (pretend) unified + must take the model-encoded branch (checked first).""" + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["something-else"] + ): + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# Loadbalancing branch (-> llm_router) and its precedence vs model-encoded. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__loadbalancing_routes_to_router(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_create(harness) + + harness.is_known_model.assert_called_once_with( + model="lb-model", llm_router=harness.router + ) + assert harness.router_acreate.call_count == 1 + harness.litellm_acreate.assert_not_called() + harness.creds_resolver.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_encoded_beats_loadbalancing(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# Team-level batch expiry enforcement (independent of routing). +# =========================================================================== # + + +def _user_with_expiry(expiry: Any) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + team_metadata={"enforced_batch_output_expires_after": expiry}, + ) + + +@pytest.mark.asyncio +async def test_create__team_expiry_injected(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create( + harness, user=_user_with_expiry({"anchor": "created_at", "seconds": 3600}) + ) + + assert harness.acreate_kwargs()["output_expires_after"] == { + "anchor": "created_at", + "seconds": 3600, + } + + +@pytest.mark.asyncio +async def test_create__no_team_expiry_not_injected(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, user=UserAPIKeyAuth(api_key="sk-test")) + + assert "output_expires_after" not in harness.acreate_kwargs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "expiry", + [ + {"seconds": 3600}, # missing anchor + {"anchor": "created_at"}, # missing seconds + {"anchor": "completed_at", "seconds": 3600}, # wrong anchor + ], +) +async def test_create__team_expiry_malformed_500(harness, expiry): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + with pytest.raises(ProxyException) as exc: + await call_create(harness, user=_user_with_expiry(expiry)) + + assert exc.value.code == "500" + harness.litellm_acreate.assert_not_called() + harness.router_acreate.assert_not_called() + + +# =========================================================================== # +# Cross-cutting: enrichment route_type, metadata sanitization, failure hook. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__uses_acreate_batch_route_type(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + assert harness.pre_call.call_args.kwargs["route_type"] == "acreate_batch" + + +@pytest.mark.asyncio +async def test_create__metadata_sanitized_before_forwarding(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": {"user_key": "user_val", "spend_logs_metadata": {"x": 1}}, + }, + ) + + await call_create(harness) + + # provider-internal-only key dropped, string key kept (real sanitize runs) + assert harness.acreate_kwargs()["metadata"] == {"user_key": "user_val"} + + +@pytest.mark.asyncio +async def test_create__exception_calls_failure_hook(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.litellm_acreate.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_create(harness) + + harness.logging.post_call_failure_hook.assert_called_once() + assert ( + harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# GET /v1/batches/{batch_id} - retrieve_batch routing-contract tests # +# # +# Same discipline as create_batch above. retrieve_batch has more seams: it # +# first consults the ManagedObjectTable (get_batch_from_database) and may # +# short-circuit on a terminal-status row WITHOUT ever calling a provider, # +# then on a miss/non-terminal row routes to one of three downstream seams # +# (litellm.aretrieve_batch via model creds, llm_router.aretrieve_batch, or # +# litellm.aretrieve_batch via env-var provider) and writes the fresh state # +# back via update_batch_in_database. Every test below locks exactly which # +# of those seams fired and asserts the siblings did NOT, so a reordered or # +# negated branch - or a dropped DB short-circuit / write-back - fails loud. # +# # +# The DB seams (get_batch_from_database, update_batch_in_database, # +# resolve_*_to_unified) are true prisma I/O boundaries and are mocked. The # +# encode/decode/credential-merge helpers run for real, so payload and id # +# round-trip assertions reflect production exactly. # +# =========================================================================== # + + +# A real model-encoded BATCH id: decodes to "azure/gpt-4o", strips to +# "batch_orig123". Distinct from AZURE_FILE_ID so retrieve tests can't pass by +# accidentally reusing the create fixture's value. +AZURE_BATCH_ID = encode_file_id_with_model( + "batch_orig123", "azure/gpt-4o", id_type="batch" +) + +# A realistic decoded unified batch id (what _is_base64_encoded_unified_file_id +# returns). model_id / llm_batch_id are parsed out of this by the real helpers. +UNIFIED_BATCH_ID = "litellm_proxy;model_id:gpt-4o-mini;llm_batch_id:batch-raw-xyz" + + +@dataclass +class RetrieveHarness: + """Seams for retrieve_batch. `data['data']` is the dict pre-call enrichment + returns; the routing branches mutate it, so it is reset per call.""" + + data: Dict[str, Any] + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_aretrieve: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + get_batch_from_db: AsyncMock + update_batch_in_db: AsyncMock + resolve_input: AsyncMock + resolve_output: AsyncMock + + @property + def router_aretrieve(self) -> AsyncMock: + return self.router.aretrieve_batch + + def aretrieve_kwargs(self) -> Dict[str, Any]: + """Exact kwargs forwarded to litellm.aretrieve_batch.""" + assert self.litellm_aretrieve.call_count == 1 + return dict(self.litellm_aretrieve.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_aretrieve.call_count == 1 + return dict(self.router_aretrieve.call_args.kwargs) + + +@pytest.fixture +def retrieve_harness(): + data_holder: Dict[str, Any] = {"data": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.aretrieve_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_aretrieve = AsyncMock(return_value=make_batch()) + # Default: DB miss -> always fall through to provider routing. + get_batch_from_db = AsyncMock(return_value=(None, None)) + update_batch_in_db = AsyncMock(return_value=None) + resolve_input = AsyncMock(return_value=None) + resolve_output = AsyncMock(return_value=None) + + with ExitStack() as stack: + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object(endpoints, "get_batch_from_database", get_batch_from_db) + ) + stack.enter_context( + patch.object(endpoints, "update_batch_in_database", update_batch_in_db) + ) + stack.enter_context( + patch.object(endpoints, "resolve_input_file_id_to_unified", resolve_input) + ) + stack.enter_context( + patch.object( + endpoints, "resolve_output_file_ids_to_unified", resolve_output + ) + ) + stack.enter_context(patch.object(litellm, "aretrieve_batch", litellm_aretrieve)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + stack.enter_context(patch.object(proxy_server, "prisma_client", MagicMock())) + + yield RetrieveHarness( + data=data_holder, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_aretrieve=litellm_aretrieve, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + get_batch_from_db=get_batch_from_db, + update_batch_in_db=update_batch_in_db, + resolve_input=resolve_input, + resolve_output=resolve_output, + ) + + +async def call_retrieve( + harness: RetrieveHarness, + batch_id: str, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, +): + # Mirror the real flow: data starts as RetrieveBatchRequest(batch_id=...). + harness.data["data"] = {"batch_id": batch_id} + return await endpoints.retrieve_batch( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + provider=provider, + batch_id=batch_id, + ) + + +# --------------------------------------------------------------------------- # +# SCENARIO 1 - batch id encoded with model. litellm.aretrieve_batch via the +# model's resolved credentials; response ids re-encoded for the client. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id(retrieve_harness): + resp = await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + # 1. DISPATCH - model-credential path fired via litellm, router did not. + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + + # 2. CREDENTIALS - resolved for the model decoded FROM the batch id. + retrieve_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # 3. SEAM PAYLOAD - exact, whole dict forwarded to the provider call. + # Note `model` is the DECODED model, not the deployment from creds: the + # endpoint overrides it (provider-config providers like bedrock need it). + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "azure", + "batch_id": "batch_orig123", # encoding stripped by this layer + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o", + } + + # 4. OUTPUT SHAPE - ids re-encoded with the model for the round-trip. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + + # write-back to the managed-object table happened, tagged as a retrieve. + assert retrieve_harness.update_batch_in_db.call_count == 1 + assert retrieve_harness.update_batch_in_db.call_args.kwargs["operation"] == "retrieve" + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id__forwards_decoded_model_not_deployment( + retrieve_harness, +): + """Regression guard for the line-483 override: the model forwarded to the + provider must be the decoded model id, never the deployment name that the + credential merge pulled in. Dropping the override silently 400s bedrock.""" + await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert retrieve_harness.aretrieve_kwargs()["model"] == "azure/gpt-4o" + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id__encodes_output_and_error_ids( + retrieve_harness, +): + retrieve_harness.litellm_aretrieve.return_value = make_batch( + id="batch-xyz", + output_file_id="file-out-raw", + error_file_id="file-err-raw", + ) + + resp = await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert resp.output_file_id == encode_file_id_with_model( + "file-out-raw", "azure/gpt-4o" + ) + assert resp.error_file_id == encode_file_id_with_model( + "file-err-raw", "azure/gpt-4o" + ) + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_beats_loadbalancing(retrieve_harness): + """Precedence: model-encoded id is checked before the loadbalancing/unified + elif, so it wins even with loadbalancing enabled.""" + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# --------------------------------------------------------------------------- # +# Unified managed batch id -> llm_router.aretrieve_batch. model_id is parsed +# out of the unified id and stamped onto hidden params; raw file ids on the +# response are resolved back to unified ids. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + resp = await call_retrieve(retrieve_harness, "batch-unified-blob") + + # DISPATCH - router fired, direct litellm did not. + assert retrieve_harness.router_aretrieve.call_count == 1 + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_not_called() + + # router receives the (still-encoded) batch id verbatim - this layer does + # not decode it for the unified path. + assert retrieve_harness.router_kwargs() == {"batch_id": "batch-unified-blob"} + + # hidden params: unified id passed through, model_id parsed from it. + assert resp._hidden_params["unified_batch_id"] == UNIFIED_BATCH_ID + assert resp._hidden_params["model_id"] == "gpt-4o-mini" + + # raw provider file ids on the response are resolved back to unified ids. + retrieve_harness.resolve_input.assert_called_once() + retrieve_harness.resolve_output.assert_called_once() + + +@pytest.mark.asyncio +async def test_retrieve__loadbalancing_raw_id_routes_to_router(retrieve_harness): + """Loadbalancing on + a plain (non-encoded, non-unified) batch id routes to + the router. Locks the current dispatch contract of the shared elif.""" + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + resp = await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.router_aretrieve.call_count == 1 + retrieve_harness.litellm_aretrieve.assert_not_called() + assert retrieve_harness.router_kwargs() == {"batch_id": "batch-raw-xyz"} + # not a unified id -> hidden param reflects that, no model_id stamped. + assert resp._hidden_params["unified_batch_id"] is False + assert "model_id" not in resp._hidden_params + # not a unified id -> no file-id resolution. + retrieve_harness.resolve_input.assert_not_called() + retrieve_harness.resolve_output.assert_not_called() + + +# --------------------------------------------------------------------------- # +# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver or the router. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__fallback_default_openai(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_not_called() # inverse-bug guard + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + assert retrieve_harness.update_batch_in_db.call_count == 1 + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_path_param(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") + + retrieve_harness.creds_resolver.assert_not_called() + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_from_header(retrieve_harness): + retrieve_harness.provider_from_headers.return_value = "bedrock" + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_from_query(retrieve_harness): + retrieve_harness.provider_from_query.return_value = "vertex_ai" + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_precedence_path_over_header( + retrieve_harness, +): + """provider path param beats the header-derived provider.""" + retrieve_harness.provider_from_headers.return_value = "bedrock" + + await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + + +# --------------------------------------------------------------------------- # +# ManagedObjectTable short-circuit. A terminal-status DB row is returned +# immediately - no provider call, no write-back. A non-terminal row falls +# through to a provider sync. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "status", ["completed", "complete", "failed", "cancelled", "expired"] +) +async def test_retrieve__db_terminal_state_short_circuits(retrieve_harness, status): + # "complete" is the DB-normalized alias of "completed"; it is not a valid + # constructor literal but reaches the endpoint via a stored row, so set it + # post-construction to exercise that exact branch. + db_response = make_batch(id="batch-from-db", status="completed") + db_response.status = status + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + resp = await call_retrieve(retrieve_harness, "batch-raw-xyz") + + # No provider seam fired, and no write-back (the row is already terminal). + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.update_batch_in_db.assert_not_called() + # The DB object is what the client gets back. + assert resp is db_response + + +@pytest.mark.asyncio +async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness): + db_response = make_batch(id="batch-from-db", status="completed") + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + await call_retrieve(retrieve_harness, "batch-unified-blob") + + # Terminal short-circuit still resolves raw provider file ids to unified. + retrieve_harness.resolve_input.assert_called_once() + retrieve_harness.resolve_output.assert_called_once() + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.router_aretrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harness): + """A non-terminal DB row must NOT short-circuit; the endpoint syncs with the + provider to refresh state.""" + db_response = make_batch(id="batch-from-db", status="validating") + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + # Provider sync happened despite the DB hit. + assert retrieve_harness.litellm_aretrieve.call_count == 1 + assert retrieve_harness.update_batch_in_db.call_count == 1 + + +# --------------------------------------------------------------------------- # +# Cross-cutting: enrichment route_type and failure-hook on provider error. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__uses_aretrieve_batch_route_type(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert ( + retrieve_harness.pre_call.call_args.kwargs["route_type"] == "aretrieve_batch" + ) + + +@pytest.mark.asyncio +async def test_retrieve__exception_calls_failure_hook(retrieve_harness): + retrieve_harness.litellm_aretrieve.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + retrieve_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + retrieve_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# GET /v1/batches - list_batches routing-contract tests # +# # +# Branch order (first match wins): # +# 1. managed_files hook present -> managed_files_obj.list_user_batches # +# 2. model from body/query/header -> litellm.alist_batches + id encode # +# 3. target_model_names (param or body) -> llm_router.alist_batches # +# 4. fallback -> litellm.alist_batches via env-var custom_llm_provider # +# # +# llm_router is required; absence is a 500 before any branch runs. # +# =========================================================================== # + + +class FakeListPage: + """Stand-in for the SyncCursorPage[Batch] that alist_batches returns. The + endpoint only touches `.data` (to encode ids) and `._hidden_params`.""" + + def __init__(self, data: Any): + self.data = data + self._hidden_params: Dict[str, Any] = {} + + +@dataclass +class ListHarness: + body: Dict[str, Any] + read_body: AsyncMock + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_alist: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + + @property + def router_alist(self) -> AsyncMock: + return self.router.alist_batches + + def set_managed_files(self, page: Any) -> AsyncMock: + """Install a managed_files hook exposing list_user_batches -> page.""" + hook = MagicMock() + hook.list_user_batches = AsyncMock(return_value=page) + self.logging.get_proxy_hook = MagicMock(return_value=hook) + return hook.list_user_batches + + def alist_kwargs(self) -> Dict[str, Any]: + assert self.litellm_alist.call_count == 1 + return dict(self.litellm_alist.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_alist.call_count == 1 + return dict(self.router_alist.call_args.kwargs) + + +@pytest.fixture +def list_harness(): + body_holder: Dict[str, Any] = {"body": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + # Default: no managed_files hook -> branches 2/3/4 are reachable. + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.alist_batches = AsyncMock(return_value=FakeListPage([])) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) + pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_alist = AsyncMock(return_value=FakeListPage([])) + + with ExitStack() as stack: + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context(patch.object(litellm, "alist_batches", litellm_alist)) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + yield ListHarness( + body=body_holder, + read_body=read_body, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_alist=litellm_alist, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + ) + + +async def call_list( + harness: ListHarness, + *, + provider: Optional[str] = None, + limit: Optional[int] = None, + after: Optional[str] = None, + target_model_names: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + body: Optional[Dict[str, Any]] = None, +): + harness.body["body"] = body if body is not None else {} + return await endpoints.list_batches( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + provider=provider, + limit=limit, + after=after, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + target_model_names=target_model_names, + ) + + +# --------------------------------------------------------------------------- # +# Branch 1 - ManagedObjectTable listing. This is the default production path +# (the managed_files hook is registered) and wins over every other branch. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__managed_files_path(list_harness): + page = FakeListPage([make_batch(id="batch-1")]) + list_user_batches = list_harness.set_managed_files(page) + + user = UserAPIKeyAuth(api_key="sk-test") + resp = await call_list( + list_harness, + user=user, + limit=7, + after="batch-cursor", + provider="openai", + target_model_names="m1,m2", + ) + + # DISPATCH - managed-files seam fired, neither provider seam did. + list_user_batches.assert_called_once_with( + user_api_key_dict=user, + limit=7, + after="batch-cursor", + provider="openai", + target_model_names="m1,m2", + llm_router=list_harness.router, + ) + list_harness.litellm_alist.assert_not_called() + list_harness.router_alist.assert_not_called() + assert resp is page + + +@pytest.mark.asyncio +async def test_list__managed_files_beats_model_param(list_harness): + """Branch 1 is checked before the model branch: a model in the body does not + divert away from managed-files listing.""" + page = FakeListPage([]) + list_user_batches = list_harness.set_managed_files(page) + + await call_list(list_harness, body={"model": "azure/gpt-4o"}) + + list_user_batches.assert_called_once() + list_harness.litellm_alist.assert_not_called() + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() + + +# --------------------------------------------------------------------------- # +# Branch 2 - model from body/query/header. CURRENTLY BROKEN: the endpoint +# forwards custom_llm_provider both explicitly and via **data (it calls +# data.update(credentials) but never pops custom_llm_provider the way +# create/retrieve do through prepare_data_with_credentials), so every call +# raises "multiple values for keyword argument 'custom_llm_provider'". +# +# The strict xfail below encodes the INTENDED contract (litellm seam fires, +# creds resolved for the body model, response ids encoded). It xfails today on +# the duplicate-kwarg TypeError; the day that branch is fixed it will XPASS and +# strict-mode turns the green into a failure, forcing whoever fixes it to drop +# the marker and adopt this as a live regression test. +# --------------------------------------------------------------------------- # + + +@pytest.mark.xfail( + strict=True, + raises=ProxyException, + reason="list_batches model branch passes custom_llm_provider twice " + "(explicit kwarg + **data after data.update(credentials)); remove when fixed", +) +@pytest.mark.asyncio +async def test_list__model_from_body_routes_and_encodes(list_harness): + list_harness.litellm_alist.return_value = FakeListPage( + [make_batch(id="batch-1"), make_batch(id="batch-2")] + ) + + resp = await call_list(list_harness, body={"model": "azure/gpt-4o"}) + + assert list_harness.litellm_alist.call_count == 1 + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + assert resp.data[0].id == encode_file_id_with_model( + "batch-1", "azure/gpt-4o", id_type="batch" + ) + assert resp.data[1].id == encode_file_id_with_model( + "batch-2", "azure/gpt-4o", id_type="batch" + ) + + +# --------------------------------------------------------------------------- # +# Branch 3 - target_model_names (function param or body) -> llm_router. Routes +# to the FIRST model in the comma list; `model` is stripped from data first. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__target_model_names_param_routes_to_router(list_harness): + await call_list(list_harness, target_model_names="m1,m2", limit=3, after="cur") + + assert list_harness.router_alist.call_count == 1 + list_harness.litellm_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() + # first model only; after/limit forwarded; nothing else (param not in data). + assert list_harness.router_kwargs() == { + "model": "m1", + "after": "cur", + "limit": 3, + } + + +@pytest.mark.asyncio +async def test_list__target_model_names_from_body(list_harness): + await call_list(list_harness, body={"target_model_names": "m1,m2"}) + + assert list_harness.router_alist.call_count == 1 + list_harness.litellm_alist.assert_not_called() + kwargs = list_harness.router_kwargs() + assert kwargs["model"] == "m1" + # body-sourced target_model_names stays in the forwarded data. + assert kwargs["target_model_names"] == "m1,m2" + + +@pytest.mark.asyncio +async def test_list__target_model_names_takes_first_only(list_harness): + """Locks the current behavior: with multiple target models, only the first + is routed to (silently, unlike create which 400s on >1). A change here - + intentional or not - must update this test.""" + await call_list(list_harness, target_model_names="alpha,beta,gamma") + + assert list_harness.router_kwargs()["model"] == "alpha" + + +# --------------------------------------------------------------------------- # +# Branch 4 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver or the router. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__fallback_default_openai(list_harness): + await call_list(list_harness) + + assert list_harness.litellm_alist.call_count == 1 + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() # inverse-bug guard + assert list_harness.alist_kwargs() == { + "custom_llm_provider": "openai", + "after": None, + "limit": None, + } + + +@pytest.mark.asyncio +async def test_list__fallback_provider_path_param(list_harness): + await call_list(list_harness, provider="anthropic") + + list_harness.creds_resolver.assert_not_called() + assert list_harness.alist_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_list__fallback_provider_from_header(list_harness): + list_harness.provider_from_headers.return_value = "bedrock" + + await call_list(list_harness) + + assert list_harness.alist_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_list__fallback_provider_from_query(list_harness): + list_harness.provider_from_query.return_value = "vertex_ai" + + await call_list(list_harness) + + assert list_harness.alist_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_list__fallback_after_and_limit_forwarded(list_harness): + await call_list(list_harness, after="cursor-9", limit=42) + + kwargs = list_harness.alist_kwargs() + assert kwargs["after"] == "cursor-9" + assert kwargs["limit"] == 42 + + +# --------------------------------------------------------------------------- # +# Cross-cutting: router requirement, route_type, failure hook. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__no_router_raises_500(list_harness): + with patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_list(list_harness) + + assert exc.value.code == "500" + list_harness.litellm_alist.assert_not_called() + + +@pytest.mark.asyncio +async def test_list__uses_alist_batches_route_type(list_harness): + await call_list(list_harness) + + assert list_harness.pre_call.call_args.kwargs["route_type"] == "alist_batches" + + +@pytest.mark.asyncio +async def test_list__exception_calls_failure_hook(list_harness): + list_harness.litellm_alist.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_list(list_harness) + + list_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + list_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# POST /v1/batches/{batch_id}/cancel - cancel_batch routing-contract tests # +# # +# Three branches, first match wins: # +# 1. model-encoded batch id -> litellm.acancel_batch via model creds # +# 2. unified batch id -> llm_router.acancel_batch (model+batch_id # +# parsed out of the unified id) # +# 3. fallback -> litellm.acancel_batch via env-var provider # +# Every branch then writes state back via update_batch_in_database( # +# operation="cancel"). There is NO ManagedObjectTable read short-circuit # +# here (unlike retrieve). # +# # +# These tests pin CURRENT behavior so a refactor can't silently change it. # +# Two current-behavior quirks are locked deliberately and noted inline: # +# - SCENARIO 1 forwards the DEPLOYMENT model from creds, not the decoded # +# model (retrieve overrides it; cancel does not). # +# - SCENARIO 3 rebuilds a CancelBatchRequest and forwards only # +# {custom_llm_provider, batch_id}, dropping enrichment keys. # +# =========================================================================== # + + +@dataclass +class CancelHarness: + data: Dict[str, Any] + pre_call: AsyncMock + add_data: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_acancel: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + update_batch_in_db: AsyncMock + + @property + def router_acancel(self) -> AsyncMock: + return self.router.acancel_batch + + def acancel_kwargs(self) -> Dict[str, Any]: + assert self.litellm_acancel.call_count == 1 + return dict(self.litellm_acancel.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_acancel.call_count == 1 + return dict(self.router_acancel.call_args.kwargs) + + +@pytest.fixture +def cancel_harness(): + data_holder: Dict[str, Any] = {"data": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.acancel_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) + # add_litellm_data_to_request is a passthrough that returns the data it got. + add_data = AsyncMock(side_effect=lambda **kw: kw["data"]) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_acancel = AsyncMock(return_value=make_batch()) + update_batch_in_db = AsyncMock(return_value=None) + + with ExitStack() as stack: + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object(endpoints, "update_batch_in_database", update_batch_in_db) + ) + stack.enter_context(patch.object(litellm, "acancel_batch", litellm_acancel)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + stack.enter_context(patch.object(proxy_server, "prisma_client", MagicMock())) + stack.enter_context( + patch.object(proxy_server, "add_litellm_data_to_request", add_data) + ) + + yield CancelHarness( + data=data_holder, + pre_call=pre_call, + add_data=add_data, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_acancel=litellm_acancel, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + update_batch_in_db=update_batch_in_db, + ) + + +async def call_cancel( + harness: CancelHarness, + batch_id: str, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + data_extra: Optional[Dict[str, Any]] = None, +): + harness.data["data"] = {"batch_id": batch_id, **(data_extra or {})} + return await endpoints.cancel_batch( + request=FakeRequest(headers=headers, query=query), + batch_id=batch_id, + fastapi_response=Response(), + provider=provider, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + ) + + +# --------------------------------------------------------------------------- # +# SCENARIO 1 - model-encoded batch id -> litellm.acancel_batch via model creds. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_id(cancel_harness): + resp = await call_cancel(cancel_harness, AZURE_BATCH_ID) + + # DISPATCH - model-credential path via litellm; router untouched. + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + + # CREDENTIALS - resolved for the model decoded from the batch id. + cancel_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # SEAM PAYLOAD - exact dict. NOTE current behavior: `model` is the + # DEPLOYMENT name from creds, NOT the decoded model (cancel, unlike + # retrieve, does not override it). Locking this guards the difference. + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "azure", + "batch_id": "batch_orig123", # decoded/stripped original id + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + } + + # OUTPUT SHAPE - response id re-encoded with the DECODED model. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + + # write-back tagged as a cancel. + assert cancel_harness.update_batch_in_db.call_count == 1 + assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel" + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_id_forwards_deployment_model(cancel_harness): + """Pin the current contract: cancel forwards the creds' deployment model. + If someone adds a decoded-model override (as retrieve has), this flips and + must be reviewed.""" + await call_cancel(cancel_harness, AZURE_BATCH_ID) + + assert cancel_harness.acancel_kwargs()["model"] == "azure/gpt-4o-deployment" + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_beats_unified(cancel_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + await call_cancel(cancel_harness, AZURE_BATCH_ID) + + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# --------------------------------------------------------------------------- # +# SCENARIO 2 - unified batch id -> llm_router.acancel_batch. model and batch_id +# are parsed out of the unified id; hidden params stamped. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__unified_batch_id_routes_to_router(cancel_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + resp = await call_cancel(cancel_harness, "batch-unified-blob") + + # DISPATCH - router fired, litellm did not, no creds lookup. + assert cancel_harness.router_acancel.call_count == 1 + cancel_harness.litellm_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_not_called() + + # model + batch_id are extracted from the unified id and forwarded. + assert cancel_harness.router_kwargs() == { + "batch_id": "batch-raw-xyz", + "model": "gpt-4o-mini", + } + + # hidden params: unified id passed through, model_id stamped from data. + assert resp._hidden_params["unified_batch_id"] == UNIFIED_BATCH_ID + assert resp._hidden_params["model_id"] == "gpt-4o-mini" + + assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel" + + +@pytest.mark.asyncio +async def test_cancel__unified_missing_model_id_400(cancel_harness): + # unified id with no model_id segment -> get_model_id returns None -> 400. + with patch.object( + endpoints, + "_is_base64_encoded_unified_file_id", + return_value="litellm_proxy;llm_batch_id:batch-xyz", + ): + with pytest.raises(ProxyException) as exc: + await call_cancel(cancel_harness, "batch-unified-blob") + + assert exc.value.code == "400" + cancel_harness.router_acancel.assert_not_called() + cancel_harness.litellm_acancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_cancel__unified_no_router_500(cancel_harness): + with patch.object(proxy_server, "llm_router", None), patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + with pytest.raises(ProxyException) as exc: + await call_cancel(cancel_harness, "batch-unified-blob") + + assert exc.value.code == "500" + + +# --------------------------------------------------------------------------- # +# SCENARIO 3 - fallback to custom_llm_provider. Rebuilds a CancelBatchRequest +# and forwards only {custom_llm_provider, batch_id}. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__fallback_default_openai(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_not_called() # inverse-bug guard + # current behavior: enrichment keys dropped; only these two forwarded. + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + assert cancel_harness.update_batch_in_db.call_count == 1 + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_path_param(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz", provider="anthropic") + + cancel_harness.creds_resolver.assert_not_called() + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_data_body(cancel_harness): + await call_cancel( + cancel_harness, "batch-raw-xyz", data_extra={"custom_llm_provider": "bedrock"} + ) + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_header(cancel_harness): + cancel_harness.provider_from_headers.return_value = "vertex_ai" + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_query(cancel_harness): + cancel_harness.provider_from_query.return_value = "azure" + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "azure" + + +@pytest.mark.xfail( + strict=True, + raises=ProxyException, + reason="cancel SCENARIO 3: `provider or data.pop('custom_llm_provider')` " + "short-circuits when provider (path param) is set, so a body " + "custom_llm_provider is left in data and forwarded twice -> duplicate-kwarg " + "TypeError. Intended: path param wins cleanly. Remove marker when fixed.", +) +@pytest.mark.asyncio +async def test_cancel__fallback_provider_precedence_path_over_body(cancel_harness): + """Intended contract: provider path param beats a body custom_llm_provider. + CURRENTLY raises because the `or` short-circuit skips the data.pop, leaving + the body value to collide with the explicit kwarg.""" + await call_cancel( + cancel_harness, + "batch-raw-xyz", + provider="anthropic", + data_extra={"custom_llm_provider": "bedrock"}, + ) + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + + +# --------------------------------------------------------------------------- # +# Cross-cutting: enrichment route_type and failure-hook on provider error. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__uses_acancel_batch_route_type(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.pre_call.call_args.kwargs["route_type"] == "acancel_batch" + + +@pytest.mark.asyncio +async def test_cancel__exception_calls_failure_hook(cancel_harness): + cancel_harness.litellm_acancel.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_cancel(cancel_harness, "batch-raw-xyz") + + cancel_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + cancel_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# Router-required 500 guards (one per endpoint branch that calls the router). +# These pin the defensive checks that fire when llm_router is unset. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__loadbalancing_no_router_500(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object( + litellm, "enable_loadbalancing_on_batch_endpoints", True + ), patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "500" + harness.router_acreate.assert_not_called() + harness.litellm_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__unified_no_router_500(harness): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"] + ), patch.object( + proxy_server, "llm_router", None + ): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "500" + + +@pytest.mark.asyncio +async def test_retrieve__unified_no_router_500(retrieve_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ), patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_retrieve(retrieve_harness, "batch-unified-blob") + + assert exc.value.code == "500" + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.litellm_aretrieve.assert_not_called() From 8e30cfbeb189b8655886c6ad486eecdade15c4f1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 29 Jun 2026 09:32:39 +0530 Subject: [PATCH 11/37] feat(a2a): support a2a-sdk 1.x proxy routing for 0.3 and 1.0 agents (#30950) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(a2a): support a2a-sdk 1.x proxy routing for 0.3 and 1.0 agents Bump a2a-sdk to 1.x and wire send/stream through compat conversions so the proxy accepts A2A 1.0 JSON-RPC while preserving 0.3 wire clients. Co-authored-by: Cursor * Add user controlled protocol version in agents * Fix exeception mapping * Fix a2a base url * Add e2e test for a2a * Fix lint * Fix lint * fix(a2a): harden card version detection and header isolation coverage Use protocolVersion when inferring agent card wire format, assert distinct httpx cache keys in the header-isolation test, and suppress targeted basedpyright errors for optional SDK imports. Co-authored-by: Cursor * fix(a2a): suppress reportArgumentType for SDK compat types and fix streaming trace ID - Add pyright: ignore[reportArgumentType] to SendMessageSuccessResponse id= and result= args in _send_message, and SendStreamingMessageResponse root= in _stream_messages, where a2a-sdk compat types diverge from basedpyright's inferred signature, reducing the reportArgumentType count back within budget. - Fix streaming trace ID in astream_a2a_message to use str(request.id) when available instead of always generating a new uuid4(), restoring JSON-RPC request-ID correlation for observability. Co-Authored-By: Claude Sonnet 4.6 * style(a2a): expand SendStreamingMessageResponse for black formatting Move pyright: ignore comment to the root= argument line so Black accepts the expanded multi-line form. Co-Authored-By: Claude Sonnet 4.6 * fix(a2a): fix 2 reportArgumentType errors without suppression - main.py: narrow logging_obj from object|None to Optional[Logging] via isinstance check before A2AStreamingIterator call, fixing the "Logging | object" argument type mismatch at line 699. - a2a_endpoints.py: extract response_dict with explicit isinstance(dict) guard before passing to normalize_jsonrpc_response, fixing the "LLMResponseTypes | dict[str, Any]" type mismatch at line 835. - Remove spurious pyright: ignore comments added in previous commits that were not suppressing the actual errors. Co-Authored-By: Claude Sonnet 4.6 * fix(a2a): rewrite upstream URL for 1.0 agent cards in getAuthenticatedExtendedCard 1.0 upstream agent cards store the endpoint URL in supportedInterfaces[0].url rather than a top-level url field. The previous guard only rewrote url when it existed at the top level, so after normalize_agent_card lowered a 1.0 card to 0.3 the upstream internal address leaked into the url field of the 0.3 response. Fix: rewrite both url and supportedInterfaces[0].url to the proxy address before calling normalize_agent_card, ensuring the upstream address is never visible to downstream clients regardless of the upstream card's wire format. Co-Authored-By: Claude Sonnet 4.6 * fix: extend _served_version to all PascalCase methods; add direct httpx-client isolation proof - _served_version now checks `_PASCAL_TO_WIRE` membership instead of two hardcoded names, so GetTask/CancelTask/etc. are promoted to 1.0 wire format alongside SendMessage — prevents mixed wire formats mid-session - test_create_a2a_client_uses_fresh_httpx_client now asserts a2a_client_a._litellm_httpx_client is not a2a_client_b._litellm_httpx_client (direct proof that header bleed cannot occur), in addition to the cache-key inequality check Co-Authored-By: Claude Sonnet 4.6 * fix: id:0 silently dropped in version_convert; explicit continue in stream retry - version_convert.py: replace `request_id or ""` with `str(request_id) if request_id is not None else ""` in both _send_result_to and _stream_result_to; id=0 is valid JSON-RPC and must not be coerced to "" which breaks response correlation - main.py: add explicit `continue` after the A2ALocalhostURLError retry in _execute_a2a_stream_with_retry so the control flow (retry → next iteration → stream_succeeded guard) is unambiguous Co-Authored-By: Claude Sonnet 4.6 * fix: preserve a2a retry and discovery card urls * Fix black * Fix test * fix(a2a): avoid KeyError in discovery log after 0.3→1.0 card normalization When a 0.3-style agent card is normalized to 1.0, the top-level url key is replaced by supportedInterfaces; log the already-computed proxy_url instead. Co-authored-by: Cursor * fix(a2a): preserve taskId when lowering push notification config set params Flatten 1.x create envelope fields before parsing into TaskPushNotificationConfig so 1.0 clients forwarding to 0.3 upstream keep taskId and config. Co-authored-by: Cursor * fix(a2a): ignore unknown fields in message/send proto fallback ParseDict in _build_message_send_params now matches other inbound paths so 1.0 clients with extra proto fields are not rejected with -32602. Co-authored-by: Cursor * fix(a2a): normalize tasks/list params and response across protocol versions Convert list task entries on the response path and lower ListTasksRequest params including status filters when forwarding 1.0 clients to 0.3 upstream. Co-authored-by: Cursor * fix(a2a): avoid reportArgumentType in _lower_list_tasks_params; use local var instead of _parse return Co-authored-by: Cursor * refactor(a2a): drop private SDK symbol in tasks/list status lowering _lower_list_tasks_params imported _CORE_TO_COMPAT_TASK_STATE, a private a2a-sdk symbol that could disappear on a patch release and silently break status-filter lowering. Derive the 0.3 wire string from the public protobuf enum name instead (TASK_STATE_ maps to the 0.3 value once the prefix is dropped and underscores become dashes) and validate the result against the 0.3 TaskState enum's own values via a fully-typed pure helper. Behavior is unchanged for every state; unspecified or unrecognized states still drop the filter. Adds parametrized regression tests covering dashed wire values (input-required, auth-required) and the unspecified drop. * fix(a2a): drop redundant push-notification envelope key; unify MessageToDict import _flatten_create_push_notification_params used `config or pushNotificationConfig`, which short-circuits so a co-present pushNotificationConfig key was never popped and leaked into the flattened params. Pop both keys unconditionally and prefer config when present. Adds a regression test on the helper that fails on the old leak. Also import MessageToDict from a2a.compat.v0_3.conversions in _lower_list_tasks_params to match every other conversion helper in the module instead of pulling it straight from google.protobuf.json_format. * fix(a2a): reject invalid message/stream params early with -32602 _handle_stream_message built MessageSendParams lazily inside the stream_response() generator, so malformed 1.0 params surfaced as a generic -32603 after the 200 status line was already committed. The non-streaming path validates up front and returns -32602 (Invalid params). Validate eagerly before returning the StreamingResponse and emit -32602 on failure so both paths reject malformed params identically. Adds a regression test asserting the streamed error code is -32602. * fix(a2a): raise clear error when non-streaming send ends on an update event _send_message fed the SDK iterator's last event straight into SendMessageSuccessResponse, whose result only accepts Message or Task. A non-standard upstream whose final event is a TaskStatusUpdateEvent or TaskArtifactUpdateEvent made the response construction raise an opaque pydantic ValidationError. Guard the converted result and raise a clear RuntimeError instead, consistent with the no-response guard above it. Adds regression tests for the Message happy path and the update-event rejection via an injected fake client. * test(a2a): lock in clean merged agent-card URL without PROXY_BASE_URL Regression coverage proving _build_merged_agent_card produces no double slash in supportedInterfaces[0].url when PROXY_BASE_URL is unset and request.base_url carries a trailing slash. get_custom_url routes through join_paths, which rstrips the base, so the f-string join stays clean. * style(a2a): modernize type annotations to satisfy strict ruff budget After merging the black->ruff-format migration from base, the A2A files owned by this PR still used Optional[X]/quoted annotations that pushed UP037/UP045 over their lowered ceilings. Convert to X | None, drop the now-unnecessary quoted local annotation in _send_message, and remove the imports left unused by the rewrite. Type semantics are unchanged. * style(a2a): type a2a_endpoints dict params as dict[str, Any] The merge with the formatter-migration baseline tightened the reportUnknownArgumentType ceiling; bare dict annotations made every value Unknown and pushed the codebase total over cap. Annotate the JSON-RPC params, body, metadata, and litellm_params dicts as dict[str, Any] so their values are typed, dropping the unknown-argument count back under the ceiling. No behavior change. * fix(a2a): guard localhost retry against a missing agent card handle_a2a_localhost_retry rewrote the card URL and called create_client with whatever agent_card it received. The caller resolves the card from the SDK client (Optional), so a None card reached set_agent_card_url and create_client, surfacing an opaque SDK error instead of a clear one. Add an early RuntimeError guard mirroring the httpx-client check, drop the now always-true card None-check on the stash line, and cover it with a regression test. * style(a2a): disable reportUnknownArgumentType in a2a-sdk boundary modules The lint env type-checks without the optional a2a-sdk/protobuf installed, so every call into the protobuf-generated compat conversions counts as an Unknown-typed argument and the new A2A code pushed the codebase reportUnknownArgumentType total over its ceiling. These three modules are the A2A SDK boundary; turn the rule off file-wide with a documented reason instead of scattering dozens of per-line ignores across every SDK call. * fix(a2a): tolerate unknown fields when lowering 1.0->0.3; align streaming trace id Two issues greptile flagged: version_convert: the 1.0->0.3 lowering paths (_send_result_to, _task_to, _stream_result_to) called ParseDict without ignore_unknown_fields=True, so a 1.0 upstream response carrying vendor extensions raised and best-effort fell back to passing the un-lowered 1.0 shape to a 0.3 client. Set the flag to match the agent-card path and every inbound path; unknown fields are now dropped and the result is correctly lowered. main.py: asend_message_streaming derived X-LiteLLM-Trace-Id from the JSON-RPC request id, unlike asend_message which uses the logging object's litellm_trace_id. Prefer the logging trace id (then request id, then a uuid) so streamed and non-streamed calls correlate under the same trace. Adds regression tests for both, including the stream-event lowering path. * style(a2a): apply ruff format to a2a protocol and proxy modules Co-authored-by: Cursor --------- Co-authored-by: Cursor Co-authored-by: Claude Sonnet 4.6 Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- README.md | 38 +- litellm/a2a_protocol/card_resolver.py | 37 +- .../a2a_protocol/exception_mapping_utils.py | 50 +- .../litellm_completion_bridge/README.md | 2 + litellm/a2a_protocol/main.py | 390 +++++++++------ litellm/proxy/a2a/agent_card.py | 28 +- litellm/proxy/a2a/version_convert.py | 449 ++++++++++++++++++ .../proxy/agent_endpoints/a2a_endpoints.py | 227 +++++++-- litellm/proxy/agent_endpoints/endpoints.py | 46 +- .../public_endpoints/public_endpoints.py | 4 +- pyproject.toml | 9 +- tests/agent_tests/test_a2a_agent.py | 23 +- .../test_a2a_exception_mapping_utils.py | 153 ++++++ .../a2a_protocol/test_card_resolver.py | 28 +- .../a2a_protocol/test_cost_calculator.py | 267 ++++++----- tests/test_litellm/a2a_protocol/test_main.py | 106 +++++ .../test_litellm/proxy/a2a/test_agent_card.py | 20 +- .../proxy/a2a/test_version_convert.py | 315 ++++++++++++ .../agent_endpoints/test_a2a_endpoints.py | 387 +++++++++++++++ .../agent_endpoints/test_a2a_version_e2e.py | 326 +++++++++++++ .../test_agent_header_isolation.py | 137 +++--- .../proxy/agent_endpoints/test_endpoints.py | 29 ++ .../src/components/agents/agent_config.ts | 12 +- .../components/agents/agent_form_fields.tsx | 9 + uv.lock | 58 ++- 25 files changed, 2701 insertions(+), 449 deletions(-) create mode 100644 litellm/proxy/a2a/version_convert.py create mode 100644 tests/test_litellm/a2a_protocol/test_a2a_exception_mapping_utils.py create mode 100644 tests/test_litellm/a2a_protocol/test_main.py create mode 100644 tests/test_litellm/proxy/a2a/test_version_convert.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py diff --git a/README.md b/README.md index 3d0f7282d7c..90d3e944fcc 100644 --- a/README.md +++ b/README.md @@ -156,35 +156,41 @@ response = await client.send_message(request) ### AI Gateway (Proxy Server) -**Step 1.** [Add your Agent to the AI Gateway](https://docs.litellm.ai/docs/a2a#adding-your-agent) +**Step 1.** [Add your Agent to the AI Gateway](https://docs.litellm.ai/docs/a2a#adding-your-agent) — set `protocolVersion` to `1.0` or `0.3` per agent -**Step 2.** Call Agent via A2A SDK +**Step 2.** Call Agent via A2A SDK (requires `a2a-sdk>=1.1.0`) ```python -from a2a.client import A2ACardResolver, A2AClient -from a2a.types import MessageSendParams, SendMessageRequest -from uuid import uuid4 import httpx +from a2a.client import A2ACardResolver, ClientConfig, ClientFactory +from a2a.types import Message, Part, Role, SendMessageRequest +from a2a.utils.constants import TransportProtocol +from uuid import uuid4 base_url = "http://localhost:4000/a2a/my-agent" # LiteLLM proxy + agent name headers = {"Authorization": "Bearer sk-1234"} # LiteLLM Virtual Key -async with httpx.AsyncClient(headers=headers) as httpx_client: - resolver = A2ACardResolver(httpx_client=httpx_client, base_url=base_url) +async with httpx.AsyncClient(headers=headers, timeout=60.0) as http_client: + resolver = A2ACardResolver(httpx_client=http_client, base_url=base_url) agent_card = await resolver.get_agent_card() - client = A2AClient(httpx_client=httpx_client, agent_card=agent_card) + config = ClientConfig( + httpx_client=http_client, + streaming=False, + supported_protocol_bindings=[TransportProtocol.JSONRPC, TransportProtocol.HTTP_JSON], + ) + client = ClientFactory(config).create(agent_card) request = SendMessageRequest( - id=str(uuid4()), - params=MessageSendParams( - message={ - "role": "user", - "parts": [{"kind": "text", "text": "Hello!"}], - "messageId": uuid4().hex, - } + message=Message( + message_id=uuid4().hex, + role=Role.ROLE_USER, + parts=[Part(text="Hello!")], ) ) - response = await client.send_message(request) + async for event in client.send_message(request): + populated = event.ListFields() + if populated and populated[0][0].name in ("message", "msg"): + print("".join(getattr(p, "text", "") or "" for p in populated[0][1].parts)) ``` [**Docs: A2A Agent Gateway**](https://docs.litellm.ai/docs/a2a) diff --git a/litellm/a2a_protocol/card_resolver.py b/litellm/a2a_protocol/card_resolver.py index 1955b5268e1..412c7a0897d 100644 --- a/litellm/a2a_protocol/card_resolver.py +++ b/litellm/a2a_protocol/card_resolver.py @@ -4,7 +4,7 @@ Custom A2A Card Resolver for LiteLLM. Extends the A2A SDK's card resolver to support multiple well-known paths. """ -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any, Dict from litellm._logging import verbose_logger from litellm.constants import LOCALHOST_URL_PATTERNS @@ -27,7 +27,7 @@ except ImportError: pass -def is_localhost_or_internal_url(url: Optional[str]) -> bool: +def is_localhost_or_internal_url(url: str | None) -> bool: """ Check if a URL is a localhost or internal URL. @@ -48,6 +48,29 @@ def is_localhost_or_internal_url(url: Optional[str]) -> bool: return any(pattern in url_lower for pattern in LOCALHOST_URL_PATTERNS) +def get_agent_card_url(agent_card: "AgentCard") -> str | None: + """Return the agent endpoint URL from the resolved SDK card.""" + url = getattr(agent_card, "url", None) + if url: + return url + + interfaces = getattr(agent_card, "supported_interfaces", None) + if interfaces: + return getattr(interfaces[0], "url", None) + return None + + +def set_agent_card_url(agent_card: "AgentCard", url: str) -> None: + """Set the agent endpoint URL on the resolved SDK card.""" + normalized = url.rstrip("/") + "/" + if hasattr(agent_card, "url"): + agent_card.url = normalized + + interfaces = getattr(agent_card, "supported_interfaces", None) + if interfaces: + interfaces[0].url = normalized + + def fix_agent_card_url(agent_card: "AgentCard", base_url: str) -> "AgentCard": """ Fix the agent card URL if it contains a localhost/internal address. @@ -70,6 +93,12 @@ def fix_agent_card_url(agent_card: "AgentCard", base_url: str) -> "AgentCard": fixed_url = base_url.rstrip("/") + "/" agent_card.url = fixed_url + interfaces = getattr(agent_card, "supported_interfaces", None) + if interfaces: + interface_url = getattr(interfaces[0], "url", None) + if interface_url and is_localhost_or_internal_url(interface_url): + interfaces[0].url = base_url.rstrip("/") + "/" + return agent_card @@ -84,8 +113,8 @@ class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc] async def get_agent_card( self, - relative_card_path: Optional[str] = None, - http_kwargs: Optional[Dict[str, Any]] = None, + relative_card_path: str | None = None, + http_kwargs: Dict[str, Any] | None = None, ) -> "AgentCard": """ Fetch the agent card, trying multiple well-known paths. diff --git a/litellm/a2a_protocol/exception_mapping_utils.py b/litellm/a2a_protocol/exception_mapping_utils.py index 99706e15cee..89b831351ab 100644 --- a/litellm/a2a_protocol/exception_mapping_utils.py +++ b/litellm/a2a_protocol/exception_mapping_utils.py @@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_logger from litellm.a2a_protocol.card_resolver import ( - fix_agent_card_url, is_localhost_or_internal_url, + set_agent_card_url, ) from litellm.a2a_protocol.exceptions import ( A2AAgentCardError, @@ -20,17 +20,18 @@ from litellm.a2a_protocol.exceptions import ( from litellm.constants import CONNECTION_ERROR_PATTERNS if TYPE_CHECKING: - from a2a.client import A2AClient as A2AClientType + from a2a.client import Client as A2AClientType -# Runtime import -A2A_SDK_AVAILABLE = False try: - from a2a.client import A2AClient as _A2AClient # type: ignore[no-redef] + from a2a.client import Client, ClientConfig, create_client A2A_SDK_AVAILABLE = True except ImportError: - _A2AClient = None # type: ignore[assignment, misc] + A2A_SDK_AVAILABLE = False + Client = None # type: ignore[misc, assignment] + ClientConfig = None # type: ignore[misc, assignment] + create_client = None # type: ignore[misc, assignment] class A2AExceptionCheckers: @@ -156,7 +157,7 @@ def map_a2a_exception( ) -def handle_a2a_localhost_retry( +async def handle_a2a_localhost_retry( error: A2ALocalhostURLError, agent_card: Any, a2a_client: "A2AClientType", @@ -180,8 +181,14 @@ def handle_a2a_localhost_retry( Raises: ImportError: If the A2A SDK is not installed """ - if not A2A_SDK_AVAILABLE or _A2AClient is None: - raise ImportError("A2A SDK is required for localhost retry handling. Install it with: pip install a2a") + if not A2A_SDK_AVAILABLE: + raise ImportError("A2A SDK is required for localhost retry handling. Install it with: pip install a2a-sdk") + + if agent_card is None: + raise RuntimeError( + "Cannot retry A2A localhost URL fix: no agent card is available to " + "rewrite, so the upstream URL cannot be corrected." + ) request_type = "streaming " if is_streaming else "" verbose_logger.warning( @@ -191,10 +198,25 @@ def handle_a2a_localhost_retry( ) # Fix the agent card URL - fix_agent_card_url(agent_card, error.base_url) + set_agent_card_url(agent_card, error.base_url) - # Create a new client with the fixed agent card (transport caches URL) - return _A2AClient( - httpx_client=a2a_client._transport.httpx_client, # type: ignore[union-attr] - agent_card=agent_card, + # Reuse the httpx client LiteLLM attached at creation. It carries this agent's + # trace-id and auth headers, so a fresh client would drop them. Only clients built + # by ``create_a2a_client`` have it; an externally-supplied client cannot be retried. + httpx_client = getattr(a2a_client, "_litellm_httpx_client", None) + if httpx_client is None: + raise RuntimeError( + "Cannot retry A2A localhost URL fix: the client was not created by " + "create_a2a_client, so no LiteLLM httpx client is attached." + ) + + new_client = await create_client( # pyright: ignore[reportOptionalCall] + agent_card, + client_config=ClientConfig( # pyright: ignore[reportOptionalCall] + httpx_client=httpx_client, + streaming=is_streaming, + ), ) + new_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined] + new_client._litellm_agent_card = agent_card # type: ignore[attr-defined] + return new_client diff --git a/litellm/a2a_protocol/litellm_completion_bridge/README.md b/litellm/a2a_protocol/litellm_completion_bridge/README.md index a809e9bf55e..3359e75f6df 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/README.md +++ b/litellm/a2a_protocol/litellm_completion_bridge/README.md @@ -67,6 +67,8 @@ When an A2A request hits `/a2a/{agent_id}/message/send`, the bridge: 3. Calls `litellm.acompletion(model="langgraph/agent", api_base="http://localhost:2024")` 4. Transforms response → A2A format +The proxy then normalizes the client-facing response to the agent's pinned `protocolVersion` (`0.3` or `1.0`). No extra provider config is required for completion-bridge agents — pin `protocolVersion` only if your client expects a specific wire format. + ## Classes - `A2ACompletionBridgeTransformation` - Static methods for message format conversion diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 6694c5c4af3..37bf7c34f02 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -1,3 +1,8 @@ +# pyright: reportUnknownArgumentType=false +# a2a-sdk (and its protobuf-generated compat conversions) ships no usable types for +# the call surface used here, so SDK calls take Unknown-typed arguments. This module +# is dedicated to the A2A SDK boundary; the rule is off file-wide instead of +# scattering per-line ignores across every SDK call. """ LiteLLM A2A SDK functions. @@ -7,7 +12,16 @@ Provides standalone functions with @client decorator for LiteLLM logging integra import asyncio import datetime import uuid -from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union +from typing import ( + TYPE_CHECKING, + Any, + AsyncIterator, + Coroutine, + Dict, + Optional, + Union, + cast, +) import litellm from litellm._logging import verbose_logger, verbose_proxy_logger @@ -23,23 +37,45 @@ from litellm.types.agents import LiteLLMSendMessageResponse from litellm.utils import client if TYPE_CHECKING: - from a2a.client import A2AClient as A2AClientType - from a2a.types import AgentCard, SendMessageRequest, SendStreamingMessageRequest + from a2a.client import Client as A2AClientType + from a2a.compat.v0_3.types import ( + AgentCard, + Message, + SendMessageRequest, + SendMessageResponse, + SendStreamingMessageRequest, + SendStreamingMessageResponse, + Task, + ) -# Runtime imports with availability check +# Runtime imports — requires a2a-sdk>=1.1.0 A2A_SDK_AVAILABLE = False -A2ACardResolver: Any = None -_A2AClient: Any = None +_a2a_conversions: Any = None try: - from a2a.client import A2AClient as _A2AClient # type: ignore[no-redef] + from a2a.client import Client, ClientConfig, create_client + from a2a.compat.v0_3 import conversions as _a2a_conversions + from a2a.compat.v0_3.types import ( + Message, + SendMessageRequest, + SendMessageResponse, + SendMessageSuccessResponse, + SendStreamingMessageRequest, + SendStreamingMessageResponse, + Task, + ) A2A_SDK_AVAILABLE = True except ImportError: - pass + Client = None # type: ignore[misc, assignment] + ClientConfig = None # type: ignore[misc, assignment] + create_client = None # type: ignore[misc, assignment] # Import our custom card resolver that supports multiple well-known paths -from litellm.a2a_protocol.card_resolver import LiteLLMA2ACardResolver +from litellm.a2a_protocol.card_resolver import ( + LiteLLMA2ACardResolver, + get_agent_card_url, +) from litellm.a2a_protocol.exception_mapping_utils import ( handle_a2a_localhost_retry, map_a2a_exception, @@ -75,7 +111,7 @@ def _set_usage_on_logging_obj( def _set_agent_id_on_logging_obj( kwargs: Dict[str, Any], - agent_id: Optional[str], + agent_id: str | None, ) -> None: """ Set agent_id on litellm_logging_obj for SpendLogs tracking. @@ -102,10 +138,7 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str: """ agent_name = "unknown" - # Try to get agent card from our stored attribute first, then fallback to SDK attribute - agent_card = getattr(a2a_client, "_litellm_agent_card", None) - if agent_card is None: - agent_card = getattr(a2a_client, "agent_card", None) + agent_card = _get_a2a_client_agent_card(a2a_client) if agent_card is not None: agent_name = getattr(agent_card, "name", "unknown") or "unknown" @@ -125,12 +158,22 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str: return agent_name +def _get_a2a_client_agent_card(a2a_client: Any) -> Optional["AgentCard"]: + agent_card = cast(Optional["AgentCard"], getattr(a2a_client, "_litellm_agent_card", None)) + if agent_card is not None: + return agent_card + agent_card = cast(Optional["AgentCard"], getattr(a2a_client, "agent_card", None)) + if agent_card is not None: + return agent_card + return cast(Optional["AgentCard"], getattr(a2a_client, "_card", None)) + + async def _send_message_via_completion_bridge( request: "SendMessageRequest", custom_llm_provider: str, - api_base: Optional[str], + api_base: str | None, litellm_params: Dict[str, Any], - agent_extra_headers: Optional[Dict[str, str]] = None, + agent_extra_headers: Dict[str, str] | None = None, ) -> LiteLLMSendMessageResponse: """ Route a send_message through the LiteLLM completion bridge (e.g. LangGraph, Bedrock AgentCore). @@ -156,39 +199,71 @@ async def _send_message_via_completion_bridge( return LiteLLMSendMessageResponse.from_dict(response_dict, request_id=str(request.id)) +async def _send_message(a2a_client: "A2AClientType", request: "SendMessageRequest") -> "SendMessageResponse": + """Send a non-streaming message via a2a-sdk 1.x and return JSON-RPC response.""" + if _a2a_conversions is None: + raise ImportError( + "The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk" + ) + + pb_request = _a2a_conversions.to_core_send_message_request(request) + last_event = None + async for event in a2a_client.send_message(pb_request): + last_event = event + if last_event is None: + raise RuntimeError("A2A send_message failed: no response received from agent.") + + stream_compat = _a2a_conversions.to_compat_stream_response( + last_event, + request_id=request.id, + ) + result = stream_compat.result + if not isinstance(result, (Message, Task)): + raise RuntimeError( + "A2A send_message failed: non-streaming message/send expects the " + "agent's final event to be a Message or Task result." + ) + return SendMessageResponse( + root=SendMessageSuccessResponse( + id=request.id, + result=result, + ) + ) + + async def _execute_a2a_send_with_retry( - a2a_client: Any, - request: Any, - agent_card: Any, - card_url: Optional[str], - api_base: Optional[str], - agent_name: Optional[str], -) -> Any: + a2a_client: "A2AClientType", + request: "SendMessageRequest", + agent_card: Optional["AgentCard"], + card_url: str | None, + api_base: str | None, + agent_name: str | None, +) -> "SendMessageResponse": """Send an A2A message with retry logic for localhost URL errors.""" a2a_response = None for _ in range(2): # max 2 attempts: original + 1 retry try: - a2a_response = await a2a_client.send_message(request) + a2a_response = await _send_message(a2a_client, request) break # success, exit retry loop except A2ALocalhostURLError as e: - a2a_client = handle_a2a_localhost_retry( + a2a_client = await handle_a2a_localhost_retry( error=e, agent_card=agent_card, a2a_client=a2a_client, is_streaming=False, ) - card_url = agent_card.url if agent_card else None + card_url = get_agent_card_url(agent_card) if agent_card else None except Exception as e: try: map_a2a_exception(e, card_url, api_base, model=agent_name) except A2ALocalhostURLError as localhost_err: - a2a_client = handle_a2a_localhost_retry( + a2a_client = await handle_a2a_localhost_retry( error=localhost_err, agent_card=agent_card, a2a_client=a2a_client, is_streaming=False, ) - card_url = agent_card.url if agent_card else None + card_url = get_agent_card_url(agent_card) if agent_card else None continue except Exception: raise @@ -197,14 +272,80 @@ async def _execute_a2a_send_with_retry( return a2a_response +async def _stream_messages( + a2a_client: "A2AClientType", request: "SendStreamingMessageRequest" +) -> AsyncIterator["SendStreamingMessageResponse"]: + """Stream message events via a2a-sdk 1.x and yield JSON-RPC chunks.""" + if _a2a_conversions is None: + raise ImportError( + "The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk" + ) + + pb_request = _a2a_conversions.to_core_send_message_request(request) + async for event in a2a_client.send_message(pb_request): + compat_chunk = _a2a_conversions.to_compat_stream_response( + event, + request_id=request.id, + ) + yield SendStreamingMessageResponse(root=compat_chunk) + + +async def _execute_a2a_stream_with_retry( + a2a_client: "A2AClientType", + request: "SendStreamingMessageRequest", + agent_card: Optional["AgentCard"], + card_url: str | None, + api_base: str | None, + agent_name: str | None, +) -> AsyncIterator["SendStreamingMessageResponse"]: + """Stream an A2A message with retry logic for localhost URL errors.""" + response_started = False + stream_succeeded = False + for _ in range(2): # max 2 attempts: original + 1 retry + try: + async for chunk in _stream_messages(a2a_client, request): + response_started = True + yield chunk + stream_succeeded = True + return + except A2ALocalhostURLError as e: + if response_started: + raise + a2a_client = await handle_a2a_localhost_retry( + error=e, + agent_card=agent_card, + a2a_client=a2a_client, + is_streaming=True, + ) + card_url = get_agent_card_url(agent_card) if agent_card else None + continue + except Exception as e: + if response_started: + raise + try: + map_a2a_exception(e, card_url, api_base, model=agent_name) + except A2ALocalhostURLError as localhost_err: + a2a_client = await handle_a2a_localhost_retry( + error=localhost_err, + agent_card=agent_card, + a2a_client=a2a_client, + is_streaming=True, + ) + card_url = get_agent_card_url(agent_card) if agent_card else None + continue + raise + if not stream_succeeded: + raise RuntimeError("A2A send_message_streaming failed: no response received after retry attempts.") + + @client async def asend_message( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendMessageRequest"] = None, - api_base: Optional[str] = None, - litellm_params: Optional[Dict[str, Any]] = None, - agent_id: Optional[str] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, + api_base: str | None = None, + litellm_params: Dict[str, Any] | None = None, + agent_id: str | None = None, + agent_extra_headers: Dict[str, str] | None = None, **kwargs: Any, ) -> LiteLLMSendMessageResponse: """ @@ -301,8 +442,8 @@ async def asend_message( verbose_logger.info(f"A2A send_message request_id={request.id}, agent={agent_name}") # Get agent card URL for localhost retry logic - agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None) - card_url = getattr(agent_card, "url", None) if agent_card else None + agent_card = _get_a2a_client_agent_card(a2a_client) + card_url = get_agent_card_url(agent_card) if agent_card else None a2a_response = await _execute_a2a_send_with_retry( a2a_client=a2a_client, @@ -375,10 +516,10 @@ def send_message( def _build_streaming_logging_obj( request: "SendStreamingMessageRequest", agent_name: str, - agent_id: Optional[str], - litellm_params: Optional[Dict[str, Any]], - metadata: Optional[Dict[str, Any]], - proxy_server_request: Optional[Dict[str, Any]], + agent_id: str | None, + litellm_params: Dict[str, Any] | None, + metadata: Dict[str, Any] | None, + proxy_server_request: Dict[str, Any] | None, ) -> Logging: """Build logging object for streaming A2A requests.""" start_time = datetime.datetime.now() @@ -417,12 +558,13 @@ def _build_streaming_logging_obj( async def asend_message_streaming( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendStreamingMessageRequest"] = None, - api_base: Optional[str] = None, - litellm_params: Optional[Dict[str, Any]] = None, - agent_id: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - proxy_server_request: Optional[Dict[str, Any]] = None, - agent_extra_headers: Optional[Dict[str, str]] = None, + api_base: str | None = None, + litellm_params: Dict[str, Any] | None = None, + agent_id: str | None = None, + metadata: Dict[str, Any] | None = None, + proxy_server_request: Dict[str, Any] | None = None, + agent_extra_headers: Dict[str, str] | None = None, + **kwargs: object, ) -> AsyncIterator[Any]: """ Async: Send a streaming message to an A2A agent. @@ -491,99 +633,72 @@ async def asend_message_streaming( yield chunk return - # Standard A2A client flow if request is None: raise ValueError("request is required") - # Create A2A client if not provided but api_base is available + _raw_logging_obj = kwargs.get("litellm_logging_obj") + logging_obj: Logging | None = _raw_logging_obj if isinstance(_raw_logging_obj, Logging) else None + if a2a_client is None: if api_base is None: raise ValueError("Either a2a_client or api_base is required for standard A2A flow") - # Mirror the non-streaming path: always include trace and agent-id headers - streaming_extra_headers: Dict[str, str] = { - "X-LiteLLM-Trace-Id": str(request.id), - } + logging_trace_id = getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None + trace_id = logging_trace_id or (str(request.id) if request.id else str(uuid.uuid4())) + extra_headers: dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id} if agent_id: - streaming_extra_headers["X-LiteLLM-Agent-Id"] = agent_id + extra_headers["X-LiteLLM-Agent-Id"] = agent_id if agent_extra_headers: - streaming_extra_headers.update(agent_extra_headers) - a2a_client = await create_a2a_client(base_url=api_base, extra_headers=streaming_extra_headers) - - # Type assertion: a2a_client is guaranteed to be non-None here - assert a2a_client is not None - - verbose_logger.info(f"A2A send_message_streaming request_id={request.id}") - - # Build logging object for streaming completion callbacks - agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None) - card_url = getattr(agent_card, "url", None) if agent_card else None - agent_name = getattr(agent_card, "name", "unknown") if agent_card else "unknown" - - logging_obj = _build_streaming_logging_obj( - request=request, - agent_name=agent_name, - agent_id=agent_id, - litellm_params=litellm_params, - metadata=metadata, - proxy_server_request=proxy_server_request, - ) - - # Retry loop: if connection fails due to localhost URL in agent card, retry with fixed URL - # Connection errors in streaming typically occur on first chunk iteration - first_chunk = True - for attempt in range(2): # max 2 attempts: original + 1 retry - stream = a2a_client.send_message_streaming(request) - iterator = A2AStreamingIterator( - stream=stream, - request=request, - logging_obj=logging_obj, - agent_name=agent_name, + extra_headers.update(agent_extra_headers) + a2a_client = await create_a2a_client( + base_url=api_base, + extra_headers=extra_headers, + streaming=True, ) - try: - first_chunk = True - async for chunk in iterator: - if first_chunk: - first_chunk = False # connection succeeded - yield chunk - return # stream completed successfully - except A2ALocalhostURLError as e: - # Only retry on first chunk, not mid-stream - if first_chunk and attempt == 0: - a2a_client = handle_a2a_localhost_retry( - error=e, - agent_card=agent_card, - a2a_client=a2a_client, - is_streaming=True, - ) - card_url = agent_card.url if agent_card else None - else: - raise - except Exception as e: - # Only map exception on first chunk - if first_chunk and attempt == 0: - try: - map_a2a_exception(e, card_url, api_base, model=agent_name) - except A2ALocalhostURLError as localhost_err: - # Localhost URL error - fix and retry - a2a_client = handle_a2a_localhost_retry( - error=localhost_err, - agent_card=agent_card, - a2a_client=a2a_client, - is_streaming=True, - ) - card_url = agent_card.url if agent_card else None - continue - except Exception: - # Re-raise the mapped exception - raise - raise + assert a2a_client is not None + + agent_name = _get_a2a_model_info(a2a_client, kwargs) + + if logging_obj is None: + logging_obj = _build_streaming_logging_obj( + request=request, + agent_name=agent_name, + agent_id=agent_id, + litellm_params=litellm_params, + metadata=metadata, + proxy_server_request=proxy_server_request, + ) + + verbose_logger.info(f"A2A send_message_streaming request_id={request.id}, agent={agent_name}") + + agent_card = _get_a2a_client_agent_card(a2a_client) + card_url = get_agent_card_url(agent_card) if agent_card else None + + stream = _execute_a2a_stream_with_retry( + a2a_client=a2a_client, + request=request, + agent_card=agent_card, + card_url=card_url, + api_base=api_base, + agent_name=agent_name, + ) + + _set_agent_id_on_logging_obj(kwargs=kwargs, agent_id=agent_id) + + async for chunk in A2AStreamingIterator( + stream=stream, + request=request, + logging_obj=logging_obj, + agent_name=agent_name, + ): + yield chunk async def create_a2a_client( base_url: str, timeout: float = DEFAULT_A2A_AGENT_TIMEOUT, - extra_headers: Optional[Dict[str, str]] = None, + extra_headers: Dict[str, str] | None = None, + streaming: bool = False, ) -> "A2AClientType": """ Create an A2A client for the given agent URL. @@ -640,23 +755,20 @@ async def create_a2a_client( httpx_client.headers.update(extra_headers) verbose_proxy_logger.debug(f"A2A client created with extra_headers={list(extra_headers.keys())}") - # Resolve agent card - resolver = A2ACardResolver( - httpx_client=httpx_client, - base_url=base_url, + a2a_client = await create_client( # pyright: ignore[reportOptionalCall] + base_url, + client_config=ClientConfig( # pyright: ignore[reportOptionalCall] + httpx_client=httpx_client, + streaming=streaming, + ), ) - agent_card = await resolver.get_agent_card() - - verbose_logger.debug(f"Resolved agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}") - - # Create A2A client - a2a_client = _A2AClient( - httpx_client=httpx_client, - agent_card=agent_card, - ) - - # Store agent_card on client for later retrieval (SDK doesn't expose it) - a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined] + # Stash LiteLLM-owned handles on the client so the localhost-retry path can reuse + # the configured httpx client (with this agent's trace-id/auth headers) without + # excavating a2a-sdk private internals. + a2a_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined] + agent_card = getattr(a2a_client, "_card", None) + if agent_card is not None: + a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined] verbose_logger.info(f"A2A client created for {base_url}") @@ -666,7 +778,7 @@ async def create_a2a_client( async def aget_agent_card( base_url: str, timeout: float = DEFAULT_A2A_AGENT_TIMEOUT, - extra_headers: Optional[Dict[str, str]] = None, + extra_headers: Dict[str, str] | None = None, ) -> "AgentCard": """ Fetch the agent card from an A2A agent. diff --git a/litellm/proxy/a2a/agent_card.py b/litellm/proxy/a2a/agent_card.py index 79638341ac1..e97ab4a01ae 100644 --- a/litellm/proxy/a2a/agent_card.py +++ b/litellm/proxy/a2a/agent_card.py @@ -8,11 +8,24 @@ and uses LiteLLM auth. """ from copy import deepcopy -from typing import Any, Dict, List, Mapping, Optional +from typing import Any, Dict, List, Mapping -# Protocol version LiteLLM speaks. Bump when the proxy's A2A surface changes. +# Protocol versions LiteLLM can serve to A2A clients. The admin pins one per agent; +# responses are normalized to it regardless of the upstream agent's own version. +SUPPORTED_A2A_PROTOCOL_VERSIONS = ("0.3", "1.0") + +# Default served version when the agent card does not pin one. LITELLM_A2A_PROTOCOL_VERSION = "1.0" + +def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str: + """Return the validated protocol version an agent card pins, else the default.""" + version = card.get("protocolVersion") if card else None + if version in SUPPORTED_A2A_PROTOCOL_VERSIONS: + return version + return LITELLM_A2A_PROTOCOL_VERSION + + # Security scheme exposed by the LiteLLM-fronted agent card. Always replaces # whatever upstream advertised — the client must authenticate to the proxy, # not the upstream agent. @@ -106,12 +119,12 @@ def _default_litellm_provider(proxy_base_url: str) -> Dict[str, str]: def merge_agent_card( - upstream_card: Optional[Mapping[str, Any]], + upstream_card: Mapping[str, Any] | None, *, proxy_url: str, proxy_base_url: str, - name: Optional[str] = None, - description: Optional[str] = None, + name: str | None = None, + description: str | None = None, ) -> Dict[str, Any]: """ Build the LiteLLM-fronted agent card. @@ -139,7 +152,8 @@ def merge_agent_card( # proxy requests. The public well-known endpoint rewrites this field # to the proxy URL before exposing the card to clients. - base["protocolVersion"] = LITELLM_A2A_PROTOCOL_VERSION + served_version = resolve_served_protocol_version(upstream_card) + base["protocolVersion"] = served_version if name: base["name"] = name @@ -165,7 +179,7 @@ def merge_agent_card( { "url": proxy_url, "protocolBinding": "JSONRPC", - "protocolVersion": LITELLM_A2A_PROTOCOL_VERSION, + "protocolVersion": served_version, } ] diff --git a/litellm/proxy/a2a/version_convert.py b/litellm/proxy/a2a/version_convert.py new file mode 100644 index 00000000000..e8f49e6f6a9 --- /dev/null +++ b/litellm/proxy/a2a/version_convert.py @@ -0,0 +1,449 @@ +# pyright: reportUnknownArgumentType=false +# a2a-sdk's compat conversions (pb2_v10, ParseDict, MessageToDict, to_compat_*) +# are protobuf-generated/untyped, so every conversion call here takes Unknown-typed +# arguments. This module is the A2A 0.3<->1.0 boundary; the rule is off file-wide +# rather than scattering per-line ignores across every SDK call. +""" +Normalize A2A JSON-RPC payloads to the protocol version LiteLLM serves for an agent. + +LiteLLM fronts upstream agents and lets an admin pin the protocol version it speaks +to clients (``0.3`` or ``1.0``) per agent. Upstream responses may arrive in either +wire shape, so every response, stream event, forwarded request and extended card is +converted to the served version here. Conversion is shape-detecting (we infer the +payload's current version rather than trusting a stored one) and best-effort: any +failure falls back to returning the input unchanged so a conversion bug can never +break an otherwise-valid response. + +The two wire shapes: + +- ``0.3``: JSON dump of the compat pydantic types, discriminated by a ``kind`` field + (``message`` / ``task`` / ``status-update`` / ``artifact-update``). A send result is + the bare object. +- ``1.0``: protobuf JSON (``MessageToDict``), a oneof envelope keyed by + ``message`` / ``task`` / ``statusUpdate`` / ``artifactUpdate`` with no ``kind``. A + ``Task`` result is a bare object without ``kind``. +""" + +from types import ModuleType +from typing import Callable, Literal, Union + +from pydantic import BaseModel + +from litellm._logging import verbose_proxy_logger + +A2AVersion = Literal["0.3", "1.0"] +RequestId = Union[str, int, None] +JsonDict = dict[str, object] + +_V1_SEND_ENVELOPE_KEYS = frozenset({"message", "task"}) +_V1_STREAM_ENVELOPE_KEYS = frozenset({"message", "task", "statusUpdate", "artifactUpdate"}) + + +def _dump_03(model: BaseModel) -> JsonDict: + """Dump a compat (0.3) pydantic model to its camelCase wire dict.""" + return model.model_dump(by_alias=True, exclude_none=True, mode="json") + + +def _best_effort(convert: Callable[[], JsonDict], fallback: JsonDict, *, label: str) -> JsonDict: + """Run a conversion, returning ``fallback`` unchanged if it raises.""" + try: + return convert() + except Exception as e: # noqa: BLE001 - best-effort passthrough + verbose_proxy_logger.debug("A2A %s conversion failed: %s", label, e) + return fallback + + +def normalize_jsonrpc_response(content: JsonDict, target: A2AVersion, *, method: str) -> JsonDict: + """Convert a JSON-RPC response's ``result`` to ``target``. + + Errors and non-dict results pass through untouched. + """ + if content.get("error") is not None: + return content + result = content.get("result") + if not isinstance(result, dict): + return content + + converted = _convert_result(result, target, method=method, request_id=_as_request_id(content.get("id"))) + if converted is result: + return content + return {**content, "result": converted} + + +def normalize_stream_event(event: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict: + """Convert a single streamed JSON-RPC event's ``result`` to ``target``.""" + if event.get("error") is not None: + return event + result = event.get("result") + if not isinstance(result, dict): + return event + + converted = _convert_stream_result(result, target, request_id=request_id) + if converted is result: + return event + return {**event, "result": converted} + + +def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: str) -> JsonDict: + """Down-convert forwarded request ``params`` from the served version to 0.3. + + Upstream agents in this proxy pivot on 0.3 wire format, so when LiteLLM serves + 1.0 the inbound params must be lowered before forwarding. A no-op when the served + version is already 0.3. + """ + if served == "0.3": + return params + return _best_effort( + lambda: _lower_request_params(params, method=method), + params, + label=f"request params ({method})", + ) + + +def _detect_card_version(card: JsonDict) -> A2AVersion: + """Infer the wire version of an agent card dict. + + ``protocolVersion`` is the authoritative indicator; fall back to presence of + ``supportedInterfaces`` (a 1.0-only field) only when the explicit field is absent. + Cards that set ``protocolVersion: "0.3"`` or carry neither signal are treated as 0.3. + """ + pv = card.get("protocolVersion") + if pv == "1.0": + return "1.0" + if pv == "0.3": + return "0.3" + # No protocolVersion field: use structural heuristic. + return "1.0" if "supportedInterfaces" in card else "0.3" + + +def normalize_agent_card(card: JsonDict, target: A2AVersion) -> JsonDict: + """Convert an extended agent card to ``target``. + + When lowering to 0.3, ``additionalInterfaces`` is stripped so the conversion never + re-exposes upstream backend URLs that the LiteLLM-fronting merge deliberately drops. + """ + if not isinstance(card, dict): + return card + + current = _detect_card_version(card) + if current == target and not (target == "0.3" and "supportedInterfaces" in card): + return card + return _best_effort(lambda: _convert_agent_card(card, target), card, label="agent card") + + +def _as_request_id(value: object) -> RequestId: + return value if isinstance(value, (str, int)) else None + + +def _convert_result( + result: JsonDict, + target: A2AVersion, + *, + method: str, + request_id: RequestId, +) -> JsonDict: + if method == "message/send": + return _convert_send_result(result, target, request_id=request_id) + if method in ("tasks/get", "tasks/cancel"): + return _convert_task(result, target) + if method == "tasks/list": + return _convert_list_tasks_result(result, target) + return result + + +def _detect_send_version(result: JsonDict) -> A2AVersion | None: + if "kind" in result: + return "0.3" + if result.keys() & _V1_SEND_ENVELOPE_KEYS: + return "1.0" + return None + + +def _convert_send_result(result: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict: + current = _detect_send_version(result) + if current is None or current == target: + return result + return _best_effort( + lambda: _send_result_to(result, target, request_id), + result, + label="send result", + ) + + +def _send_result_to(result: JsonDict, target: A2AVersion, request_id: RequestId) -> JsonDict: + from a2a.compat.v0_3.conversions import ( + MessageToDict, + ParseDict, + pb2_v10, + to_compat_send_message_response, + to_core_send_message_response, + types_v03, + ) + + if target == "1.0": + compat_result = _validate_message_or_task(result, types_v03) + response = types_v03.SendMessageResponse( + root=types_v03.SendMessageSuccessResponse( + id=str(request_id) if request_id is not None else "", + result=compat_result, # pyright: ignore[reportArgumentType] + ) + ) + return MessageToDict( + to_core_send_message_response(response), + preserving_proto_field_name=False, + ) + + pb = pb2_v10.SendMessageResponse() + ParseDict(result, pb, ignore_unknown_fields=True) + return _dump_03(to_compat_send_message_response(pb, request_id).root.result) + + +def _convert_task(result: JsonDict, target: A2AVersion) -> JsonDict: + current: A2AVersion = "0.3" if "kind" in result else "1.0" + if current == target: + return result + return _best_effort(lambda: _task_to(result, target), result, label="task") + + +def _detect_list_tasks_version(result: JsonDict) -> A2AVersion | None: + tasks = result.get("tasks") + if not isinstance(tasks, list) or not tasks: + return None + first = tasks[0] + if not isinstance(first, dict): + return None + return "0.3" if "kind" in first else "1.0" + + +def _convert_list_tasks_result(result: JsonDict, target: A2AVersion) -> JsonDict: + current = _detect_list_tasks_version(result) + if current is None or current == target: + return result + return _best_effort( + lambda: _list_tasks_result_to(result, target), + result, + label="list tasks result", + ) + + +def _list_tasks_result_to(result: JsonDict, target: A2AVersion) -> JsonDict: + tasks = result.get("tasks") + if not isinstance(tasks, list): + return result + return { + **result, + "tasks": [_task_to(item, target) if isinstance(item, dict) else item for item in tasks], + } + + +def _task_to(result: JsonDict, target: A2AVersion) -> JsonDict: + from a2a.compat.v0_3.conversions import ( + MessageToDict, + ParseDict, + pb2_v10, + to_compat_task, + to_core_task, + types_v03, + ) + + if target == "1.0": + core = to_core_task(types_v03.Task.model_validate(result)) + return MessageToDict(core, preserving_proto_field_name=False) + + pb = pb2_v10.Task() + ParseDict(result, pb, ignore_unknown_fields=True) + return _dump_03(to_compat_task(pb)) + + +def _detect_stream_version(result: JsonDict) -> A2AVersion | None: + if "kind" in result: + return "0.3" + if result.keys() & _V1_STREAM_ENVELOPE_KEYS: + return "1.0" + return None + + +def _convert_stream_result(result: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict: + current = _detect_stream_version(result) + if current is None or current == target: + return result + return _best_effort( + lambda: _stream_result_to(result, target, request_id), + result, + label="stream event", + ) + + +def _stream_result_to(result: JsonDict, target: A2AVersion, request_id: RequestId) -> JsonDict: + from a2a.compat.v0_3.conversions import ( + MessageToDict, + ParseDict, + pb2_v10, + to_compat_stream_response, + to_core_stream_response, + types_v03, + ) + + if target == "1.0": + event = _validate_stream_event(result, types_v03) + wrapper = types_v03.SendStreamingMessageSuccessResponse( + id=str(request_id) if request_id is not None else "", + result=event, # pyright: ignore[reportArgumentType] + ) + return MessageToDict(to_core_stream_response(wrapper), preserving_proto_field_name=False) + + pb = pb2_v10.StreamResponse() + ParseDict(result, pb, ignore_unknown_fields=True) + return _dump_03(to_compat_stream_response(pb, request_id).result) + + +def _convert_agent_card(card: JsonDict, target: A2AVersion) -> JsonDict: + from a2a.compat.v0_3.conversions import ( + MessageToDict, + ParseDict, + pb2_v10, + to_compat_agent_card, + to_core_agent_card, + types_v03, + ) + + if target == "0.3": + pb = pb2_v10.AgentCard() + ParseDict(card, pb, ignore_unknown_fields=True) + lowered = _dump_03(to_compat_agent_card(pb)) + lowered.pop("additionalInterfaces", None) + return lowered + + core = to_core_agent_card(types_v03.AgentCard.model_validate(card)) + return MessageToDict(core, preserving_proto_field_name=False) + + +def _validate_message_or_task(result: JsonDict, types_v03: ModuleType) -> BaseModel: + if result.get("kind") == "task": + return types_v03.Task.model_validate(result) + return types_v03.Message.model_validate(result) + + +def _validate_stream_event(result: JsonDict, types_v03: ModuleType) -> BaseModel: + kind = result.get("kind") + if kind == "task": + return types_v03.Task.model_validate(result) + if kind == "status-update": + return types_v03.TaskStatusUpdateEvent.model_validate(result) + if kind == "artifact-update": + return types_v03.TaskArtifactUpdateEvent.model_validate(result) + return types_v03.Message.model_validate(result) + + +def _lower_request_params(params: JsonDict, *, method: str) -> JsonDict: + if method == "tasks/list": + return _lower_list_tasks_params(params) + + from a2a.compat.v0_3.conversions import ( + ParseDict, + pb2_v10, + to_compat_cancel_task_request, + to_compat_create_task_push_notification_config_request, + to_compat_delete_task_push_notification_config_request, + to_compat_get_task_push_notification_config_request, + to_compat_get_task_request, + to_compat_list_task_push_notification_config_request, + to_compat_subscribe_to_task_request, + ) + + lowerings: dict[str, Callable[[JsonDict], BaseModel]] = { + "tasks/get": lambda p: to_compat_get_task_request(_parse(ParseDict, p, pb2_v10.GetTaskRequest()), "").params, + "tasks/cancel": lambda p: ( + to_compat_cancel_task_request(_parse(ParseDict, p, pb2_v10.CancelTaskRequest()), "").params + ), + "tasks/resubscribe": lambda p: ( + to_compat_subscribe_to_task_request(_parse(ParseDict, p, pb2_v10.SubscribeToTaskRequest()), "").params + ), + "tasks/pushNotificationConfig/set": lambda p: ( + to_compat_create_task_push_notification_config_request( + _parse( + ParseDict, + _flatten_create_push_notification_params(p), + pb2_v10.TaskPushNotificationConfig(), + ), + "", + ).params + ), + "tasks/pushNotificationConfig/get": lambda p: ( + to_compat_get_task_push_notification_config_request( + _parse(ParseDict, p, pb2_v10.GetTaskPushNotificationConfigRequest()), "" + ).params + ), + "tasks/pushNotificationConfig/list": lambda p: ( + to_compat_list_task_push_notification_config_request( + _parse(ParseDict, p, pb2_v10.ListTaskPushNotificationConfigsRequest()), "" + ).params + ), + "tasks/pushNotificationConfig/delete": lambda p: ( + to_compat_delete_task_push_notification_config_request( + _parse(ParseDict, p, pb2_v10.DeleteTaskPushNotificationConfigRequest()), "" + ).params + ), + } + lower = lowerings.get(method) + if lower is None: + return params + return _dump_03(lower(params)) + + +def _lower_list_tasks_params(params: JsonDict) -> JsonDict: + from a2a.compat.v0_3.conversions import ( + MessageToDict, + ParseDict, + pb2_v10, + types_v03, + ) + + proto = pb2_v10.ListTasksRequest() + _parse(ParseDict, params, proto) + lowered = MessageToDict(proto, preserving_proto_field_name=False) + status_name = str(pb2_v10.TaskState.Name(proto.status)) + valid_0_3_values = frozenset(str(member.value) for member in types_v03.TaskState) + compat_status = _proto_task_state_name_to_0_3(status_name, valid_0_3_values) + if compat_status is None: + lowered.pop("status", None) + else: + lowered["status"] = compat_status + return lowered + + +def _proto_task_state_name_to_0_3(name: str, valid_0_3_values: frozenset[str]) -> str | None: + """Map a 1.0 protobuf ``TaskState`` enum name to its 0.3 wire string. + + The ``TASK_STATE_`` enum names line up with the 0.3 wire values once the + prefix is dropped and underscores become dashes, so no private SDK mapping is + needed. The result is validated against the 0.3 enum's own values; an unspecified + or unrecognized state yields ``None`` so the status filter is dropped. + """ + base = name.removeprefix("TASK_STATE_") + if base == "UNSPECIFIED": + return None + candidate = base.lower().replace("_", "-") + return candidate if candidate in valid_0_3_values else None + + +def _flatten_create_push_notification_params(params: JsonDict) -> JsonDict: + """Merge 1.x create envelope fields (parent/configId/config) into flat pb fields.""" + flat = dict(params) + config = flat.pop("config", None) + push_config = flat.pop("pushNotificationConfig", None) + nested = config if config is not None else push_config + if not isinstance(nested, dict): + return params + parent = flat.pop("parent", None) + if isinstance(parent, str) and parent.startswith("tasks/") and "taskId" not in flat: + flat["taskId"] = parent.removeprefix("tasks/").split("/")[0] + if (config_id := flat.pop("configId", None)) and "id" not in nested: + nested["id"] = config_id + flat.update(nested) + return flat + + +def _parse(parse_dict: Callable[..., object], data: JsonDict, message: object) -> object: + parse_dict(data, message, ignore_unknown_fields=True) + return message diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 78be6231ae0..76423fbebde 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -1,3 +1,8 @@ +# pyright: reportUnknownArgumentType=false +# This module forwards JSON-RPC payloads through the untyped a2a-sdk compat +# conversions (pb2_v10/ParseDict/MessageToDict/to_compat_*), so SDK and decoded-JSON +# values flow in as Unknown. The rule is off file-wide rather than scattering per-line +# ignores across every SDK and JSON-RPC call. """ A2A Protocol endpoints for LiteLLM Proxy. @@ -6,26 +11,43 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM """ import json -from typing import Any, AsyncGenerator, Dict, List, Optional +from copy import deepcopy +from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse, StreamingResponse +from pydantic import ValidationError from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.a2a.version_convert import ( + A2AVersion, + normalize_agent_card, + normalize_jsonrpc_response, + normalize_request_params, + normalize_stream_event, +) from litellm.proxy.agent_endpoints.databricks_oauth import ( DATABRICKS_OAUTH_PARAM, resolve_databricks_app_auth_header, ) from litellm.proxy.agent_endpoints.utils import merge_agent_headers from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.utils import get_custom_url from litellm.types.utils import all_litellm_params +if TYPE_CHECKING: + from a2a.compat.v0_3.types import MessageSendParams + + from litellm.types.agents import AgentResponse + router = APIRouter() _PASCAL_TO_WIRE: Dict[str, str] = { + "SendMessage": "message/send", + "SendStreamingMessage": "message/stream", "GetTask": "tasks/get", "ListTasks": "tasks/list", "CancelTask": "tasks/cancel", @@ -38,6 +60,39 @@ _PASCAL_TO_WIRE: Dict[str, str] = { } +def _build_message_send_params(params: dict[str, Any]) -> "MessageSendParams": + """Build MessageSendParams from wire (0.3) or A2A 1.0 JSON-RPC params.""" + from a2a.compat.v0_3.types import MessageSendParams + + try: + return MessageSendParams(**params) + except ValidationError: + from a2a.compat.v0_3.conversions import pb2_v10, to_compat_send_message_request + from google.protobuf.json_format import ParseDict, ParseError + + pb = pb2_v10.SendMessageRequest() + try: + ParseDict(params, pb, ignore_unknown_fields=True) + except ParseError as e: + raise ValueError(f"Invalid message/send params: {e}") from e + return to_compat_send_message_request(pb, "").params + + +def _served_version(agent: "AgentResponse", request: Request, original_method: str | None = None) -> A2AVersion: + """Protocol version LiteLLM serves for this agent. + + The agent's configured version governs. For agents that pin no version, fall back + to the client's signal: PascalCase JSON-RPC methods and an ``a2a-version: 1.x`` + header both mark a 1.0 caller; otherwise default to 0.3. + """ + configured = (agent.agent_card_params or {}).get("protocolVersion") + if configured in ("0.3", "1.0"): + return configured + if original_method in _PASCAL_TO_WIRE: + return "1.0" + return "1.0" if request.headers.get("a2a-version", "").startswith("1.") else "0.3" + + def _validate_push_notification_url(url: str) -> None: parsed = urlparse(url) if parsed.scheme != "https": @@ -62,9 +117,9 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str def _forwarding_headers( user_api_key_dict: UserAPIKeyAuth, - request_data: dict, - agent_extra_headers: Optional[Dict[str, str]], -) -> Optional[Dict[str, str]]: + request_data: dict[str, Any], + agent_extra_headers: Dict[str, str] | None, +) -> Dict[str, str] | None: sanitized = ( {k: v for k, v in agent_extra_headers.items() if not k.lower().startswith("x-litellm-")} if agent_extra_headers @@ -80,7 +135,7 @@ def _forwarding_headers( def _jsonrpc_error( - request_id: Optional[Any], + request_id: Any | None, code: int, message: str, status_code: int = 400, @@ -125,9 +180,9 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None: async def _forward_jsonrpc( agent_url: str, - body: dict, - extra_headers: Optional[Dict[str, str]] = None, -) -> dict: + body: dict[str, Any], + extra_headers: Dict[str, str] | None = None, +) -> dict[str, Any]: from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider @@ -149,9 +204,10 @@ async def _forward_jsonrpc( async def _a2a_sse_event_source( agent_url: str, - body: dict, - request_id: Optional[Any] = None, - extra_headers: Optional[Dict[str, str]] = None, + body: dict[str, Any], + request_id: Any | None = None, + extra_headers: Dict[str, str] | None = None, + served_version: A2AVersion = "0.3", ) -> AsyncGenerator[dict, None]: """Stream an upstream A2A SSE response as parsed JSON-RPC event dicts. @@ -177,7 +233,7 @@ async def _a2a_sse_event_source( try: if not resp.is_success: error_body = await resp.aread() - error_event: Optional[dict] = None + error_event: dict[str, Any] | None = None try: parsed = json.loads(error_body) if isinstance(parsed, dict) and "error" in parsed: @@ -198,23 +254,33 @@ async def _a2a_sse_event_source( if not payload: continue try: - yield json.loads(payload) + event = json.loads(payload) except Exception: continue + if isinstance(event, dict): + event = normalize_stream_event(event, served_version, request_id=request_id) + yield event finally: await resp.aclose() async def _forward_jsonrpc_sse( agent_url: str, - body: dict, - request_id: Optional[Any] = None, - extra_headers: Optional[Dict[str, str]] = None, - proxy_logging_obj: Optional[Any] = None, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[dict] = None, + body: dict[str, Any], + request_id: Any | None = None, + extra_headers: Dict[str, str] | None = None, + proxy_logging_obj: Any | None = None, + user_api_key_dict: Any | None = None, + request_data: dict[str, Any] | None = None, + served_version: A2AVersion = "0.3", ) -> StreamingResponse: - event_source = _a2a_sse_event_source(agent_url, body, request_id=request_id, extra_headers=extra_headers) + event_source = _a2a_sse_event_source( + agent_url, + body, + request_id=request_id, + extra_headers=extra_headers, + served_version=served_version, + ) def _serialize_chunk(chunk: Any) -> str: return f"data: {json.dumps(chunk)}\n\n" @@ -263,18 +329,19 @@ async def _forward_jsonrpc_sse( async def _handle_stream_message( - api_base: Optional[str], + api_base: str | None, request_id: Any, - params: dict, - litellm_params: Optional[dict] = None, - agent_id: Optional[str] = None, - metadata: Optional[dict] = None, - proxy_server_request: Optional[dict] = None, + params: dict[str, Any], + litellm_params: dict[str, Any] | None = None, + agent_id: str | None = None, + metadata: dict[str, Any] | None = None, + proxy_server_request: dict[str, Any] | None = None, *, - agent_extra_headers: Optional[Dict[str, str]] = None, - user_api_key_dict: Optional[UserAPIKeyAuth] = None, - request_data: Optional[dict] = None, - proxy_logging_obj: Optional[Any] = None, + agent_extra_headers: Dict[str, str] | None = None, + user_api_key_dict: UserAPIKeyAuth | None = None, + request_data: dict[str, Any] | None = None, + proxy_logging_obj: Any | None = None, + served_version: A2AVersion = "0.3", ) -> StreamingResponse: """Handle message/stream method via SDK functions. @@ -304,15 +371,34 @@ async def _handle_stream_message( return StreamingResponse(_error_stream(), media_type="application/x-ndjson") - from a2a.types import MessageSendParams, SendStreamingMessageRequest + from a2a.compat.v0_3.types import SendStreamingMessageRequest use_proxy_hooks = user_api_key_dict is not None and request_data is not None and proxy_logging_obj is not None + try: + message_send_params = _build_message_send_params(params) + except (ValidationError, ValueError) as e: + invalid_params_message = f"Invalid params: {e}" + + async def _invalid_params_stream(): + yield ( + json.dumps( + { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32602, "message": invalid_params_message}, + } + ) + + "\n" + ) + + return StreamingResponse(_invalid_params_stream(), media_type="application/x-ndjson") + async def stream_response(): try: a2a_request = SendStreamingMessageRequest( id=request_id, - params=MessageSendParams(**params), + params=message_send_params, ) a2a_stream = asend_message_streaming( request=a2a_request, @@ -339,6 +425,8 @@ async def _handle_stream_message( obj = chunk.model_dump(mode="json", exclude_none=True) else: obj = chunk + if isinstance(obj, dict): + obj = normalize_stream_event(obj, served_version, request_id=request_id) return json.dumps(obj) + "\n" def _ndjson_error(proxy_exc: Any) -> str: @@ -372,9 +460,12 @@ async def _handle_stream_message( else: async for chunk in a2a_stream: if hasattr(chunk, "model_dump"): - yield (json.dumps(chunk.model_dump(mode="json", exclude_none=True)) + "\n") + obj = chunk.model_dump(mode="json", exclude_none=True) else: - yield json.dumps(chunk) + "\n" + obj = chunk + if isinstance(obj, dict): + obj = normalize_stream_event(obj, served_version, request_id=request_id) + yield json.dumps(obj) + "\n" except Exception as e: verbose_proxy_logger.exception(f"Error streaming A2A response: {e}") if ( @@ -460,13 +551,16 @@ async def get_agent_card( detail=f"Agent '{agent_id}' has no agent card configured", ) - # Copy and rewrite URL to point to LiteLLM proxy - agent_card = { - **agent.agent_card_params, - "url": f"{str(request.base_url).rstrip('/')}/a2a/{agent_id}", - } + proxy_url = get_custom_url(str(request.base_url), route=f"a2a/{agent_id}") + agent_card = deepcopy(agent.agent_card_params) + agent_card["url"] = proxy_url + interfaces = agent_card.get("supportedInterfaces") + if isinstance(interfaces, list) and interfaces: + interfaces[0]["url"] = proxy_url + served_version = _served_version(agent, request) + agent_card = normalize_agent_card(agent_card, served_version) - verbose_proxy_logger.debug(f"Returning agent card for '{agent_id}' with proxy URL: {agent_card['url']}") + verbose_proxy_logger.debug(f"Returning agent card for '{agent_id}' with proxy URL: {proxy_url}") return JSONResponse(content=agent_card) except HTTPException: @@ -526,8 +620,9 @@ async def invoke_agent_a2a( if body.get("jsonrpc") != "2.0": return _jsonrpc_error(body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'") - request_id: Optional[Any] = body.get("id") - method: Optional[str] = body.get("method") + request_id: Any | None = body.get("id") + original_method: str | None = body.get("method") + method: str | None = original_method params = body.get("params", {}) if method: @@ -553,6 +648,8 @@ async def invoke_agent_a2a( if agent is None: return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) + served_version = _served_version(agent, request, original_method) + is_allowed = await AgentRequestHandler.is_agent_allowed( agent_id=agent.agent_id, user_api_key_auth=user_api_key_dict, @@ -691,11 +788,16 @@ async def invoke_agent_a2a( "Server error: 'a2a' package not installed. Please install 'a2a-sdk'.", 500, ) - from a2a.types import MessageSendParams, SendMessageRequest + from a2a.compat.v0_3.types import SendMessageRequest + + try: + message_send_params = _build_message_send_params(params) + except (ValidationError, ValueError) as e: + return _jsonrpc_error(request_id, -32602, f"Invalid params: {e}") a2a_request = SendMessageRequest( id=request_id if request_id is not None else "", - params=MessageSendParams(**params), + params=message_send_params, ) # Defer spend-log until after post_call_success_hook so guardrail # results written by the unified_guardrail hook are captured. @@ -723,11 +825,18 @@ async def invoke_agent_a2a( logging_obj._enqueue_deferred_logging = None # type: ignore[union-attr] _enqueue_fn() + response_dict: Dict[str, Any] = ( + response.model_dump(mode="json", exclude_none=True) # type: ignore + if hasattr(response, "model_dump") + else response + if isinstance(response, dict) + else {} + ) return JSONResponse( - content=( - response.model_dump(mode="json", exclude_none=True) # type: ignore - if hasattr(response, "model_dump") - else response + content=normalize_jsonrpc_response( + response_dict, + served_version, + method="message/send", ) ) @@ -744,6 +853,7 @@ async def invoke_agent_a2a( user_api_key_dict=user_api_key_dict, request_data=data, proxy_logging_obj=proxy_logging_obj, + served_version=served_version, ) elif method in { "tasks/get", @@ -757,6 +867,8 @@ async def invoke_agent_a2a( }: if not agent_url: return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) + if isinstance(params, dict): + params = normalize_request_params(params, served_version, method=method) if method == "tasks/pushNotificationConfig/set": if not isinstance(params, dict): raise HTTPException( @@ -791,8 +903,20 @@ async def invoke_agent_a2a( ) result = await _forward_jsonrpc(agent_url, forward_body, extra_headers=caller_headers) if method == "agent/getAuthenticatedExtendedCard": - if isinstance(result.get("result"), dict) and "url" in result["result"]: - result["result"]["url"] = f"{str(request.base_url).rstrip('/')}/a2a/{agent_id}" + if isinstance(result.get("result"), dict): + card = result["result"] + proxy_url = get_custom_url(str(request.base_url), route=f"a2a/{agent_id}") + # Rewrite the upstream agent URL in both 0.3 (top-level `url`) + # and 1.0 (`supportedInterfaces[0].url`) wire formats so that + # downstream clients never see the upstream internal address. + if "url" in card: + card["url"] = proxy_url + interfaces = card.get("supportedInterfaces") + if isinstance(interfaces, list) and interfaces: + interfaces[0]["url"] = proxy_url + result["result"] = normalize_agent_card(card, served_version) + else: + result = normalize_jsonrpc_response(result, served_version, method=method) from litellm.types.agents import LiteLLMSendMessageResponse response = LiteLLMSendMessageResponse.from_dict(result, request_id=request_id) @@ -810,6 +934,8 @@ async def invoke_agent_a2a( elif method == "tasks/resubscribe": if not agent_url: return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) + if isinstance(params, dict): + params = normalize_request_params(params, served_version, method=method) forward_body = { "jsonrpc": "2.0", "id": request_id, @@ -829,6 +955,7 @@ async def invoke_agent_a2a( proxy_logging_obj=proxy_logging_obj, user_api_key_dict=user_api_key_dict, request_data=data, + served_version=served_version, ) else: diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 450c6414270..51e451efbb9 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -11,7 +11,7 @@ Follows the A2A Spec. import asyncio import os import uuid -from typing import Any, Dict, List, Mapping, Optional +from typing import Any, Dict, List, Mapping from fastapi import APIRouter, Depends, HTTPException, Query, Request @@ -20,10 +20,14 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import _get_masked_values from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.a2a.agent_card import merge_agent_card +from litellm.proxy.a2a.agent_card import ( + SUPPORTED_A2A_PROTOCOL_VERSIONS, + merge_agent_card, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity +from litellm.proxy.utils import get_custom_url from litellm.types.agents import ( AgentConfig, AgentKeySummary, @@ -40,19 +44,33 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import ( def _proxy_base_url(http_request: Request) -> str: - """Return the proxy's base URL as seen by the caller, without trailing slash.""" - return str(http_request.base_url).rstrip("/") + """Return the proxy's public base URL, preferring PROXY_BASE_URL when set.""" + return get_custom_url(str(http_request.base_url), route=None) + + +def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None: + """Reject an agent card pinning an unsupported A2A protocol version.""" + version = upstream_card.get("protocolVersion") if upstream_card else None + if version is not None and version not in SUPPORTED_A2A_PROTOCOL_VERSIONS: + raise HTTPException( + status_code=400, + detail=( + f"Unsupported protocolVersion '{version}'. " + f"Supported versions: {', '.join(SUPPORTED_A2A_PROTOCOL_VERSIONS)}." + ), + ) def _build_merged_agent_card( - upstream_card: Optional[Mapping[str, Any]], + upstream_card: Mapping[str, Any] | None, *, agent_id: str, http_request: Request, - agent_name: Optional[str] = None, + agent_name: str | None = None, ) -> Dict[str, Any]: """Apply the LiteLLM-fronting merge to ``upstream_card`` for ``agent_id``.""" proxy_base = _proxy_base_url(http_request) + _validate_protocol_version(upstream_card) # Prefer a card-supplied ``name`` (the discovery UI exposes an editable # "Name (shown to API clients)" field that flows into # ``agent_card_params.name``) over the internal ``agent_name`` identifier. @@ -382,7 +400,7 @@ async def create_agent( # schemes, default skills) the agent doesn't actually expose. upstream_card = request.get("agent_card_params") agent_to_create: AgentConfig = request - new_agent_id: Optional[str] = None + new_agent_id: str | None = None if upstream_card is not None: # Pre-generate the agent_id so the merged card can reference it # in ``supportedInterfaces`` before the DB row exists. @@ -988,14 +1006,14 @@ async def make_agents_public( response_model=SpendAnalyticsPaginatedResponse, ) async def get_agent_daily_activity( - agent_ids: Optional[str] = None, - start_date: Optional[str] = None, - end_date: Optional[str] = None, - model: Optional[str] = None, - api_key: Optional[str] = None, + agent_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, page: int = 1, page_size: int = 10, - exclude_agent_ids: Optional[str] = None, + exclude_agent_ids: str | None = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -1012,7 +1030,7 @@ async def get_agent_daily_activity( ) agent_ids_list = agent_ids.split(",") if agent_ids else None - exclude_agent_ids_list: Optional[List[str]] = None + exclude_agent_ids_list: List[str] | None = None if exclude_agent_ids: exclude_agent_ids_list = exclude_agent_ids.split(",") if exclude_agent_ids else None diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 91332480d75..76d211d2c8f 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -17,6 +17,7 @@ from litellm.litellm_core_utils.get_blog_posts import ( from litellm.proxy._types import ( CommonProxyErrors, ) +from litellm.proxy.utils import get_custom_url from litellm.repositories.table_repositories import ClaudeCodePluginRepository from litellm.types.agents import AgentCard from litellm.types.mcp import MCPPublicServer @@ -213,11 +214,10 @@ async def get_agents(request: Request): if litellm.public_agent_groups is None: return [] - proxy_base = str(request.base_url).rstrip("/") return [ { **(agent.agent_card_params or {}), - "url": f"{proxy_base}/a2a/{agent.agent_id}", + "url": get_custom_url(str(request.base_url), route=f"a2a/{agent.agent_id}"), } for agent in agents if agent.agent_id in litellm.public_agent_groups diff --git a/pyproject.toml b/pyproject.toml index 23f02cb5762..2ad96c4936b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -90,7 +90,7 @@ extra_proxy = [ # Not in PyPI proxy extra. "resend>=2.23.0,<3.0", "redisvl>=0.4.1,<1.0; python_version < '3.14'", - "a2a-sdk>=0.3.24,<1.0", + "a2a-sdk>=1.1.0,<2.0", ] utils = [ # Not in Docker or PyPI proxy extra. @@ -193,7 +193,7 @@ proxy-dev = [ "opentelemetry-exporter-otlp==1.28.0", "opentelemetry-instrumentation-fastapi==0.49b0", "azure-identity==1.25.2", - "a2a-sdk==0.3.24", + "a2a-sdk==1.1.0", ] ci = [ # These are lazily imported at call sites; keep them out of core deps to @@ -241,6 +241,11 @@ build-backend = "uv_build" constraint-dependencies = [ "tornado>=6.5.6", "aiohttp>=3.14.1,<4.0", + "packaging>=24.0", +] +override-dependencies = [ + # a2a-sdk 1.x requires packaging>=24.0; lunary 1.4.x still caps at <24.0. + "packaging>=24.0", ] default-groups = ["dev"] required-version = ">=0.10.9" diff --git a/tests/agent_tests/test_a2a_agent.py b/tests/agent_tests/test_a2a_agent.py index ed7a5ab9823..f5ad9601369 100644 --- a/tests/agent_tests/test_a2a_agent.py +++ b/tests/agent_tests/test_a2a_agent.py @@ -41,21 +41,24 @@ class MockA2AClient: ) async def send_message(self, request): - return MockA2AResponse(text="hello") + from a2a.compat.v0_3.conversions import pb2_v10 - def send_message_streaming(self, request): - async def _stream(): - yield MockA2AStreamingChunk(text="hel", state="in_progress") - yield MockA2AStreamingChunk(text="hello", state="completed") - - return _stream() + for text in ("hel", "hello"): + event = pb2_v10.StreamResponse() + message = event.message + message.message_id = uuid4().hex + message.role = pb2_v10.ROLE_AGENT + message.parts.add().text = text + yield event @pytest.fixture def mock_a2a_client(monkeypatch): import litellm.a2a_protocol.main as a2a_main - async def _fake_create_a2a_client(base_url, timeout=60.0, extra_headers=None): + async def _fake_create_a2a_client( + base_url, timeout=60.0, extra_headers=None, streaming=False + ): return MockA2AClient() monkeypatch.setattr(a2a_main, "create_a2a_client", _fake_create_a2a_client) @@ -64,7 +67,7 @@ def mock_a2a_client(monkeypatch): @pytest.mark.asyncio async def test_a2a_non_streaming(mock_a2a_client): """Test non-streaming A2A request.""" - from a2a.types import MessageSendParams, SendMessageRequest + from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest from litellm.a2a_protocol import asend_message request = SendMessageRequest( @@ -90,7 +93,7 @@ async def test_a2a_non_streaming(mock_a2a_client): @pytest.mark.asyncio async def test_a2a_streaming(mock_a2a_client): """Test streaming A2A request.""" - from a2a.types import MessageSendParams, SendStreamingMessageRequest + from a2a.compat.v0_3.types import MessageSendParams, SendStreamingMessageRequest from litellm.a2a_protocol import asend_message_streaming request = SendStreamingMessageRequest( diff --git a/tests/test_litellm/a2a_protocol/test_a2a_exception_mapping_utils.py b/tests/test_litellm/a2a_protocol/test_a2a_exception_mapping_utils.py new file mode 100644 index 00000000000..44d2803a260 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/test_a2a_exception_mapping_utils.py @@ -0,0 +1,153 @@ +"""Tests for litellm/a2a_protocol/exception_mapping_utils.py.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.a2a_protocol import exception_mapping_utils as emu +from litellm.a2a_protocol.exceptions import A2ALocalhostURLError + + +def _localhost_error() -> A2ALocalhostURLError: + return A2ALocalhostURLError( + localhost_url="http://localhost:10001/", + base_url="https://agent.example", + original_error=ConnectionError("boom"), + ) + + +@pytest.mark.asyncio +async def test_localhost_retry_reuses_stashed_httpx_client(): + """The retry must reuse the httpx client LiteLLM attached at creation (it carries + the agent's trace-id/auth headers), passing it straight into the new ClientConfig. + """ + stashed_httpx_client = object() + a2a_client = MagicMock() + a2a_client._litellm_httpx_client = stashed_httpx_client + new_client = MagicMock() + + captured = {} + + def fake_client_config(*, httpx_client, streaming): + captured["httpx_client"] = httpx_client + captured["streaming"] = streaming + return MagicMock() + + with ( + patch.object(emu, "A2A_SDK_AVAILABLE", True), + patch.object(emu, "set_agent_card_url") as mock_set_url, + patch.object(emu, "ClientConfig", side_effect=fake_client_config), + patch.object( + emu, "create_client", new=AsyncMock(return_value=new_client) + ) as mock_create, + ): + result = await emu.handle_a2a_localhost_retry( + error=_localhost_error(), + agent_card=MagicMock(), + a2a_client=a2a_client, + is_streaming=True, + ) + + assert result is new_client + mock_set_url.assert_called_once() + # The exact stashed client is threaded through, not a freshly built one. + assert captured["httpx_client"] is stashed_httpx_client + assert captured["streaming"] is True + assert new_client._litellm_httpx_client is stashed_httpx_client + assert mock_create.await_count == 1 + + +@pytest.mark.asyncio +async def test_localhost_retry_raises_when_no_stashed_client(): + """An externally-supplied client has no LiteLLM httpx handle; the retry must fail + with a clear error instead of excavating a2a-sdk internals.""" + a2a_client = MagicMock(spec=[]) # no _litellm_httpx_client attribute + + with ( + patch.object(emu, "A2A_SDK_AVAILABLE", True), + patch.object(emu, "set_agent_card_url"), + patch.object(emu, "create_client", new=AsyncMock()) as mock_create, + ): + with pytest.raises(RuntimeError, match="not created by create_a2a_client"): + await emu.handle_a2a_localhost_retry( + error=_localhost_error(), + agent_card=MagicMock(), + a2a_client=a2a_client, + is_streaming=False, + ) + + mock_create.assert_not_called() + + +@pytest.mark.asyncio +async def test_localhost_retry_raises_when_agent_card_is_none(): + """With no agent card to rewrite, the retry must fail with a clear error instead + of calling create_client(None, ...) and surfacing an opaque SDK TypeError.""" + a2a_client = MagicMock() + a2a_client._litellm_httpx_client = MagicMock() + + with ( + patch.object(emu, "A2A_SDK_AVAILABLE", True), + patch.object(emu, "set_agent_card_url") as mock_set_url, + patch.object(emu, "create_client", new=AsyncMock()) as mock_create, + ): + with pytest.raises(RuntimeError, match="no agent card is available"): + await emu.handle_a2a_localhost_retry( + error=_localhost_error(), + agent_card=None, + a2a_client=a2a_client, + is_streaming=False, + ) + + mock_set_url.assert_not_called() + mock_create.assert_not_called() + + +def test_get_a2a_client_agent_card_reads_sdk_private_card(): + from litellm.a2a_protocol.main import _get_a2a_client_agent_card + + sdk_card = SimpleNamespace(name="Test Agent", url="http://localhost:10001/") + a2a_client = SimpleNamespace(_card=sdk_card) + + assert _get_a2a_client_agent_card(a2a_client) is sdk_card + + +@pytest.mark.asyncio +async def test_stream_with_retry_raises_after_localhost_retries_exhausted(): + """Exhausted localhost retries must not return a silent empty stream.""" + from litellm.a2a_protocol.main import _execute_a2a_stream_with_retry + + localhost_err = _localhost_error() + mock_request = MagicMock() + mock_request.id = "req-1" + mock_a2a_client = MagicMock() + + async def _always_fail_stream(a2a_client, request): + raise localhost_err + yield # pragma: no cover - makes this an async generator + + with ( + patch( + "litellm.a2a_protocol.main._stream_messages", + new=_always_fail_stream, + ), + patch( + "litellm.a2a_protocol.main.handle_a2a_localhost_retry", + new=AsyncMock(return_value=mock_a2a_client), + ), + ): + stream = _execute_a2a_stream_with_retry( + a2a_client=mock_a2a_client, + request=mock_request, + agent_card=MagicMock(), + card_url="http://localhost:10001/", + api_base="https://agent.example", + agent_name="test-agent", + ) + with pytest.raises( + RuntimeError, + match="no response received after retry attempts", + ): + async for _chunk in stream: + pytest.fail("expected retry exhaustion to raise before yielding") diff --git a/tests/test_litellm/a2a_protocol/test_card_resolver.py b/tests/test_litellm/a2a_protocol/test_card_resolver.py index 1bdab50860c..053f28c940f 100644 --- a/tests/test_litellm/a2a_protocol/test_card_resolver.py +++ b/tests/test_litellm/a2a_protocol/test_card_resolver.py @@ -4,7 +4,8 @@ Mock tests for LiteLLMA2ACardResolver. Tests that the card resolver tries both old and new well-known paths. """ -from unittest.mock import AsyncMock, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import MagicMock, patch import pytest @@ -12,6 +13,7 @@ from litellm.a2a_protocol.card_resolver import ( LiteLLMA2ACardResolver, fix_agent_card_url, is_localhost_or_internal_url, + set_agent_card_url, ) @@ -88,3 +90,27 @@ def test_fix_agent_card_url_replaces_localhost(): # Verify localhost URL was replaced with base_url assert result.url == "https://my-public-agent.example.com/" + + +def test_set_agent_card_url_updates_top_level_and_supported_interface(): + card = SimpleNamespace( + url="http://localhost:10001/", + supported_interfaces=[SimpleNamespace(url="http://0.0.0.0:10001/")], + ) + + set_agent_card_url(card, "https://my-public-agent.example.com") + + assert card.url == "https://my-public-agent.example.com/" + assert card.supported_interfaces[0].url == "https://my-public-agent.example.com/" + + +def test_fix_agent_card_url_updates_interface_when_top_level_is_localhost(): + card = SimpleNamespace( + url="http://localhost:10001/", + supported_interfaces=[SimpleNamespace(url="http://0.0.0.0:10001/")], + ) + + result = fix_agent_card_url(card, "https://my-public-agent.example.com") + + assert result.url == "https://my-public-agent.example.com/" + assert result.supported_interfaces[0].url == "https://my-public-agent.example.com/" diff --git a/tests/test_litellm/a2a_protocol/test_cost_calculator.py b/tests/test_litellm/a2a_protocol/test_cost_calculator.py index d7bacaf39eb..bf03562f8ae 100644 --- a/tests/test_litellm/a2a_protocol/test_cost_calculator.py +++ b/tests/test_litellm/a2a_protocol/test_cost_calculator.py @@ -3,8 +3,8 @@ Test A2A cost calculator with cost_per_query parameter. """ import asyncio -from typing import Optional -from unittest.mock import AsyncMock, MagicMock +from typing import Any, AsyncIterator, Optional +from unittest.mock import MagicMock, patch import pytest @@ -12,6 +12,102 @@ import litellm from litellm.integrations.custom_logger import CustomLogger +def _make_send_message_request(request_id: str, user_text: str = "Hello"): + from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest + + return SendMessageRequest( + id=request_id, + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": user_text}], + "messageId": "msg-1", + } + ), + ) + + +async def _mock_execute_a2a_send( + a2a_client: Any, + request: Any, + **kwargs: Any, +) -> Any: + mock_response = MagicMock() + mock_response.model_dump = MagicMock( + return_value={ + "id": request.id, + "jsonrpc": "2.0", + "result": {"status": "completed"}, + } + ) + return mock_response + + +async def _mock_execute_a2a_send_with_assistant_reply( + a2a_client: Any, + request: Any, + **kwargs: Any, +) -> Any: + mock_response = MagicMock() + mock_response.model_dump = MagicMock( + return_value={ + "id": request.id, + "jsonrpc": "2.0", + "result": { + "status": {"state": "completed"}, + "message": { + "role": "assistant", + "parts": [ + { + "kind": "text", + "text": "Hello! I am your assistant. How can I help you today?", + } + ], + "messageId": "msg-456", + }, + }, + } + ) + return mock_response + + +def _make_streaming_request(request_id: str): + from a2a.compat.v0_3.types import MessageSendParams, SendStreamingMessageRequest + + return SendStreamingMessageRequest( + id=request_id, + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Hello"}], + "messageId": "msg-1", + } + ), + ) + + +async def _mock_stream_messages(a2a_client: Any, request: Any) -> AsyncIterator[Any]: + from a2a.compat.v0_3.types import ( + Message, + Part, + Role, + SendStreamingMessageResponse, + SendStreamingMessageSuccessResponse, + TextPart, + ) + + msg = Message( + message_id="msg-agent", + role=Role.agent, + parts=[Part(root=TextPart(kind="text", text="hello"))], + kind="message", + ) + for _ in range(2): + yield SendStreamingMessageResponse( + root=SendStreamingMessageSuccessResponse(id=request.id, result=msg) + ) + + class CostLogger(CustomLogger): """Custom logger to capture response_cost.""" @@ -46,27 +142,18 @@ async def test_asend_message_uses_cost_per_query(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - # Mock response with required fields - mock_response = MagicMock() - mock_response.model_dump = MagicMock( - return_value={ - "id": "test-123", - "jsonrpc": "2.0", - "result": {"status": "completed"}, - } - ) - mock_client.send_message = AsyncMock(return_value=mock_response) - - # Mock request - mock_request = MagicMock() - mock_request.id = "test-123" + mock_request = _make_send_message_request("test-123") # Call asend_message with cost_per_query - await asend_message( - a2a_client=mock_client, - request=mock_request, - cost_per_query=0.05, - ) + with patch( + "litellm.a2a_protocol.main._execute_a2a_send_with_retry", + new=_mock_execute_a2a_send, + ): + await asend_message( + a2a_client=mock_client, + request=mock_request, + cost_per_query=0.05, + ) await asyncio.sleep(0.1) @@ -120,49 +207,24 @@ async def test_asend_message_uses_input_output_cost_per_token(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - # Realistic A2A response with message parts - mock_response = MagicMock() - mock_response.model_dump = MagicMock( - return_value={ - "id": "test-123", - "jsonrpc": "2.0", - "result": { - "status": {"state": "completed"}, - "message": { - "role": "assistant", - "parts": [ - { - "kind": "text", - "text": "Hello! I am your assistant. How can I help you today?", - } - ], - "messageId": "msg-456", - }, - }, - } + mock_request = _make_send_message_request( + "test-123", user_text="Hello, what can you do?" ) - mock_client.send_message = AsyncMock(return_value=mock_response) - - # Mock request with message parts - mock_request = MagicMock() - mock_request.id = "test-123" - mock_request.params = MagicMock() - mock_request.params.message = { - "role": "user", - "parts": [{"kind": "text", "text": "Hello, what can you do?"}], - "messageId": "msg-123", - } # Define specific cost per token values input_cost_per_token = 0.00001 # $0.01 per 1000 tokens output_cost_per_token = 0.00002 # $0.02 per 1000 tokens - await asend_message( - a2a_client=mock_client, - request=mock_request, - input_cost_per_token=input_cost_per_token, - output_cost_per_token=output_cost_per_token, - ) + with patch( + "litellm.a2a_protocol.main._execute_a2a_send_with_retry", + new=_mock_execute_a2a_send_with_assistant_reply, + ): + await asend_message( + a2a_client=mock_client, + request=mock_request, + input_cost_per_token=input_cost_per_token, + output_cost_per_token=output_cost_per_token, + ) await asyncio.sleep(0.1) @@ -225,29 +287,20 @@ async def test_asend_message_passes_agent_id_to_callback(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - # Mock response - mock_response = MagicMock() - mock_response.model_dump = MagicMock( - return_value={ - "id": "test-123", - "jsonrpc": "2.0", - "result": {"status": "completed"}, - } - ) - mock_client.send_message = AsyncMock(return_value=mock_response) - - # Mock request - mock_request = MagicMock() - mock_request.id = "test-123" + mock_request = _make_send_message_request("test-123") test_agent_id = "agent-uuid-12345" # Call asend_message with agent_id - await asend_message( - a2a_client=mock_client, - request=mock_request, - agent_id=test_agent_id, - ) + with patch( + "litellm.a2a_protocol.main._execute_a2a_send_with_retry", + new=_mock_execute_a2a_send, + ): + await asend_message( + a2a_client=mock_client, + request=mock_request, + agent_id=test_agent_id, + ) await asyncio.sleep(0.1) @@ -294,21 +347,7 @@ async def test_asend_message_streaming_propagates_metadata(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - # Mock streaming response - async def mock_stream(): - yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 1}) - yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 2}) - - mock_client.send_message_streaming = MagicMock(return_value=mock_stream()) - - # Mock request - mock_request = MagicMock() - mock_request.id = "test-stream-metadata" - mock_request.params = MagicMock() - mock_request.params.message = { - "role": "user", - "parts": [{"kind": "text", "text": "Hello"}], - } + mock_request = _make_streaming_request("test-stream-metadata") # Metadata from proxy (contains user_api_key, user_id, team_id for SpendLogs) test_metadata = { @@ -319,12 +358,16 @@ async def test_asend_message_streaming_propagates_metadata(): # Consume streaming response with metadata chunks = [] - async for chunk in asend_message_streaming( - a2a_client=mock_client, - request=mock_request, - metadata=test_metadata, + with patch( + "litellm.a2a_protocol.main._stream_messages", + new=_mock_stream_messages, ): - chunks.append(chunk) + async for chunk in asend_message_streaming( + a2a_client=mock_client, + request=mock_request, + metadata=test_metadata, + ): + chunks.append(chunk) await asyncio.sleep(0.2) @@ -352,32 +395,22 @@ async def test_asend_message_streaming_triggers_callbacks(): mock_client._litellm_agent_card = MagicMock() mock_client._litellm_agent_card.name = "test-agent" - # Mock streaming response - async def mock_stream(): - yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 1}) - yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 2}) - - mock_client.send_message_streaming = MagicMock(return_value=mock_stream()) - - # Mock request - mock_request = MagicMock() - mock_request.id = "test-stream-123" - mock_request.params = MagicMock() - mock_request.params.message = { - "role": "user", - "parts": [{"kind": "text", "text": "Hello"}], - } + mock_request = _make_streaming_request("test-stream-123") test_agent_id = "test-agent-id-streaming" # Consume streaming response chunks = [] - async for chunk in asend_message_streaming( - a2a_client=mock_client, - request=mock_request, - agent_id=test_agent_id, + with patch( + "litellm.a2a_protocol.main._stream_messages", + new=_mock_stream_messages, ): - chunks.append(chunk) + async for chunk in asend_message_streaming( + a2a_client=mock_client, + request=mock_request, + agent_id=test_agent_id, + ): + chunks.append(chunk) await asyncio.sleep(0.2) diff --git a/tests/test_litellm/a2a_protocol/test_main.py b/tests/test_litellm/a2a_protocol/test_main.py new file mode 100644 index 00000000000..2a675616245 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/test_main.py @@ -0,0 +1,106 @@ +"""Tests for litellm/a2a_protocol/main.py non-streaming send behavior.""" + +import pytest + +pytest.importorskip("a2a.compat.v0_3.conversions") + +from a2a.compat.v0_3 import conversions as _conv +from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest + +from litellm.a2a_protocol.main import _send_message + + +def _request() -> SendMessageRequest: + params = MessageSendParams( + message={ + "messageId": "m1", + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + } + ) + return SendMessageRequest(id="r1", params=params) + + +def _message_stream_response(): + sr = _conv.pb2_v10.StreamResponse() + sr.message.message_id = "reply-1" + sr.message.role = _conv.pb2_v10.Role.ROLE_AGENT + sr.message.parts.add().text = "hello back" + return sr + + +def _status_update_stream_response(): + sr = _conv.pb2_v10.StreamResponse() + sr.status_update.task_id = "t1" + sr.status_update.context_id = "c1" + return sr + + +class _FakeClient: + def __init__(self, *events): + self._events = events + + async def send_message(self, _pb_request): + for event in self._events: + yield event + + +@pytest.mark.asyncio +async def test_send_message_returns_message_result(): + response = await _send_message(_FakeClient(_message_stream_response()), _request()) + result = response.root.result + assert type(result).__name__ == "Message" + assert response.root.id == "r1" + + +@pytest.mark.asyncio +async def test_send_message_rejects_update_event_final_with_runtime_error(): + with pytest.raises(RuntimeError, match="Message or Task"): + await _send_message(_FakeClient(_status_update_stream_response()), _request()) + + +@pytest.mark.asyncio +async def test_streaming_trace_id_prefers_logging_trace_id(): + """The streaming X-LiteLLM-Trace-Id must use the logging object's trace id (same + as the non-streaming path), not the JSON-RPC request id, so traces correlate.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from a2a.compat.v0_3.types import ( + MessageSendParams, + SendStreamingMessageRequest, + ) + + from litellm.a2a_protocol import main as a2a_main + from litellm.litellm_core_utils.litellm_logging import Logging + + request = SendStreamingMessageRequest( + id="rpc-1", + params=MessageSendParams( + message={ + "messageId": "m1", + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + } + ), + ) + logging_obj = MagicMock(spec=Logging) + logging_obj.litellm_trace_id = "trace-from-logging" + + captured: dict = {} + + async def _capture(*, base_url, extra_headers=None, streaming=False, **_): + captured["extra_headers"] = extra_headers + raise RuntimeError("stop") + + with patch.object( + a2a_main, "create_a2a_client", new=AsyncMock(side_effect=_capture) + ): + with pytest.raises(RuntimeError, match="stop"): + async for _ in a2a_main.asend_message_streaming( + request=request, + api_base="http://upstream.local", + litellm_logging_obj=logging_obj, + ): + pass + + assert captured["extra_headers"]["X-LiteLLM-Trace-Id"] == "trace-from-logging" diff --git a/tests/test_litellm/proxy/a2a/test_agent_card.py b/tests/test_litellm/proxy/a2a/test_agent_card.py index 0022053d8d1..d302bde7895 100644 --- a/tests/test_litellm/proxy/a2a/test_agent_card.py +++ b/tests/test_litellm/proxy/a2a/test_agent_card.py @@ -49,13 +49,31 @@ def test_preserves_top_level_url_for_runtime_invocation(): assert merged["url"] == "http://internal:9999/" -def test_overrides_protocol_version(): +def test_unsupported_protocol_version_defaults_to_1_0(): + # The fixture card pins "0.9", which LiteLLM does not serve; it falls back to + # the default rather than advertising a version the proxy can't honor. merged = merge_agent_card( _full_upstream_card(), proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE ) assert merged["protocolVersion"] == LITELLM_A2A_PROTOCOL_VERSION +def test_serves_pinned_protocol_version(): + for version in ("0.3", "1.0"): + card = _full_upstream_card() + card["protocolVersion"] = version + merged = merge_agent_card(card, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE) + assert merged["protocolVersion"] == version + assert merged["supportedInterfaces"][0]["protocolVersion"] == version + + +def test_absent_protocol_version_defaults_to_1_0(): + card = _full_upstream_card() + card.pop("protocolVersion", None) + merged = merge_agent_card(card, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE) + assert merged["protocolVersion"] == "1.0" + + def test_overrides_name_and_description_when_provided(): merged = merge_agent_card( _full_upstream_card(), diff --git a/tests/test_litellm/proxy/a2a/test_version_convert.py b/tests/test_litellm/proxy/a2a/test_version_convert.py new file mode 100644 index 00000000000..f3c51ca6b72 --- /dev/null +++ b/tests/test_litellm/proxy/a2a/test_version_convert.py @@ -0,0 +1,315 @@ +"""Unit tests for A2A protocol version normalization in +litellm/proxy/a2a/version_convert.py. + +These assert the conversion actually changes wire shape in the right direction and +preserves core fields on a round trip, so a mutation that no-ops or flips the direction +fails the suite. +""" + +import pytest + +from litellm.proxy.a2a.version_convert import ( + normalize_agent_card, + normalize_jsonrpc_response, + normalize_request_params, + normalize_stream_event, +) + +a2a = pytest.importorskip("a2a.compat.v0_3.conversions") + + +def _rpc(result: dict, request_id: str = "1") -> dict: + return {"jsonrpc": "2.0", "id": request_id, "result": result} + + +V03_MESSAGE = { + "kind": "message", + "messageId": "m1", + "role": "agent", + "parts": [{"kind": "text", "text": "hi"}], +} + +V03_TASK = { + "kind": "task", + "id": "t1", + "contextId": "c1", + "status": {"state": "completed"}, +} + +V03_STATUS_UPDATE = { + "kind": "status-update", + "taskId": "t1", + "contextId": "c1", + "status": {"state": "working"}, + "final": False, +} + +V03_ARTIFACT_UPDATE = { + "kind": "artifact-update", + "taskId": "t1", + "contextId": "c1", + "artifact": {"artifactId": "a1", "parts": [{"kind": "text", "text": "out"}]}, +} + + +def test_send_result_0_3_to_1_0_wraps_in_envelope(): + out = normalize_jsonrpc_response(_rpc(V03_MESSAGE), "1.0", method="message/send") + assert "message" in out["result"] + assert "kind" not in out["result"] + assert out["result"]["message"]["messageId"] == "m1" + + +def test_send_result_1_0_to_0_3_unwraps_to_bare_kind(): + v1 = normalize_jsonrpc_response(_rpc(V03_MESSAGE), "1.0", method="message/send") + out = normalize_jsonrpc_response(v1, "0.3", method="message/send") + assert out["result"]["kind"] == "message" + assert out["result"]["messageId"] == "m1" + assert out["result"]["parts"][0]["text"] == "hi" + + +def test_send_result_same_version_is_identity_passthrough(): + rpc = _rpc(V03_MESSAGE) + out = normalize_jsonrpc_response(rpc, "0.3", method="message/send") + assert out is rpc + + +def test_task_result_round_trip_preserves_ids(): + v1 = normalize_jsonrpc_response(_rpc(V03_TASK), "1.0", method="tasks/get") + assert "kind" not in v1["result"] + assert v1["result"]["id"] == "t1" + back = normalize_jsonrpc_response(v1, "0.3", method="tasks/get") + assert back["result"]["kind"] == "task" + assert back["result"]["id"] == "t1" + assert back["result"]["contextId"] == "c1" + + +def test_error_response_passes_through_untouched(): + err = {"jsonrpc": "2.0", "id": "1", "error": {"code": -32600, "message": "bad"}} + assert normalize_jsonrpc_response(err, "1.0", method="message/send") is err + + +def test_malformed_result_falls_back_to_passthrough(): + # A 0.3 message missing required fields can't validate; conversion must not raise. + rpc = _rpc({"kind": "message"}) + out = normalize_jsonrpc_response(rpc, "1.0", method="message/send") + assert out["result"] == {"kind": "message"} + + +def test_unknown_shape_passes_through(): + rpc = _rpc({"unexpected": "shape"}) + out = normalize_jsonrpc_response(rpc, "1.0", method="message/send") + assert out is rpc + + +@pytest.mark.parametrize("event", [V03_STATUS_UPDATE, V03_ARTIFACT_UPDATE]) +def test_stream_event_round_trip_preserves_kind(event): + v1 = normalize_stream_event(_rpc(event), "1.0", request_id="1") + assert "kind" not in v1["result"] + back = normalize_stream_event(v1, "0.3", request_id="1") + assert back["result"]["kind"] == event["kind"] + assert back["result"]["taskId"] == "t1" + + +def test_stream_event_envelope_key_for_status_update(): + v1 = normalize_stream_event(_rpc(V03_STATUS_UPDATE), "1.0", request_id="1") + assert "statusUpdate" in v1["result"] + + +def test_request_params_lowering_is_noop_for_0_3(): + params = {"id": "t1", "historyLength": 5} + assert normalize_request_params(params, "0.3", method="tasks/get") is params + + +def test_request_params_lowering_get_task_to_0_3(): + out = normalize_request_params( + {"id": "t1", "historyLength": 5}, "1.0", method="tasks/get" + ) + assert out["id"] == "t1" + assert out["historyLength"] == 5 + + +def test_request_params_lowering_create_push_notification_config_preserves_task_id(): + out = normalize_request_params( + { + "parent": "tasks/task-1", + "configId": "cfg-1", + "config": {"url": "https://webhook.example.com"}, + }, + "1.0", + method="tasks/pushNotificationConfig/set", + ) + assert out["taskId"] == "task-1" + assert out["pushNotificationConfig"]["url"] == "https://webhook.example.com" + assert out["pushNotificationConfig"]["id"] == "cfg-1" + + +def test_flatten_create_push_notification_drops_redundant_envelope_key(): + from litellm.proxy.a2a.version_convert import ( + _flatten_create_push_notification_params, + ) + + flat = _flatten_create_push_notification_params( + { + "parent": "tasks/task-1", + "config": {"url": "https://chosen.example.com"}, + "pushNotificationConfig": {"url": "https://ignored.example.com"}, + } + ) + assert flat["url"] == "https://chosen.example.com" + assert "pushNotificationConfig" not in flat + assert "config" not in flat + + +def test_request_params_lowering_list_tasks_to_0_3(): + out = normalize_request_params( + { + "contextId": "ctx-1", + "pageSize": 10, + "status": "TASK_STATE_COMPLETED", + }, + "1.0", + method="tasks/list", + ) + assert out["contextId"] == "ctx-1" + assert out["pageSize"] == 10 + assert out["status"] == "completed" + + +@pytest.mark.parametrize( + "proto_status, expected", + [ + ("TASK_STATE_COMPLETED", "completed"), + ("TASK_STATE_INPUT_REQUIRED", "input-required"), + ("TASK_STATE_AUTH_REQUIRED", "auth-required"), + ("TASK_STATE_CANCELED", "canceled"), + ], +) +def test_list_tasks_status_filter_lowers_to_0_3_wire_value(proto_status, expected): + out = normalize_request_params( + {"status": proto_status}, + "1.0", + method="tasks/list", + ) + assert out["status"] == expected + + +def test_list_tasks_unspecified_status_is_dropped(): + out = normalize_request_params( + {"contextId": "ctx-1", "status": "TASK_STATE_UNSPECIFIED"}, + "1.0", + method="tasks/list", + ) + assert "status" not in out + assert out["contextId"] == "ctx-1" + + +@pytest.mark.parametrize( + "method, result", + [ + ( + "message/send", + { + "task": { + "id": "t1", + "contextId": "c1", + "status": {"state": "completed"}, + }, + "vendorExtraField": "x", + }, + ), + ( + "tasks/get", + { + "id": "t1", + "contextId": "c1", + "status": {"state": "completed"}, + "vendorExtraField": "x", + }, + ), + ], +) +def test_lowering_1_0_to_0_3_tolerates_unknown_upstream_fields(method, result): + out = normalize_jsonrpc_response(_rpc(result), "0.3", method=method) + lowered = out["result"] + assert lowered["kind"] == "task" + assert lowered["id"] == "t1" + assert "vendorExtraField" not in lowered + + +def test_stream_event_lowering_1_0_to_0_3_tolerates_unknown_fields(): + event = { + "task": {"id": "t1", "contextId": "c1", "status": {"state": "completed"}}, + "vendorExtraField": "x", + } + out = normalize_stream_event(_rpc(event), "0.3", request_id="1") + lowered = out["result"] + assert lowered["kind"] == "task" + assert lowered["id"] == "t1" + + +def test_list_tasks_result_round_trip_preserves_task_ids(): + rpc = _rpc( + { + "tasks": [ + { + "kind": "task", + "id": "t1", + "contextId": "c1", + "status": {"state": "completed"}, + } + ], + "nextPageToken": "tok", + } + ) + v1 = normalize_jsonrpc_response(rpc, "1.0", method="tasks/list") + assert "kind" not in v1["result"]["tasks"][0] + assert v1["result"]["tasks"][0]["id"] == "t1" + back = normalize_jsonrpc_response(v1, "0.3", method="tasks/list") + assert back["result"]["tasks"][0]["kind"] == "task" + assert back["result"]["tasks"][0]["id"] == "t1" + + +def _extended_card_1_0() -> dict: + return { + "name": "Card", + "description": "d", + "version": "1.0.0", + "supportedInterfaces": [ + { + "url": "https://upstream.example", + "protocolBinding": "JSONRPC", + "protocolVersion": "0.3", + }, + { + "url": "http://internal:9999", + "protocolBinding": "JSONRPC", + "protocolVersion": "0.3", + }, + ], + } + + +def test_agent_card_lowered_to_0_3_drops_additional_interfaces(): + # A 1.0 card with multiple interfaces would lower into a 0.3 card carrying the + # secondary backend URLs in ``additionalInterfaces``; those must be stripped so + # the conversion never re-exposes an upstream backend to A2A clients. + out = normalize_agent_card(_extended_card_1_0(), "0.3") + assert out["url"] == "https://upstream.example" + assert "additionalInterfaces" not in out + assert "supportedInterfaces" not in out + assert "http://internal:9999" not in str(out) + + +def test_agent_card_with_0_3_pin_and_supported_interfaces_is_lowered(): + card = _extended_card_1_0() + card["protocolVersion"] = "0.3" + + out = normalize_agent_card(card, "0.3") + + assert out["protocolVersion"] == "0.3" + assert "supportedInterfaces" not in out + + +def test_agent_card_same_version_passthrough(): + card = _extended_card_1_0() + assert normalize_agent_card(card, "1.0") is card diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 07e878401e0..09a73e076bb 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -967,6 +967,193 @@ async def test_get_extended_agent_card_rewrites_url(): assert body["result"]["name"] == "Test Agent" +@pytest.mark.asyncio +async def test_get_agent_card_uses_proxy_base_url_when_set(monkeypatch): + """Regression: discovery must expose the public proxy URL, not the internal one.""" + from litellm.proxy._types import UserAPIKeyAuth + + monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com") + agent = _make_agent_mock() + agent.agent_card_params["protocolVersion"] = "1.0" + agent.agent_card_params["supportedInterfaces"] = [ + { + "url": "http://old-proxy.example.com/a2a/test-agent", + "protocolBinding": "JSONRPC", + "protocolVersion": "1.0", + } + ] + mock_request = MagicMock() + mock_request.base_url = "http://litellm-internal:4000/" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card + + response = await get_agent_card( + agent_id="test-agent", + request=mock_request, + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["url"] == "https://litellm.example.com/a2a/test-agent" + assert ( + body["supportedInterfaces"][0]["url"] + == "https://litellm.example.com/a2a/test-agent" + ) + + +@pytest.mark.asyncio +async def test_get_agent_card_normalizes_0_3_discovery_card(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + agent.agent_card_params["protocolVersion"] = "0.3" + agent.agent_card_params["supportedInterfaces"] = [ + { + "url": "http://localhost:4000/a2a/test-agent", + "protocolBinding": "JSONRPC", + "protocolVersion": "0.3", + } + ] + mock_request = MagicMock() + mock_request.base_url = "http://localhost:4000/" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card + + response = await get_agent_card( + agent_id="test-agent", + request=mock_request, + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["protocolVersion"] == "0.3" + assert body["url"] == "http://localhost:4000/a2a/test-agent" + assert "supportedInterfaces" not in body + + +@pytest.mark.asyncio +async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header(): + """Regression: 0.3 card normalized to 1.0 must not KeyError on debug log.""" + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + agent.agent_card_params = { + "name": "Test Agent", + "description": "A test agent", + "url": "http://backend-agent:10001", + "version": "1.0.0", + "capabilities": {"streaming": True}, + "skills": [ + {"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]} + ], + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + } + mock_request = MagicMock() + mock_request.base_url = "http://localhost:4000/" + mock_request.headers = {"a2a-version": "1.0"} + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card + + response = await get_agent_card( + agent_id="test-agent", + request=mock_request, + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert "url" not in body + assert body["supportedInterfaces"][0]["url"] == ( + "http://localhost:4000/a2a/test-agent" + ) + + +@pytest.mark.asyncio +async def test_get_extended_agent_card_uses_proxy_base_url_when_set(monkeypatch): + """Regression: proxied extended cards must rewrite url to the public proxy base.""" + from litellm.proxy._types import UserAPIKeyAuth + + monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com") + agent = _make_agent_mock() + mock_request = _make_request_mock("GetExtendedAgentCard", {}) + mock_request.base_url = "http://litellm-internal:4000/" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + upstream_card = { + "name": "Test Agent", + "url": "http://backend-agent:10001", + "description": "A test agent", + } + upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": upstream_card} + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + mock_http_response.raise_for_status = MagicMock() + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = MagicMock() + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["result"]["url"] == "https://litellm.example.com/a2a/test-agent" + + +def test_build_merged_agent_card_uses_proxy_base_url_for_supported_interfaces( + monkeypatch, +): + """Regression: agent create/update must front supportedInterfaces with the public base.""" + from litellm.proxy.agent_endpoints.endpoints import _build_merged_agent_card + + monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com") + mock_request = MagicMock() + mock_request.base_url = "http://litellm-internal:4000/" + + merged = _build_merged_agent_card( + {"name": "My Agent", "url": "http://upstream:8080"}, + agent_id="jenkins_agent", + http_request=mock_request, + ) + + assert merged["supportedInterfaces"][0]["url"] == ( + "https://litellm.example.com/a2a/jenkins_agent" + ) + + @pytest.mark.asyncio async def test_unknown_method_returns_jsonrpc_error(): from litellm.proxy._types import UserAPIKeyAuth @@ -1076,6 +1263,173 @@ async def test_pascal_method_names_normalize_to_wire_format( ) +@pytest.mark.parametrize( + "params", + [ + { + "message": { + "messageId": "msg-1", + "role": "ROLE_USER", + "parts": [{"text": "hello"}], + }, + "configuration": {}, + }, + { + "message": { + "messageId": "msg-2", + "role": "user", + "parts": [{"kind": "text", "text": "hello"}], + }, + }, + ], +) +def test_build_message_send_params_accepts_wire_and_a2a_10(params): + from litellm.proxy.agent_endpoints.a2a_endpoints import _build_message_send_params + + result = _build_message_send_params(params) + assert result.message.role.value == "user" + assert result.message.parts[0].root.text == "hello" + + +def test_build_message_send_params_proto_fallback_ignores_unknown_fields(): + from litellm.proxy.agent_endpoints.a2a_endpoints import _build_message_send_params + + result = _build_message_send_params( + { + "message": { + "messageId": "msg-1", + "role": "ROLE_USER", + "parts": [{"text": "hello"}], + }, + "configuration": {}, + "futureField": "ignored", + } + ) + assert result.message.role.value == "user" + + +@pytest.mark.asyncio +async def test_handle_stream_message_rejects_invalid_params_with_32602(): + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={"message": 12345}, + ) + chunks = [chunk async for chunk in response.body_iterator] + body = "".join( + chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks + ) + payload = json.loads(body.strip()) + assert payload["error"]["code"] == -32602 + assert payload["id"] == "req-1" + + +@pytest.mark.asyncio +async def test_send_message_pascal_case_routes_to_asend_message(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + params = { + "message": { + "messageId": "msg-123", + "role": "ROLE_USER", + "parts": [{"text": "Hello"}], + }, + "configuration": {}, + } + mock_request = _make_request_mock("SendMessage", params) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + captured = {} + + async def capture_asend_message(request, **kwargs): + captured["method"] = request.method + captured["role"] = request.params.message.role.value + response = MagicMock() + response.model_dump.return_value = { + "jsonrpc": "2.0", + "id": request.id, + "result": { + "contextId": "ctx-1", + "kind": "message", + "messageId": "msg-123", + "parts": [{"kind": "text", "text": "Hello"}], + "role": "agent", + }, + } + return response + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch( + "litellm.a2a_protocol.asend_message", + new=AsyncMock(side_effect=capture_asend_message), + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert "error" not in body, f"Got error: {body}" + assert captured["method"] == "message/send" + assert captured["role"] == "user" + assert "message" in body["result"] + assert body["result"]["message"]["role"] == "ROLE_AGENT" + + +def test_normalize_response_wraps_flat_message_result_for_1_0(): + from litellm.proxy.a2a.version_convert import normalize_jsonrpc_response + + wire_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": { + "contextId": "ctx-1", + "kind": "message", + "messageId": "msg-1", + "parts": [{"kind": "text", "text": "hello"}], + "role": "agent", + "taskId": "task-1", + }, + } + formatted = normalize_jsonrpc_response(wire_response, "1.0", method="message/send") + assert "message" in formatted["result"] + assert formatted["result"]["message"]["role"] == "ROLE_AGENT" + assert formatted["result"]["message"]["parts"] == [{"text": "hello"}] + assert "contextId" not in formatted["result"] + + +def test_normalize_response_keeps_wire_format_for_0_3(): + from litellm.proxy.a2a.version_convert import normalize_jsonrpc_response + + wire_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": { + "contextId": "ctx-1", + "kind": "message", + "messageId": "msg-1", + "parts": [{"kind": "text", "text": "hello"}], + "role": "agent", + }, + } + assert ( + normalize_jsonrpc_response(wire_response, "0.3", method="message/send") + is wire_response + ) + + @pytest.mark.asyncio async def test_task_method_upstream_jsonrpc_error_on_http_4xx_is_relayed(): """When upstream returns HTTP 4xx with a JSON-RPC error body, the error body @@ -1560,3 +1914,36 @@ async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers() assert ( posted_headers.get("X-LiteLLM-Team-Id") == "real-team" ), "authenticated team id must not be overridden by forwarded client headers" + + +def _agent(protocol_version): + agent = MagicMock() + agent.agent_card_params = ( + {"protocolVersion": protocol_version} if protocol_version is not None else {} + ) + return agent + + +def _request_with_a2a_header(value): + request = MagicMock() + request.headers = {"a2a-version": value} if value is not None else {} + return request + + +def test_served_version_config_governs_over_header(): + from litellm.proxy.agent_endpoints.a2a_endpoints import _served_version + + # A 0.3-configured agent serves 0.3 even when the client asks for 1.0. + agent = _agent("0.3") + request = _request_with_a2a_header("1.0") + assert _served_version(agent, request) == "0.3" + + # A 1.0-configured agent serves 1.0 even when the client asks for 0.3. + assert _served_version(_agent("1.0"), _request_with_a2a_header("0.3")) == "1.0" + + +def test_served_version_falls_back_to_header_when_unconfigured(): + from litellm.proxy.agent_endpoints.a2a_endpoints import _served_version + + assert _served_version(_agent(None), _request_with_a2a_header("1.0")) == "1.0" + assert _served_version(_agent(None), _request_with_a2a_header(None)) == "0.3" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py new file mode 100644 index 00000000000..069c72af53a --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py @@ -0,0 +1,326 @@ +""" +Near-E2E tests for A2A 0.3/1.0 version routing through the proxy. + +Runs invoke_agent_a2a -> asend_message -> a2a-sdk 1.x -> ASGI mock upstream. +Only proxy auth/registry/pre-call plumbing is patched; version normalization +runs on the real response path. +""" + +from __future__ import annotations + +import json +from contextlib import ExitStack +from typing import Any, AsyncIterator, Dict, List, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from httpx import ASGITransport +from starlette.applications import Starlette +from starlette.responses import JSONResponse, StreamingResponse +from starlette.routing import Route + +pytest.importorskip("a2a.compat.v0_3.types") + +from litellm.proxy._types import UserAPIKeyAuth + +UPSTREAM_BASE = "http://testserver" + +_UPSTREAM_CALLS: List[Dict[str, Any]] = [] + + +def _upstream_card_payload() -> Dict[str, Any]: + return { + "protocolVersion": "0.3", + "name": "mock-agent", + "url": f"{UPSTREAM_BASE}/", + "capabilities": {"streaming": True}, + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "skills": [], + } + + +def _message_result(request_id: Any) -> Dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": request_id, + "result": { + "kind": "message", + "role": "agent", + "messageId": "m-out", + "parts": [{"kind": "text", "text": "pong"}], + }, + } + + +def _sse_stream(request_id: Any) -> AsyncIterator[bytes]: + events = [ + { + "jsonrpc": "2.0", + "id": request_id, + "result": { + "kind": "task", + "id": "t1", + "contextId": "c1", + "status": {"state": "submitted"}, + }, + }, + _message_result(request_id), + ] + + async def _gen() -> AsyncIterator[bytes]: + for event in events: + yield f"data: {json.dumps(event)}\n\n".encode() + + return _gen() + + +async def _serve_upstream_agent_card(request: Any) -> JSONResponse: + return JSONResponse(_upstream_card_payload()) + + +async def _upstream_jsonrpc(request: Any) -> JSONResponse | StreamingResponse: + body = await request.json() + _UPSTREAM_CALLS.append(body) + request_id = body.get("id", "req-1") + method = body.get("method") + + if method == "message/stream": + return StreamingResponse( + _sse_stream(request_id), + media_type="text/event-stream", + ) + + return JSONResponse(_message_result(request_id)) + + +def _build_upstream_app() -> Starlette: + return Starlette( + routes=[ + Route( + "/.well-known/agent-card.json", + _serve_upstream_agent_card, + methods=["GET"], + ), + Route("/.well-known/agent.json", _serve_upstream_agent_card, methods=["GET"]), + Route("/", _upstream_jsonrpc, methods=["POST"]), + ] + ) + + +def _fake_get_async_httpx_client( + llm_provider: Any = None, params: Optional[Dict[str, Any]] = None +) -> MagicMock: + handler = MagicMock() + handler.client = httpx.AsyncClient( + transport=ASGITransport(_build_upstream_app()), + base_url=UPSTREAM_BASE, + ) + return handler + + +def _make_agent(*, protocol_version: str) -> MagicMock: + agent = MagicMock() + agent.agent_id = "test-agent" + agent.agent_name = "test-agent" + agent.agent_card_params = { + "url": f"{UPSTREAM_BASE}/", + "name": "Test Agent", + "protocolVersion": protocol_version, + } + agent.litellm_params = {} + agent.static_headers = None + agent.extra_headers = None + return agent + + +def _make_request( + method: str, + params: Dict[str, Any], + *, + headers: Optional[Dict[str, str]] = None, + request_id: str = "req-1", +) -> MagicMock: + request = MagicMock() + request.headers = headers or {} + request.json = AsyncMock( + return_value={ + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + ) + return request + + +async def _add_proxy_data(data: Dict[str, Any], **_: Any) -> Dict[str, Any]: + data["proxy_server_request"] = { + "url": "http://localhost:4000/a2a/test-agent", + "method": "POST", + "headers": {}, + "body": {}, + } + data.setdefault("metadata", {}) + return data + + +def _proxy_patches(agent: MagicMock) -> List[Any]: + from litellm.proxy.agent_endpoints import a2a_endpoints as a2a_endpoints_mod + + return [ + patch.object(a2a_endpoints_mod, "_get_agent", return_value=agent), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler" + ".AgentRequestHandler.is_agent_allowed", + new=AsyncMock(return_value=True), + ), + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=_add_proxy_data), + ), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), + patch("litellm.proxy.proxy_server.version", "1.0.0"), + patch( + "litellm.a2a_protocol.main.get_async_httpx_client", + side_effect=_fake_get_async_httpx_client, + ), + patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), + ] + + +def _wire_send_params() -> Dict[str, Any]: + return { + "message": { + "role": "user", + "messageId": "m-in", + "parts": [{"kind": "text", "text": "ping"}], + } + } + + +def _a2a10_send_params() -> Dict[str, Any]: + return { + "message": { + "role": "ROLE_USER", + "messageId": "m-in", + "parts": [{"text": "ping"}], + }, + "configuration": {}, + } + + +@pytest.fixture(autouse=True) +def _clear_upstream_calls() -> None: + _UPSTREAM_CALLS.clear() + + +@pytest.mark.asyncio +async def test_proxy_serves_1_0_when_agent_pinned_and_upstream_speaks_03(): + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + agent = _make_agent(protocol_version="1.0") + request = _make_request( + "SendMessage", + _a2a10_send_params(), + headers={"a2a-version": "1.0"}, + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", user_id="test-user", team_id="test-team" + ) + + with ExitStack() as stack: + for item in _proxy_patches(agent): + stack.enter_context(item) + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert "error" not in body, body + assert "message" in body["result"] + assert "kind" not in body["result"] + assert body["result"]["message"]["parts"][0]["text"] == "pong" + assert _UPSTREAM_CALLS, "expected upstream to receive a JSON-RPC call" + assert _UPSTREAM_CALLS[0]["method"] == "message/send" + + +@pytest.mark.asyncio +async def test_proxy_serves_0_3_when_agent_pinned_passthrough(): + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + agent = _make_agent(protocol_version="0.3") + request = _make_request("message/send", _wire_send_params()) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", user_id="test-user", team_id="test-team" + ) + + with ExitStack() as stack: + for item in _proxy_patches(agent): + stack.enter_context(item) + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert "error" not in body, body + assert body["result"]["kind"] == "message" + assert body["result"]["parts"][0]["text"] == "pong" + assert "message" not in body["result"] + + +@pytest.mark.asyncio +async def test_proxy_streaming_serves_1_0_envelopes(): + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + agent = _make_agent(protocol_version="1.0") + request = _make_request( + "SendStreamingMessage", + _a2a10_send_params(), + headers={"a2a-version": "1.0"}, + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", user_id="test-user", team_id="test-team" + ) + + with ExitStack() as stack: + for item in _proxy_patches(agent): + stack.enter_context(item) + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + lines: List[Dict[str, Any]] = [] + async for raw_line in response.body_iterator: + line = ( + raw_line.decode().strip() + if isinstance(raw_line, (bytes, bytearray)) + else str(raw_line).strip() + ) + if line: + lines.append(json.loads(line)) + + assert lines, "expected at least one streamed JSON-RPC event" + message_events = [ + line + for line in lines + if isinstance(line.get("result"), dict) and "message" in line["result"] + ] + assert message_events, f"expected a 1.0 message envelope, got: {lines}" + assert message_events[-1]["result"]["message"]["parts"][0]["text"] == "pong" + assert _UPSTREAM_CALLS, "expected upstream streaming call" + assert _UPSTREAM_CALLS[0]["method"] == "message/stream" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py index 554f98d7209..a51d6f6abcc 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_header_isolation.py @@ -10,7 +10,7 @@ per call; default timeout uses DEFAULT_A2A_AGENT_TIMEOUT). """ import sys -from unittest.mock import AsyncMock, MagicMock, call, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -231,57 +231,95 @@ async def test_each_agent_gets_only_its_own_static_headers(): # --------------------------------------------------------------------------- +def _fake_get_async_httpx_client_factory(captured_calls: list): + """Return a side_effect that records every (params, client) pair.""" + + def _fake_get_async_httpx_client(llm_provider, params, **kwargs): + client = MagicMock() + client.headers = MagicMock() + handler = MagicMock() + handler.client = client + captured_calls.append({"params": params.copy(), "client": client}) + return handler + + return _fake_get_async_httpx_client + + +async def _fake_create_client(base_url, client_config=None, **kwargs): + client = MagicMock() + if client_config is not None: + client._litellm_httpx_client = client_config.httpx_client + return client + + @pytest.mark.asyncio async def test_create_a2a_client_uses_fresh_httpx_client(): """ - Two calls to create_a2a_client with different extra_headers must NOT - share the same underlying httpx.AsyncClient instance. - """ - import httpx + Two calls to create_a2a_client with different extra_headers must produce + distinct underlying httpx clients — preventing header bleed between agents. + The test checks: + 1. get_async_httpx_client was called twice (once per create_a2a_client call). + 2. The two returned A2A clients carry distinct httpx client objects (direct + proof of header isolation, not just cache-key difference). + 3. The cache-key param differs between calls (so the real LRU cache cannot + return the same httpx client even under load). + """ + pytest.importorskip("a2a.client") from litellm.a2a_protocol.main import create_a2a_client - created_clients = [] - - fake_agent_card = MagicMock() - fake_agent_card.name = "test-agent" - - class FakeResolver: - def __init__(self, **kw): - created_clients.append(kw.get("httpx_client")) - - async def get_agent_card(self): - return fake_agent_card - - class FakeA2AClient: - def __init__(self, httpx_client, agent_card): - self._client = httpx_client - self._litellm_agent_card = agent_card + captured_calls: list = [] with ( patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), - patch("litellm.a2a_protocol.main.A2ACardResolver", FakeResolver), - patch("litellm.a2a_protocol.main._A2AClient", FakeA2AClient), + patch( + "litellm.a2a_protocol.main.get_async_httpx_client", + side_effect=_fake_get_async_httpx_client_factory(captured_calls), + ), + patch( + "litellm.a2a_protocol.main.create_client", + new=AsyncMock(side_effect=_fake_create_client), + ), ): - await create_a2a_client( + a2a_client_a = await create_a2a_client( base_url="http://agent-a:9999", extra_headers={"Authorization": "Bearer a"}, ) - await create_a2a_client( + a2a_client_b = await create_a2a_client( base_url="http://agent-b:9999", extra_headers={"Authorization": "Bearer b"}, ) - assert len(created_clients) == 2 - # Must be distinct objects assert ( - created_clients[0] is not created_clients[1] - ), "create_a2a_client reused a cached httpx client — headers will bleed between agents" + len(captured_calls) == 2 + ), "create_a2a_client should call get_async_httpx_client once per invocation" + + # Direct proof: the two A2A clients must carry distinct httpx client objects. + # If they share one, mutating agent-B's Authorization header would bleed into A. + httpx_a = getattr(a2a_client_a, "_litellm_httpx_client", None) + httpx_b = getattr(a2a_client_b, "_litellm_httpx_client", None) + assert httpx_a is not None, "a2a_client_a missing _litellm_httpx_client" + assert httpx_b is not None, "a2a_client_b missing _litellm_httpx_client" + assert httpx_a is not httpx_b, ( + "create_a2a_client returned the same httpx client for two agents with " + "different headers — Authorization header will bleed between agents" + ) + + # Also verify the cache-key param differs so the LRU cache never conflates them. + key_a = captured_calls[0]["params"].get("disable_aiohttp_transport") + key_b = captured_calls[1]["params"].get("disable_aiohttp_transport") + assert key_a is not None, "cache-key param 'disable_aiohttp_transport' missing" + assert key_b is not None, "cache-key param 'disable_aiohttp_transport' missing" + assert key_a != key_b, ( + f"create_a2a_client used the same cache key for two agents with different " + f"headers — headers will bleed: key_a={key_a!r}, key_b={key_b!r}" + ) @pytest.mark.asyncio async def test_create_a2a_client_default_timeout_matches_constant(): """When timeout is omitted, httpx client params must use DEFAULT_A2A_AGENT_TIMEOUT.""" + pytest.importorskip("a2a.client") from litellm.a2a_protocol.main import create_a2a_client captured: dict = {} @@ -293,28 +331,16 @@ async def test_create_a2a_client_default_timeout_matches_constant(): handler.client.headers = MagicMock() return handler - fake_agent_card = MagicMock() - fake_agent_card.name = "test-agent" - - class _FakeResolver: - def __init__(self, **kw): - pass - - async def get_agent_card(self): - return fake_agent_card - - class _FakeA2AClient: - def __init__(self, httpx_client, agent_card): - pass - with ( patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), patch( "litellm.a2a_protocol.main.get_async_httpx_client", side_effect=_capture_get_async_httpx_client, ), - patch("litellm.a2a_protocol.main.A2ACardResolver", _FakeResolver), - patch("litellm.a2a_protocol.main._A2AClient", _FakeA2AClient), + patch( + "litellm.a2a_protocol.main.create_client", + new=AsyncMock(side_effect=_fake_create_client), + ), ): await create_a2a_client(base_url="http://127.0.0.1:9") @@ -324,6 +350,7 @@ async def test_create_a2a_client_default_timeout_matches_constant(): @pytest.mark.asyncio async def test_create_a2a_client_explicit_timeout_overrides_default(): """Explicit timeout= must be passed through to the httpx client params.""" + pytest.importorskip("a2a.client") from litellm.a2a_protocol.main import create_a2a_client captured: dict = {} @@ -335,28 +362,16 @@ async def test_create_a2a_client_explicit_timeout_overrides_default(): handler.client.headers = MagicMock() return handler - fake_agent_card = MagicMock() - fake_agent_card.name = "test-agent" - - class _FakeResolver: - def __init__(self, **kw): - pass - - async def get_agent_card(self): - return fake_agent_card - - class _FakeA2AClient: - def __init__(self, httpx_client, agent_card): - pass - with ( patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), patch( "litellm.a2a_protocol.main.get_async_httpx_client", side_effect=_capture_get_async_httpx_client, ), - patch("litellm.a2a_protocol.main.A2ACardResolver", _FakeResolver), - patch("litellm.a2a_protocol.main._A2AClient", _FakeA2AClient), + patch( + "litellm.a2a_protocol.main.create_client", + new=AsyncMock(side_effect=_fake_create_client), + ), ): await create_a2a_client(base_url="http://127.0.0.1:9", timeout=42.5) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 0cb03023e0f..3740c01b7fc 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -786,3 +786,32 @@ class TestCheckAgentUrlHealth: ) result = await _check_agent_url_health(agent) assert result["healthy"] is True + + +@pytest.mark.parametrize( + "base_url", + ["http://0.0.0.0:4000/", "http://localhost:4000/", "https://api.example.com/"], +) +def test_merged_agent_card_url_has_no_double_slash_without_proxy_base_url( + monkeypatch, base_url +): + """Without PROXY_BASE_URL, request.base_url carries a trailing slash; the merged + card's supportedInterfaces URL must still join cleanly (no `//a2a`).""" + from litellm.proxy.agent_endpoints.endpoints import _build_merged_agent_card + + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + + http_request = MagicMock() + http_request.base_url = base_url + + merged = _build_merged_agent_card( + _sample_agent_card_params(), + agent_id="agent-xyz", + http_request=http_request, + agent_name="Test Agent", + ) + + interface_url = merged["supportedInterfaces"][0]["url"] + assert interface_url == f"{base_url.rstrip('/')}/a2a/agent-xyz" + assert "//a2a" not in interface_url diff --git a/ui/litellm-dashboard/src/components/agents/agent_config.ts b/ui/litellm-dashboard/src/components/agents/agent_config.ts index 14b6729bf93..442dcd48f66 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_config.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_config.ts @@ -6,13 +6,15 @@ export interface FieldConfig { name: string; label: string; - type: "text" | "textarea" | "url" | "switch" | "list"; + type: "text" | "textarea" | "url" | "switch" | "list" | "select"; required?: boolean; tooltip?: string; placeholder?: string; defaultValue?: any; rows?: number; validation?: any[]; + options?: string[]; + helpText?: string; } export interface SectionConfig { @@ -69,9 +71,13 @@ export const AGENT_FORM_CONFIG: { { name: "protocolVersion", label: "Protocol Version", - type: "text", - placeholder: "1.0", + type: "select", + options: ["1.0", "0.3"], defaultValue: "1.0", + tooltip: + "The A2A protocol version LiteLLM serves to clients for this agent. LiteLLM converts the upstream agent's responses to this version, so clients always see the version you pick here regardless of the original agent's version.", + helpText: + "LiteLLM serves this version to clients and converts the upstream agent's responses to match it, regardless of the original agent's version.", }, ], }, diff --git a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx index 7965af6322d..ff103218346 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx @@ -47,9 +47,18 @@ const AgentFormFields: React.FC = ({ showAgentName = true, : undefined } tooltip={field.tooltip} + extra={field.helpText} > {field.type === "textarea" ? ( + ) : field.type === "select" ? ( + ) : ( )} diff --git a/uv.lock b/uv.lock index da44ad25715..f76be69505f 100644 --- a/uv.lock +++ b/uv.lock @@ -20,23 +20,29 @@ members = [ ] constraints = [ { name = "aiohttp", specifier = ">=3.14.1,<4.0" }, + { name = "packaging", specifier = ">=24.0" }, { name = "tornado", specifier = ">=6.5.6" }, ] +overrides = [{ name = "packaging", specifier = ">=24.0" }] [[package]] name = "a2a-sdk" -version = "0.3.24" +version = "1.1.0" source = { registry = "https://pypi.org/simple" } dependencies = [ + { name = "culsans", marker = "python_full_version < '3.13'" }, { name = "google-api-core" }, + { name = "googleapis-common-protos" }, { name = "httpx" }, { name = "httpx-sse" }, + { name = "json-rpc" }, + { name = "packaging" }, { name = "protobuf" }, { name = "pydantic" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ad/76/cefa956fb2d3911cb91552a1da8ce2dbb339f1759cb475e2982f0ae2332b/a2a_sdk-0.3.24.tar.gz", hash = "sha256:3581e6e8a854cd725808f5732f90b7978e661b6d4e227a4755a8f063a3c1599d", size = 255550, upload-time = "2026-02-20T10:05:43.423Z" } +sdist = { url = "https://files.pythonhosted.org/packages/c7/7e/8ac10bbf8b15b16574355f39b17dbdf617a282c27b41c7ff2116e30336df/a2a_sdk-1.1.0.tar.gz", hash = "sha256:e8102dad1b36709dbdc3d19319e38e6dfa3b3a79c30416030eb2d482576be204", size = 375726, upload-time = "2026-05-29T09:34:43.015Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/10/6e/cae5f0caea527b39c0abd7204d9416768764573c76649ca03cc345a372be/a2a_sdk-0.3.24-py3-none-any.whl", hash = "sha256:7b248767096bb55311f57deebf6b767349388d94c1b376c60cb8f6b715e053f6", size = 145752, upload-time = "2026-02-20T10:05:41.729Z" }, + { url = "https://files.pythonhosted.org/packages/d4/ea/3a5b160cfd51c67759b08748051094d9365ceff18127633d0021950c9860/a2a_sdk-1.1.0-py3-none-any.whl", hash = "sha256:d7f5846caf18033d8bf3108b11ec827dd8dd32f867c98848ede0e39474be93be", size = 241886, upload-time = "2026-05-29T09:34:41.484Z" }, ] [[package]] @@ -165,6 +171,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/22/0a/62e7232dc9484fbec112ceb32efb6a624cc7994ec6e2b019286f17c4e8f2/aiohttp-3.14.1-cp313-cp313-win_arm64.whl", hash = "sha256:250d14af67f6b6a1a4a811049b1afa69d61d617fca6bf33149b3ab1a6dbcf7b8", size = 447723, upload-time = "2026-06-07T21:08:00.154Z" }, ] +[[package]] +name = "aiologic" +version = "0.17.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "sniffio", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "wrapt", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/53/a7/809482759f40079f4c4328c7318bf569ae25d457f5017aad30a1b9aafedc/aiologic-0.17.0.tar.gz", hash = "sha256:65aa058e858c94cd208badb188e7f00b54dcabb3ba85b34f794db98074d108b9", size = 251625, upload-time = "2026-06-14T12:24:35.367Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9e/6b/5f75d6194b597ac32bbdbb7b524a28fb1fa98bd0ddcefce94b313a818cc0/aiologic-0.17.0-py3-none-any.whl", hash = "sha256:1bf4d3e4314df2bcb06a9e696417204e206ab50e10ec98d28d157e2e57634f74", size = 161084, upload-time = "2026-06-14T12:24:34.146Z" }, +] + [[package]] name = "aiosignal" version = "1.4.0" @@ -1191,6 +1211,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5d/8c/ce3823c06c2804f194f9e64f0d67fa3f4094a39f2bb1a990cd03603af8fc/cryptography-48.0.1-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:6184ca7b174f28d7c703f1290d4b297217c45355f77a98f67e9b7f14549ac54a", size = 3742204, upload-time = "2026-06-09T22:31:34.773Z" }, ] +[[package]] +name = "culsans" +version = "0.11.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "aiologic", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d9/e3/49afa1bc180e0d28008ec6bcdf82a4072d1c7a41032b5b759b60814ca4b0/culsans-0.11.0.tar.gz", hash = "sha256:0b43d0d05dce6106293d114c86e3fb4bfc63088cfe8ff08ed3fe36891447fe33", size = 107546, upload-time = "2025-12-31T23:15:38.196Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e0/5d/9fb19fb38f6d6120422064279ea5532e22b84aa2be8831d49607194feda3/culsans-0.11.0-py3-none-any.whl", hash = "sha256:278d118f63fc75b9db11b664b436a1b83cc30d9577127848ba41420e66eb5a47", size = 21811, upload-time = "2025-12-31T23:15:37.189Z" }, +] + [[package]] name = "cycler" version = "0.12.1" @@ -2746,6 +2779,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7b/91/984aca2ec129e2757d1e4e3c81c3fcda9d0f85b74670a094cc443d9ee949/joblib-1.5.3-py3-none-any.whl", hash = "sha256:5fc3c5039fc5ca8c0276333a188bbd59d6b7ab37fe6632daa76bc7f9ec18e713", size = 309071, upload-time = "2025-12-15T08:41:44.973Z" }, ] +[[package]] +name = "json-rpc" +version = "1.15.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/6d/9e/59f4a5b7855ced7346ebf40a2e9a8942863f644378d956f68bcef2c88b90/json-rpc-1.15.0.tar.gz", hash = "sha256:e6441d56c1dcd54241c937d0a2dcd193bdf0bdc539b5316524713f554b7f85b9", size = 28854, upload-time = "2023-06-11T09:45:49.078Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/94/9e/820c4b086ad01ba7d77369fb8b11470a01fac9b4977f02e18659cf378b6b/json_rpc-1.15.0-py2.py3-none-any.whl", hash = "sha256:4a4668bbbe7116feb4abbd0f54e64a4adcf4b8f648f19ffa0848ad0f6606a9bf", size = 39450, upload-time = "2023-06-11T09:45:47.136Z" }, +] + [[package]] name = "jsonlines" version = "4.0.0" @@ -3428,7 +3470,7 @@ proxy-dev = [ [package.metadata] requires-dist = [ - { name = "a2a-sdk", marker = "extra == 'extra-proxy'", specifier = ">=0.3.24,<1.0" }, + { name = "a2a-sdk", marker = "extra == 'extra-proxy'", specifier = ">=1.1.0,<2.0" }, { name = "aiohttp", specifier = ">=3.10,<4.0" }, { name = "anthropic", extras = ["vertex"], marker = "extra == 'proxy-runtime'", specifier = ">=0.84.0,<1.0" }, { name = "apscheduler", marker = "extra == 'proxy'", specifier = ">=3.11.2,<4.0" }, @@ -3584,7 +3626,7 @@ healthcheck = [ { name = "pyyaml", specifier = "==6.0.3" }, ] proxy-dev = [ - { name = "a2a-sdk", specifier = "==0.3.24" }, + { name = "a2a-sdk", specifier = "==1.1.0" }, { name = "azure-identity", specifier = "==1.25.2" }, { name = "hypercorn", specifier = "==0.17.3" }, { name = "opentelemetry-api", specifier = "==1.28.0" }, @@ -5144,11 +5186,11 @@ wheels = [ [[package]] name = "packaging" -version = "23.2" +version = "26.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/fb/2b/9b9c33ffed44ee921d0967086d653047286054117d584f1b1a7c22ceaf7b/packaging-23.2.tar.gz", hash = "sha256:048fb0e9405036518eaaf48a55953c750c11e1a1b68e0dd1a9d62ed0c092cfc5", size = 146714, upload-time = "2023-10-01T13:50:05.279Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d7/f1/e7a6dd94a8d4a5626c03e4e99c87f241ba9e350cd9e6d75123f992427270/packaging-26.2.tar.gz", hash = "sha256:ff452ff5a3e828ce110190feff1178bb1f2ea2281fa2075aadb987c2fb221661", size = 228134, upload-time = "2026-04-24T20:15:23.917Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ec/1a/610693ac4ee14fcdf2d9bf3c493370e4f2ef7ae2e19217d7a237ff42367d/packaging-23.2-py3-none-any.whl", hash = "sha256:8c491190033a9af7e1d931d0b5dacc2ef47509b34dd0de67ed209b5203fc88c7", size = 53011, upload-time = "2023-10-01T13:50:03.745Z" }, + { url = "https://files.pythonhosted.org/packages/df/b2/87e62e8c3e2f4b32e5fe99e0b86d576da1312593b39f47d8ceef365e95ed/packaging-26.2-py3-none-any.whl", hash = "sha256:5fc45236b9446107ff2415ce77c807cee2862cb6fac22b8a73826d0693b0980e", size = 100195, upload-time = "2026-04-24T20:15:22.081Z" }, ] [[package]] From d7654d07ab949e6fc05fe2da046ab1c67735cf40 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Mon, 29 Jun 2026 20:14:22 +0200 Subject: [PATCH 12/37] feat(proxy): add AES-256-GCM at-rest credential encryption with versioned format and re-encryption migration (#31215) * feat(proxy): add AES-256-GCM at-rest credential encryption with versioned format and re-encryption migration * test(proxy): add behavior scenarios for credential migration endpoints * fix(proxy): scan covered tables in encryption check, fix CI lint and route types * fix(proxy): migrate callback_settings credentials, clear CI lint/recursion gates, add encryption endpoint+CLI tests * fix(proxy): correct dry-run/real-run migrated vs residual-legacy counters in config and SSO walkers * fix(proxy): make callback-vars residual detection gate-independent in encryption check --- .../proxy/client/cli/commands/encryption.py | 60 ++ litellm/proxy/client/cli/main.py | 3 + .../common_utils/encrypt_decrypt_utils.py | 87 ++- .../credential_migration.py | 702 ++++++++++++++++++ .../key_management_endpoints.py | 80 ++ .../test_credential_migration_endpoint.py | 71 ++ .../client/cli/test_encryption_commands.py | 71 ++ .../test_encrypt_decrypt_utils.py | 148 ++++ .../test_credential_migration.py | 516 +++++++++++++ .../test_encryption_endpoints.py | 100 +++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 98 +++ 11 files changed, 1935 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/client/cli/commands/encryption.py create mode 100644 litellm/proxy/management_endpoints/credential_migration.py create mode 100644 tests/proxy_behavior/management/test_credential_migration_endpoint.py create mode 100644 tests/test_litellm/proxy/client/cli/test_encryption_commands.py create mode 100644 tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py create mode 100644 tests/test_litellm/proxy/management_endpoints/test_credential_migration.py create mode 100644 tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py diff --git a/litellm/proxy/client/cli/commands/encryption.py b/litellm/proxy/client/cli/commands/encryption.py new file mode 100644 index 00000000000..f67c9746fa9 --- /dev/null +++ b/litellm/proxy/client/cli/commands/encryption.py @@ -0,0 +1,60 @@ +"""CLI commands for the at-rest credential encryption migration.""" + +import click +import rich + +from ...http_client import HTTPClient + + +@click.group() +def encryption(): + """Migrate at-rest credentials to AES-256-GCM and attest residual state.""" + pass + + +@encryption.command(name="migrate") +@click.option( + "--check", + "check_only", + is_flag=True, + default=False, + help="Read-only residual scan (no writes). Reports legacy values remaining.", +) +@click.option( + "--dry-run", + is_flag=True, + default=False, + help="Run the full migration walkers without writing any changes.", +) +@click.pass_context +def migrate(ctx: click.Context, check_only: bool, dry_run: bool): + """Re-encrypt at-rest credentials into the AES-256-GCM (v2:gcm:) format. + + Requires the proxy to be started with + ``general_settings.encryption_algorithm: aes-256-gcm``. Idempotent and + resumable — safe to re-run after an interruption. + + Examples: + litellm-proxy encryption migrate --check # attestation scan, no writes + litellm-proxy encryption migrate # perform the migration + """ + client = HTTPClient(ctx.obj["base_url"], ctx.obj["api_key"]) + + if check_only: + response = client.request("GET", "/credentials/migrate-encryption/check") + else: + response = client.request( + "POST", + "/credentials/migrate-encryption", + json={}, + params={"dry_run": "true"} if dry_run else None, + ) + + rich.print_json(data=response) + + report = response.get("report", {}) if isinstance(response, dict) else {} + residual = report.get("residual_legacy") + if residual is not None and residual > 0: + rich.print(f"[yellow]Residual legacy values remaining: {residual}[/yellow]") + elif residual == 0: + rich.print("[green]No legacy values remaining (residual_legacy == 0).[/green]") diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index b8c483f4b08..43b64aebd3b 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -11,6 +11,7 @@ from .commands.agents import agent_commands from .commands.auth import get_stored_api_key, login, logout, whoami from .commands.chat import chat from .commands.credentials import credentials +from .commands.encryption import encryption from .commands.http import http from .commands.keys import keys @@ -103,6 +104,8 @@ cli.add_command(whoami) cli.add_command(models) # Add the credentials command group cli.add_command(credentials) +# Add the encryption migration command group +cli.add_command(encryption) # Add the chat command group cli.add_command(chat) # Add the http command group diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index 6be56de1260..8599b3ace7f 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -1,9 +1,24 @@ import base64 import os -from typing import Literal, Optional +from typing import Literal, Optional, cast from litellm._logging import verbose_proxy_logger +# Versioned ciphertext marker for AES-256-GCM values. +# Format: "v2:gcm:" + base64url(nonce(12) || ciphertext || tag(16)). +# Legacy XSalsa20-Poly1305 (nacl) values carry no marker; the colon in the +# prefix can never appear in base64url(nacl output), so the prefix check is an +# unambiguous discriminator between the two formats on read. +_V2_GCM_PREFIX = "v2:gcm:" + +# general_settings key selecting the at-rest encryption algorithm for new writes. +# Default preserves the legacy algorithm so existing deployments are byte-for-byte +# unchanged until they explicitly opt in. Decrypt is always format-detecting, so +# flipping this flag forward (or back) never strands previously-written data. +_ENCRYPTION_ALGORITHM_SETTING = "encryption_algorithm" +_ALGO_AES_GCM = "aes-256-gcm" +_ALGO_XSALSA20 = "xsalsa20-poly1305" + def _get_salt_key(): from litellm.proxy.proxy_server import master_key @@ -16,11 +31,76 @@ def _get_salt_key(): return salt_key +def _get_encryption_algorithm() -> str: + """ + Resolve the configured at-rest encryption algorithm for *new writes*. + + Read from ``general_settings.encryption_algorithm`` at write time. Defaults to + the legacy XSalsa20-Poly1305 algorithm so deployments that have not opted in + keep producing byte-for-byte identical ciphertext. + """ + try: + from litellm.proxy.proxy_server import general_settings + + algo = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20) + except Exception: + # general_settings may not be importable in some contexts (e.g. SDK-only + # use of these helpers). Fall back to the legacy algorithm. + return _ALGO_XSALSA20 + + if isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM: + return _ALGO_AES_GCM + return _ALGO_XSALSA20 + + +def _derive_key(signing_key: str) -> bytes: + """Derive a 32-byte key from the salt/master key (shared by both algorithms). + + Known limitation: this is a single-pass, unsalted ``SHA-256`` of the key, not + a dedicated KDF (HKDF/PBKDF2). It is the *same* derivation the legacy nacl + path already uses, so the AES path introduces no new weakness and stays + interoperable with existing key sourcing; AES-256-GCM's per-value 12-byte + random nonce gives the unique (key, nonce) pairs GCM requires. Moving both + algorithms to HKDF-SHA256 would be more defensible in an audit but is a + separate, coordinated change (it must re-derive or re-encrypt existing data). + """ + import hashlib + + return hashlib.sha256(signing_key.encode()).digest() + + +def _encrypt_aes_gcm(value: str, signing_key: str) -> str: + """Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string.""" + from cryptography.hazmat.primitives.ciphers.aead import AESGCM + + nonce = os.urandom(12) + # AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that. + blob = AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), None) + return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8") + + +def _decrypt_aes_gcm(value: str, signing_key: str) -> str: + """Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`.""" + from cryptography.hazmat.primitives.ciphers.aead import AESGCM + + raw = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :]) + # An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a + # short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be + # swallowed by decrypt_value_helper (returns None/original), same as legacy. + nonce, blob = raw[:12], raw[12:] + return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8") + + def encrypt_value_helper(value: str, new_encryption_key: Optional[str] = None): signing_key = new_encryption_key or _get_salt_key() try: if isinstance(value, str): + if _get_encryption_algorithm() == _ALGO_AES_GCM: + # AES path: the v2:gcm: output is already a base64url string, so it + # is returned directly with no extra base64 wrapper. + return _encrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) + encrypted_value = encrypt_value(value=value, signing_key=signing_key) # type: ignore # Use urlsafe_b64encode for URL-safe base64 encoding (replaces + with - and / with _) encrypted_value = base64.urlsafe_b64encode(encrypted_value).decode("utf-8") @@ -46,6 +126,11 @@ def decrypt_value_helper( try: if isinstance(value, str): + # Versioned AES-256-GCM values are detected before any base64 decode. + # The prefix is the algorithm tag the legacy nacl format never carried. + if value.startswith(_V2_GCM_PREFIX): + return _decrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) + # Try URL-safe base64 decoding first (new format) # Fall back to standard base64 decoding for backwards compatibility (old format) try: diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py new file mode 100644 index 00000000000..4d51295f8dc --- /dev/null +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -0,0 +1,702 @@ +""" +At-rest credential re-encryption migration. + +Switches every encrypted-at-rest value from the legacy XSalsa20-Poly1305 (nacl) +format to the versioned AES-256-GCM (``v2:gcm:``) format produced by +``encrypt_decrypt_utils`` when ``general_settings.encryption_algorithm`` is set to +``aes-256-gcm``. + +Design properties (see case 2026-06-24 fix plan): + +* **Same key, new algorithm.** The migration does not change the encryption key; + it re-encrypts existing ciphertext under the same derived key but in the new + AES format. This is achieved by decrypting with the format-detecting reader and + re-encrypting through ``encrypt_value_helper`` with the AES gate enabled. +* **Idempotent.** A value already carrying the ``v2:gcm:`` prefix is recognised + and left untouched, so re-running the migration is a no-op on migrated rows. +* **Resumable.** Walkers commit per row (or per small table), so an interrupted + run leaves a clean mixed state that a re-run completes. +* **Skip-on-undecryptable.** A value that cannot be decrypted is never + overwritten — corrupt rows are preserved and reported, never destroyed. +* **Attestable.** :func:`check_encryption` is a read-only scan that classifies + every value as ``migrated`` / ``legacy`` / ``plaintext`` / ``undecryptable``. + A residual ``legacy == 0`` is the compliance attestation. + +Coverage. The covered tables (model table, credentials table, MCP credential/env +tables, config ``environment_variables``) already have a re-encryption path in +``_rotate_master_key``; this module delegates to it in *same-key* mode and adds +walkers for the locations that had no rotation path: team / verification-token +``callback_vars`` metadata, the ``vantage_settings`` / ``cloudzero_settings`` +config rows, and the SSO config table. +""" + +import json +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Literal, cast + +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.utils import PrismaClient +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + _ALGO_AES_GCM, + _ENCRYPTION_ALGORITHM_SETTING, + _V2_GCM_PREFIX, + _get_salt_key, + decrypt_value_helper, + encrypt_value_helper, +) + +ValueClass = Literal["migrated", "legacy", "plaintext", "undecryptable", "not-a-string"] + + +@dataclass +class LocationReport: + """Per-location counters for one migration / check pass.""" + + location: str + scanned: int = 0 + migrated: int = 0 # values rewritten to v2 this run + already_v2: int = 0 # values already migrated (skipped) + plaintext: int = 0 # legacy-plaintext values (no ciphertext to migrate) + undecryptable: int = 0 # could not decrypt — preserved, not overwritten + + # Used by --check (read-only classification): + legacy: int = 0 # nacl ciphertext still awaiting migration + + def as_dict(self) -> dict[str, int]: + return { + "scanned": self.scanned, + "migrated": self.migrated, + "already_v2": self.already_v2, + "plaintext": self.plaintext, + "undecryptable": self.undecryptable, + "legacy": self.legacy, + } + + +@dataclass +class MigrationReport: + """Aggregate report across all locations.""" + + locations: list[LocationReport] = field(default_factory=list) + + def add(self, report: LocationReport) -> None: + self.locations.append(report) + + @property + def residual_legacy(self) -> int: + """Total legacy ciphertext still un-migrated (the TRO attestation number).""" + return sum(loc.legacy for loc in self.locations) + + @property + def total_undecryptable(self) -> int: + return sum(loc.undecryptable for loc in self.locations) + + def as_dict(self) -> dict[str, object]: + return { + "residual_legacy": self.residual_legacy, + "total_undecryptable": self.total_undecryptable, + "locations": {loc.location: loc.as_dict() for loc in self.locations}, + } + + +# --------------------------------------------------------------------------- +# Pure engine — no DB I/O, fully unit-testable. +# --------------------------------------------------------------------------- + + +def is_migrated(value: object) -> bool: + """True if ``value`` is already an AES-256-GCM (``v2:gcm:``) ciphertext.""" + return isinstance(value, str) and value.startswith(_V2_GCM_PREFIX) + + +def classify_value(value: object, key: str = "scan") -> ValueClass: + """Classify a stored value for the residual scanner. + + * ``not-a-string`` — not a string (numbers/bools/None left as-is on disk). + * ``migrated`` — carries the ``v2:gcm:`` prefix. + * ``legacy`` — decrypts under the legacy nacl reader (still needs migrating). + * ``plaintext`` — a non-empty string that does not decrypt and is not v2; + treated as legacy plaintext (nothing to migrate). + * ``undecryptable`` — reserved for callers that already know a value is + ciphertext but cannot decrypt it; ``classify_value`` itself cannot tell a + corrupt ciphertext from plaintext, so it returns ``plaintext`` for both. + """ + if not isinstance(value, str): + return "not-a-string" + if value == "": + return "plaintext" + if value.startswith(_V2_GCM_PREFIX): + return "migrated" + decrypted = decrypt_value_helper( + value=value, key=key, exception_type="debug", return_original_value=False + ) + if decrypted is None: + # Did not decrypt under nacl and has no v2 marker: legacy plaintext. + return "plaintext" + return "legacy" + + +def reencrypt_value(value: object, key: str = "migrate") -> object: + """Re-encrypt a single stored string into the configured (AES) format. + + Returns the value unchanged if it is not a string, is already ``v2:``, or + cannot be decrypted (skip-on-undecryptable). Otherwise decrypts under the + format-detecting reader and re-encrypts through ``encrypt_value_helper`` + (which writes AES when the gate is on). + """ + if not isinstance(value, str) or value == "": + return value + if value.startswith(_V2_GCM_PREFIX): + return value # idempotent: already migrated + decrypted = decrypt_value_helper( + value=value, key=key, exception_type="debug", return_original_value=False + ) + if decrypted is None: + # Either legacy plaintext (no ciphertext to migrate) or corrupt. Either + # way, do not overwrite — preserve the value as stored. + return value + return encrypt_value_helper(decrypted) + + +def reencrypt_selective_dict( + data: dict[str, object], sensitive_keys: list[str] +) -> dict[str, object]: + """Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted. + + Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is. + Null/missing fields are skipped. + """ + out = dict(data) + for k in sensitive_keys: + v = out.get(k) + if v is None: + continue + out[k] = reencrypt_value(v, key=k) + return out + + +def _assert_aes_gate_enabled() -> None: + """Fail fast if the AES algorithm gate is not enabled. + + Running the migration with the gate off would decrypt then re-encrypt right + back into the legacy format — a no-op that silently fails the migration. + """ + from litellm.proxy.proxy_server import general_settings + + algo = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING) + if not (isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM): + raise RuntimeError( + "Encryption migration requires general_settings.encryption_algorithm: " + f"'{_ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " + "so re-encrypted values are written in the AES-256-GCM format." + ) + + +# --------------------------------------------------------------------------- +# Walkers for the locations with no pre-existing rotation path. +# Each walker delegates the structural transform to the existing, tested helper +# for that table and only adds the per-row re-encrypt + commit + counters. +# --------------------------------------------------------------------------- + + +async def _migrate_config_settings_row( + prisma_client: object, + param_name: str, + sensitive_fields: list[str], + dry_run: bool, +) -> LocationReport: + """Migrate a single ``LiteLLM_Config`` row whose ``param_value`` is a JSON + dict with selected sensitive fields (vantage_settings / cloudzero_settings). + """ + report = LocationReport(location=param_name) + record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": param_name} + ) + if record is None or record.param_value is None: + return report + + settings = record.param_value + if isinstance(settings, str): + settings = json.loads(settings) + if not isinstance(settings, dict): + return report + + changed = False + for fld in sensitive_fields: + v = settings.get(fld) + if v is None: + continue + report.scanned += 1 + cls = classify_value(v, key=fld) + if cls == "migrated": + report.already_v2 += 1 + continue + if cls == "legacy": + if dry_run: + # Residual: would migrate, but a dry run writes nothing, so it + # stays legacy for the attestation (never counted as migrated). + report.legacy += 1 + continue + new_v = reencrypt_value(v, key=fld) + if new_v != v: + settings[fld] = new_v + report.migrated += 1 + changed = True + else: + # Defensive: a legacy value that did not re-encrypt is still + # residual, not migrated. + report.legacy += 1 + else: # plaintext / not-a-string — nothing to migrate + report.plaintext += 1 + + if changed and not dry_run: + await prisma_client.db.litellm_config.update( + where={"param_name": param_name}, + data={"param_value": json.dumps(settings)}, + ) + return report + + +async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationReport: + """Migrate the ``LiteLLM_SSOConfig`` row. All non-null fields are encrypted + (via the same ``_encrypt_env_variables`` path used on save), so we re-encrypt + every present string field. + """ + report = LocationReport(location="sso_config") + record = await prisma_client.db.litellm_ssoconfig.find_unique( + where={"id": "sso_config"} + ) + if record is None or record.sso_settings is None: + return report + + settings = record.sso_settings + if isinstance(settings, str): + settings = json.loads(settings) + if not isinstance(settings, dict): + return report + + new_settings = dict(settings) + changed = False + for fld, v in settings.items(): + if not isinstance(v, str) or v == "": + continue + report.scanned += 1 + cls = classify_value(v, key=fld) + if cls == "migrated": + report.already_v2 += 1 + continue + if cls == "legacy": + if dry_run: + # Residual: would migrate, but a dry run writes nothing, so it + # stays legacy for the attestation (never counted as migrated). + report.legacy += 1 + continue + new_v = reencrypt_value(v, key=fld) + if new_v != v: + new_settings[fld] = new_v + report.migrated += 1 + changed = True + else: + # Defensive: a legacy value that did not re-encrypt is still + # residual, not migrated. + report.legacy += 1 + else: + report.plaintext += 1 + + if changed and not dry_run: + await prisma_client.db.litellm_ssoconfig.update( + where={"id": "sso_config"}, + data={"sso_settings": json.dumps(new_settings)}, + ) + return report + + +async def _migrate_callback_vars_table( + prisma_client: object, + table_name: Literal["team", "verification_token"], + dry_run: bool, +) -> LocationReport: + """Migrate callback-var credentials on the team or verification-token table. + + Covers both shapes the ``decrypt_callback_vars`` / ``encrypt_callback_vars`` + transforms understand: ``metadata.logging[*].callback_vars.`` and + the top-level ``metadata.callback_settings.callback_vars.``. Reuses + those proven transforms (selective, prefix-marked; legacy plaintext is left + alone until re-encrypted). + """ + from litellm.proxy.common_utils.callback_utils import ( + decrypt_callback_vars, + encrypt_callback_vars, + ) + + report = LocationReport(location=f"{table_name}.callback_vars") + + if table_name == "team": + table = prisma_client.db.litellm_teamtable + pk = "team_id" + else: + table = prisma_client.db.litellm_verificationtoken + pk = "token" + + rows = await table.find_many() + for row in rows or []: + metadata = getattr(row, "metadata", None) + if not isinstance(metadata, dict) or ( + "logging" not in metadata and "callback_settings" not in metadata + ): + continue + + # Classify every callback-var value directly (strip the litellm_enc:: + # marker, then prefix/decrypt-classify), exactly like the covered-table + # scanner. Detecting legacy this way is independent of the AES gate, so + # the check_encryption (dry-run) attestation is correct even when run + # before the gate is enabled -- a re-encrypt-delta heuristic would read + # zero residual here with the gate off. + row_legacy = 0 + for cvs in _iter_callback_var_dicts(metadata): + for v in cvs.values(): + report.scanned += 1 + cls = _classify_callback_value(v) + if cls == "migrated": + report.already_v2 += 1 + elif cls == "legacy": + row_legacy += 1 + else: # plaintext / not-a-string + report.plaintext += 1 + + if row_legacy == 0: + continue # no legacy ciphertext in this row + + if dry_run: + # Residual for the attestation; a dry run writes nothing. + report.legacy += row_legacy + continue + + # Real run: re-encrypt the legacy ciphertext to AES via the proven + # selective transforms and persist. Never drop a row on failure. + try: + re_encrypted = encrypt_callback_vars(decrypt_callback_vars(metadata)) + except Exception as e: # pragma: no cover - defensive; never drop a row + verbose_proxy_logger.warning( + "Skipping %s row %s callback_vars (transform failed): %s", + table_name, + getattr(row, pk, "?"), + str(e), + ) + report.undecryptable += row_legacy + continue + report.migrated += row_legacy + await table.update( + where={pk: getattr(row, pk)}, + data={"metadata": json.dumps(re_encrypted)}, + ) + + return report + + +def _iter_callback_var_dicts(metadata: dict[str, object]): + """Yield each ``callback_vars`` dict in a metadata structure. + + Mirrors ``_transform_callback_vars``: credentials live both under + ``logging[*].callback_vars`` and under the top-level + ``callback_settings.callback_vars``. Counting only the former would let the + walker report success while leaving ``callback_settings`` secrets in legacy + format at rest. + """ + for entry in metadata.get("logging", []) or []: + if isinstance(entry, dict): + cvs = entry.get("callback_vars") + if isinstance(cvs, dict): + yield cvs + callback_settings = metadata.get("callback_settings") + if isinstance(callback_settings, dict): + cvs = callback_settings.get("callback_vars") + if isinstance(cvs, dict): + yield cvs + + +def _classify_callback_value(value: object) -> ValueClass: + """Classify one stored callback-var value, independent of the AES gate. + + Encrypted callback vars carry the ``litellm_enc::`` marker in front of the + ciphertext; strip it, then classify the inner value the same way the + covered-table scanner does (``v2:gcm:`` prefix -> migrated, nacl-decryptable + -> legacy, otherwise plaintext). Detecting legacy by decrypt rather than by a + re-encrypt delta is what makes the ``check_encryption`` attestation correct + even when run with the AES write gate off. + """ + from litellm.proxy.common_utils.callback_utils import ( + _CALLBACK_VAR_ENCRYPTED_PREFIX, + ) + + if not isinstance(value, str): + return "not-a-string" + inner = value + if inner.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX): + inner = inner[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :] + return classify_value(inner, key="callback") + + +# --------------------------------------------------------------------------- +# Read-only scanner for the rotation-covered tables. +# +# ``_rotate_master_key`` re-encrypts these tables but returns no counts, so on +# its own it can neither attest residual legacy nor report how many rows it +# migrated. This scanner reads (never writes) the same encrypted columns the +# rotation path touches and classifies every value, giving both the attestation +# coverage and the pre/post counts the rotation path can't supply itself. +# --------------------------------------------------------------------------- + +# (location, prisma db attribute, JSON columns to walk, scalar string columns). +_COVERED_TABLE_SPECS = [ + ("model_table", "litellm_proxymodeltable", ("litellm_params",), ()), + ("credentials", "litellm_credentialstable", ("credential_values",), ()), + ("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars"), ()), + ("mcp_user_credentials", "litellm_mcpusercredentials", (), ("credential_b64",)), + ("mcp_user_env_vars", "litellm_mcpuserenvvars", (), ("values_b64",)), +] + + +def _iter_encrypted_strings(obj: object): + """Yield every string leaf in a nested dict/list/scalar structure. + + Iterative (explicit stack) on purpose: recursion here is banned by the + code-quality recursive-function detector (unbounded nesting has caused CPU + spikes in the past), and an explicit stack walks arbitrary depth safely. + """ + stack: list[object] = [obj] + while stack: + cur = stack.pop() + if isinstance(cur, str): + yield cur + elif isinstance(cur, dict): + stack.extend(cur.values()) + elif isinstance(cur, list): + stack.extend(cur) + + +def _classify_into_report(report: LocationReport, value: str) -> None: + """Classify one stored string and bump the matching read-only counter. + + Only genuine nacl ciphertext lands in ``legacy``; non-secret strings (model + names, base URLs, …) do not decrypt and fall through to ``plaintext``, so + over-scanning a column is harmless to the residual count. + """ + report.scanned += 1 + cls = classify_value(value, key="scan") + if cls == "migrated": + report.already_v2 += 1 + elif cls == "legacy": + report.legacy += 1 + else: # plaintext / not-a-string + report.plaintext += 1 + + +async def _scan_one_table( + prisma_client: object, + location: str, + db_attr: str, + json_columns: tuple, + scalar_columns: tuple, +) -> LocationReport: + report = LocationReport(location=location) + table = getattr(prisma_client.db, db_attr, None) + if table is None: + return report + try: + rows = await table.find_many() + except Exception as e: # pragma: no cover - table absent / not migrated + verbose_proxy_logger.debug("scan: %s unavailable: %s", location, str(e)) + return report + for row in rows or []: + for col in json_columns: + raw = getattr(row, col, None) + if raw is None: + continue + if isinstance(raw, str): + try: + raw = json.loads(raw) + except (ValueError, TypeError): + pass + for s in _iter_encrypted_strings(raw): + _classify_into_report(report, s) + for col in scalar_columns: + v = getattr(row, col, None) + if isinstance(v, str): + _classify_into_report(report, v) + return report + + +async def _scan_config_env_vars(prisma_client: object) -> LocationReport: + """Scan the ``environment_variables`` config row (``param_value`` dict).""" + report = LocationReport(location="config_environment_variables") + try: + record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "environment_variables"} + ) + except Exception as e: # pragma: no cover - defensive + verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e)) + return report + if record is None or record.param_value is None: + return report + value = record.param_value + if isinstance(value, str): + try: + value = json.loads(value) + except (ValueError, TypeError): + value = {} + for s in _iter_encrypted_strings(value): + _classify_into_report(report, s) + return report + + +async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]: + """Read-only classification of every rotation-covered table. No writes.""" + reports: list[LocationReport] = [] + for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS: + reports.append( + await _scan_one_table( + prisma_client, location, db_attr, json_cols, scalar_cols + ) + ) + reports.append(await _scan_config_env_vars(prisma_client)) + return reports + + +# --------------------------------------------------------------------------- +# Orchestrator +# --------------------------------------------------------------------------- + +# vantage_settings / cloudzero_settings sensitive fields (see *_endpoints.py). +_VANTAGE_SENSITIVE = ["api_key", "integration_token"] +_CLOUDZERO_SENSITIVE = ["api_key"] + + +async def _migrate_covered_tables( + prisma_client: object, user_api_key_dict: object +) -> list[LocationReport]: + """Re-encrypt the tables already covered by ``_rotate_master_key`` (model + table, credentials, MCP credential/env tables, config environment_variables) + by running that orchestrator in *same-key* mode. With the AES gate on, the + re-encrypt writes land in ``v2:`` format. + + ``_rotate_master_key`` returns no counts, so we bracket it with read-only + scans: the pre-scan's legacy total minus the post-scan's gives the number + actually migrated per location, and the post-scan supplies the residual / + already-v2 / scanned figures. Returns one report per covered location. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _rotate_master_key, + ) + + pre = {r.location: r for r in await _scan_covered_tables(prisma_client)} + + current_key = _get_salt_key() + if current_key is None: + raise RuntimeError( + "Cannot migrate covered tables: no salt key / master key is set. " + "Set LITELLM_SALT_KEY before migrating." + ) + await _rotate_master_key( + prisma_client=cast("PrismaClient", prisma_client), + user_api_key_dict=cast("UserAPIKeyAuth", user_api_key_dict), + current_master_key=current_key, + new_master_key=current_key, # same key, algorithm-only switch + ) + + post = await _scan_covered_tables(prisma_client) + for post_report in post: + pre_report = pre.get(post_report.location) + pre_legacy = pre_report.legacy if pre_report else 0 + # Everything that was legacy before and is no longer legacy now was + # converted this run. + post_report.migrated = max(0, pre_legacy - post_report.legacy) + return post + + +async def migrate_encryption( + prisma_client: object, + user_api_key_dict: object, + dry_run: bool = False, +) -> MigrationReport: + """Run the full at-rest re-encryption migration. + + Requires ``general_settings.encryption_algorithm == 'aes-256-gcm'`` so writes + are produced in the AES format. Idempotent and resumable: re-running skips + already-migrated values and finishes any partial run. + + A ``dry_run`` performs no writes: the covered tables are scanned read-only + (so their residual legacy still counts toward the attestation) and the + net-new walkers run in dry-run mode. + """ + _assert_aes_gate_enabled() + + report = MigrationReport() + + # Tables that already have a rotation path (items 1, 2, 5-10). On a real run + # delegate to the rotation path (with bracketing scans for counts); on a dry + # run only classify them read-only. + if dry_run: + for covered in await _scan_covered_tables(prisma_client): + report.add(covered) + else: + for covered in await _migrate_covered_tables(prisma_client, user_api_key_dict): + report.add(covered) + + # Net-new walkers (items 3, 4, 11, 12, 13). + report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run)) + report.add( + await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run + ) + ) + report.add(await _migrate_sso_config(prisma_client, dry_run)) + + return report + + +async def check_encryption(prisma_client: object) -> MigrationReport: + """Read-only residual scan across **every** at-rest location. No writes. + + Covers both the rotation-managed tables (model / credentials / MCP credential + and env-var tables / config ``environment_variables``) and the net-new walker + locations (team and verification-token ``callback_vars``, vantage / cloudzero + config rows, SSO config). Reports how many values are still ``legacy``; + ``residual_legacy == 0`` across this full scan is the compliance attestation. + """ + report = MigrationReport() + + # Rotation-covered tables (read-only classification). + for covered in await _scan_covered_tables(prisma_client): + report.add(covered) + + # Net-new walker locations, in dry-run (read-only) mode. + report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True)) + report.add( + await _migrate_callback_vars_table( + prisma_client, "verification_token", dry_run=True + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True + ) + ) + report.add(await _migrate_sso_config(prisma_client, dry_run=True)) + return report diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2f5c38b0131..15e228a1a9f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -4091,6 +4091,86 @@ async def _rotate_master_key( verbose_proxy_logger.debug(f"Successfully re-encrypted {len(credentials)} credentials with new master key") +def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: + from litellm.proxy._types import CommonProxyErrors + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException( + status_code=403, + detail={"error": CommonProxyErrors.not_allowed_access.value}, + ) + + +@router.post( + "/credentials/migrate-encryption", + tags=["credential management"], + dependencies=[Depends(user_api_key_auth)], +) +async def migrate_encryption_endpoint( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + dry_run: bool = Query( + False, + description="If true, scan and report without writing any changes.", + ), +): + """ + Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format. + + Admin only. Requires ``general_settings.encryption_algorithm: aes-256-gcm``. + Idempotent and resumable — re-running skips already-migrated values. Pass + ``dry_run=true`` for a non-mutating scan (equivalent to ``--check``). + """ + from litellm.proxy._types import CommonProxyErrors + from litellm.proxy.management_endpoints.credential_migration import ( + migrate_encryption, + ) + from litellm.proxy.proxy_server import prisma_client + + _require_proxy_admin(user_api_key_dict) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + report = await migrate_encryption( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + dry_run=dry_run, + ) + return {"status": "success", "dry_run": dry_run, "report": report.as_dict()} + + +@router.get( + "/credentials/migrate-encryption/check", + tags=["credential management"], + dependencies=[Depends(user_api_key_auth)], +) +async def check_encryption_endpoint( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Read-only residual scan for compliance attestation. Reports how many at-rest + values are still in the legacy format. ``residual_legacy == 0`` attests no + legacy ciphertext remains. Admin only; performs no writes. + """ + from litellm.proxy._types import CommonProxyErrors + from litellm.proxy.management_endpoints.credential_migration import ( + check_encryption, + ) + from litellm.proxy.proxy_server import prisma_client + + _require_proxy_admin(user_api_key_dict) + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + report = await check_encryption(prisma_client=prisma_client) + return {"status": "success", "report": report.as_dict()} + + async def get_new_token(data: Optional[RegenerateKeyRequest]) -> str: if data and data.new_key is not None: # Reject custom key values if disabled by admin diff --git a/tests/proxy_behavior/management/test_credential_migration_endpoint.py b/tests/proxy_behavior/management/test_credential_migration_endpoint.py new file mode 100644 index 00000000000..b0428195674 --- /dev/null +++ b/tests/proxy_behavior/management/test_credential_migration_endpoint.py @@ -0,0 +1,71 @@ +"""Behavior scenarios for the credential re-encryption migration endpoints. + +These run against the live ASGI app + DB. The migration POST is *not* exercised +end-to-end here because it mutates shared at-rest data (it delegates to the +master-key rotation path); that full flow is covered by the unit suite and a +live proxy run. Here we pin the HTTP-boundary contract: the read-only check is +admin-reachable, and both routes are admin-gated. + +Both routes are admin-only management routes, so a non-admin key is rejected by +the ``user_api_key_auth`` layer (401) before the endpoint's own admin guard runs +-- the negative scenarios assert that framework-level rejection. +""" + +import pytest + +from .conftest import MASTER_KEY + +pytestmark = pytest.mark.asyncio(loop_scope="session") + + +async def test_migrate_encryption_check_as_admin_is_read_only(proxy_client): + """GET /credentials/migrate-encryption/check returns a residual report (no writes).""" + resp = await proxy_client.get( + "/credentials/migrate-encryption/check", + headers={"Authorization": f"Bearer {MASTER_KEY}"}, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert body["status"] == "success" + assert "residual_legacy" in body["report"] + + +async def test_migrate_encryption_check_requires_admin(proxy_client, scratch): + """A non-admin key cannot reach the residual scan (auth layer rejects, 401).""" + gen = await proxy_client.post( + "/key/generate", + headers={"Authorization": f"Bearer {MASTER_KEY}"}, + json={"key_alias": scratch.tag("check"), "user_id": scratch.tag("check-user")}, + ) + assert gen.status_code == 200, gen.text + nonadmin_key = gen.json()["key"] + + resp = await proxy_client.get( + "/credentials/migrate-encryption/check", + headers={"Authorization": f"Bearer {nonadmin_key}"}, + ) + assert resp.status_code == 401, resp.text + + +async def test_migrate_encryption_requires_admin(proxy_client, scratch): + """A non-admin key cannot trigger the migration (auth layer rejects, 401). + + Rejection happens before any write: the admin-only route check fires in + ``user_api_key_auth``, ahead of the endpoint body. + """ + gen = await proxy_client.post( + "/key/generate", + headers={"Authorization": f"Bearer {MASTER_KEY}"}, + json={ + "key_alias": scratch.tag("migrate"), + "user_id": scratch.tag("migrate-user"), + }, + ) + assert gen.status_code == 200, gen.text + nonadmin_key = gen.json()["key"] + + resp = await proxy_client.post( + "/credentials/migrate-encryption", + headers={"Authorization": f"Bearer {nonadmin_key}"}, + ) + assert resp.status_code == 401, resp.text diff --git a/tests/test_litellm/proxy/client/cli/test_encryption_commands.py b/tests/test_litellm/proxy/client/cli/test_encryption_commands.py new file mode 100644 index 00000000000..43e53cf5be2 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_encryption_commands.py @@ -0,0 +1,71 @@ +"""CLI tests for the ``litellm-proxy encryption migrate`` command. + +The HTTP client is mocked, so these assert the command's request routing (GET +check vs POST migrate, dry-run param) and its residual-state messaging without a +live proxy. +""" + +import pytest +from click.testing import CliRunner + +from litellm.proxy.client.cli import main as cli_main +from litellm.proxy.client.cli.commands import encryption as enc_cli + + +class _FakeHTTPClient: + """Stand-in for HTTPClient: records the last request and returns a canned body.""" + + last = None + response = {"status": "success", "report": {"residual_legacy": 0, "locations": {}}} + + def __init__(self, base_url, api_key): + self.base_url = base_url + self.api_key = api_key + + def request(self, method, path, **kwargs): + _FakeHTTPClient.last = {"method": method, "path": path, **kwargs} + return _FakeHTTPClient.response + + +@pytest.fixture +def runner(monkeypatch): + _FakeHTTPClient.last = None + _FakeHTTPClient.response = { + "status": "success", + "report": {"residual_legacy": 0, "locations": {}}, + } + monkeypatch.setattr(enc_cli, "HTTPClient", _FakeHTTPClient) + return CliRunner() + + +def test_migrate_check_hits_check_route(runner): + result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--check"]) + assert result.exit_code == 0, result.output + assert _FakeHTTPClient.last["method"] == "GET" + assert _FakeHTTPClient.last["path"] == "/credentials/migrate-encryption/check" + assert "No legacy values remaining" in result.output + + +def test_migrate_default_posts_without_dry_run(runner): + result = runner.invoke(cli_main.cli, ["encryption", "migrate"]) + assert result.exit_code == 0, result.output + assert _FakeHTTPClient.last["method"] == "POST" + assert _FakeHTTPClient.last["path"] == "/credentials/migrate-encryption" + assert _FakeHTTPClient.last["params"] is None + + +def test_migrate_dry_run_sets_param(runner): + result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--dry-run"]) + assert result.exit_code == 0, result.output + assert _FakeHTTPClient.last["method"] == "POST" + assert _FakeHTTPClient.last["params"] == {"dry_run": "true"} + + +def test_migrate_reports_residual_legacy(runner): + _FakeHTTPClient.response = { + "status": "success", + "report": {"residual_legacy": 3, "locations": {}}, + } + result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--check"]) + assert result.exit_code == 0, result.output + assert "Residual legacy values remaining: 3" in result.output diff --git a/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py b/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py new file mode 100644 index 00000000000..bee39e01dd6 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py @@ -0,0 +1,148 @@ +""" +Tests for the at-rest credential encryption chokepoint. + +Covers the AES-256-GCM (``v2:gcm:``) path, the ``encryption_algorithm`` config +gate, and the backward-compatibility guarantees that let legacy XSalsa20-Poly1305 +(nacl) ciphertext and new AES values coexist and decrypt correctly. +""" + +import pytest + +from litellm.proxy import proxy_server +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + _V2_GCM_PREFIX, + decrypt_value_helper, + encrypt_value_helper, +) + + +def _use_aes(monkeypatch): + """Flip the write-time algorithm to AES-256-GCM for the duration of a test.""" + monkeypatch.setattr( + proxy_server, "general_settings", {"encryption_algorithm": "aes-256-gcm"} + ) + + +@pytest.fixture(autouse=True) +def _salt_key(monkeypatch): + # Dominant convention in the test_litellm/ tree: set the key via env. + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-aes-1234") + # Ensure the legacy default is in force unless a test opts into AES. + monkeypatch.setattr(proxy_server, "general_settings", {}) + yield + + +def test_aes_gcm_round_trip(monkeypatch): + """A value written under AES-256-GCM is tagged v2:gcm: and decrypts back.""" + _use_aes(monkeypatch) + + ct = encrypt_value_helper("super-secret") + + assert ct.startswith(_V2_GCM_PREFIX) + assert decrypt_value_helper(ct, key="t") == "super-secret" + + +def test_default_is_legacy_algorithm(monkeypatch): + """With no config, writes stay on the legacy algorithm (no v2: marker).""" + ct = encrypt_value_helper("legacy-secret") + + assert not ct.startswith(_V2_GCM_PREFIX) + assert decrypt_value_helper(ct, key="t") == "legacy-secret" + + +def test_legacy_nacl_value_still_decrypts_after_flag_flip(monkeypatch): + """A value written under the old algorithm decrypts unchanged once AES is on. + + This is the mixed-format readback guarantee: decrypt is format-detecting, so + flipping the flag forward never strands previously-written data. + """ + legacy = encrypt_value_helper("legacy-secret") # default = xsalsa20 + assert not legacy.startswith(_V2_GCM_PREFIX) + + _use_aes(monkeypatch) + # New writes are now AES, but the old value must still come back. + assert decrypt_value_helper(legacy, key="t") == "legacy-secret" + assert encrypt_value_helper("fresh").startswith(_V2_GCM_PREFIX) + + +def test_v2_prefix_is_idempotent_marker(monkeypatch): + """The migration's skip-check: an already-v2 value is recognized by its prefix. + + Re-encrypting an AES value yields a fresh (different nonce) AES value, but the + prefix is what lets a migration skip already-migrated rows without decrypting. + """ + _use_aes(monkeypatch) + + ct = encrypt_value_helper("secret") + assert ct.startswith(_V2_GCM_PREFIX) + + # Round-tripping does not change the plaintext, and the marker is stable. + again = encrypt_value_helper(decrypt_value_helper(ct, key="t")) + assert again.startswith(_V2_GCM_PREFIX) + assert decrypt_value_helper(again, key="t") == "secret" + + +def test_aes_decrypt_failure_returns_none_not_raise(monkeypatch): + """Decrypt contract preserved: a garbled v2 value returns None, never raises.""" + _use_aes(monkeypatch) + + garbled = _V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!" + # exception_type="debug" exercises the swallow path; must not raise. + assert decrypt_value_helper(garbled, key="t", exception_type="debug") is None + + +def test_aes_decrypt_failure_returns_original_when_requested(monkeypatch): + """With return_original_value=True a bad v2 value comes back as-is, not None.""" + _use_aes(monkeypatch) + + garbled = _V2_GCM_PREFIX + "###" + assert ( + decrypt_value_helper( + garbled, key="t", exception_type="debug", return_original_value=True + ) + == garbled + ) + + +def test_empty_string_round_trips_under_aes(monkeypatch): + """Empty string is preserved through the AES path (parity with legacy).""" + _use_aes(monkeypatch) + + ct = encrypt_value_helper("") + assert ct.startswith(_V2_GCM_PREFIX) + assert decrypt_value_helper(ct, key="t") == "" + + +def test_callback_prefix_composes_with_v2(monkeypatch): + """litellm_enc:: + v2:gcm:... round-trips through the callback read path. + + Callback vars are stored as ``litellm_enc::``; the read path + strips ``litellm_enc::`` then calls the helper, so the value handed to the + helper is ``v2:gcm:...``. Ordering must work end to end. + """ + from litellm.proxy.common_utils.callback_utils import ( + _CALLBACK_VAR_ENCRYPTED_PREFIX, + _decrypt_or_passthrough, + _encrypt_if_plaintext, + ) + + _use_aes(monkeypatch) + + # "gcs_path_service_account" is a known-sensitive callback key. + stored = _encrypt_if_plaintext("gcs_path_service_account", "my-sa-secret") + + assert stored.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX) + inner = stored[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :] + assert inner.startswith(_V2_GCM_PREFIX) + assert _decrypt_or_passthrough("gcs_path_service_account", stored) == "my-sa-secret" + + +def test_unknown_algorithm_falls_back_to_legacy(monkeypatch): + """An unrecognized encryption_algorithm value does not produce v2 writes.""" + monkeypatch.setattr( + proxy_server, "general_settings", {"encryption_algorithm": "rot13"} + ) + + ct = encrypt_value_helper("secret") + assert not ct.startswith(_V2_GCM_PREFIX) + assert decrypt_value_helper(ct, key="t") == "secret" diff --git a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py new file mode 100644 index 00000000000..81226981089 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py @@ -0,0 +1,516 @@ +""" +Tests for the at-rest credential re-encryption migration engine. + +The pure engine (classify / reencrypt / selective-dict) is tested directly; the +DB walkers are tested against an AsyncMock Prisma client. Live end-to-end +proof-of-fix (real proxy + DB) is performed separately on the repro server. +""" + +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + _V2_GCM_PREFIX, + encrypt_value_helper, +) +from litellm.proxy.management_endpoints import credential_migration as cm + + +@pytest.fixture +def salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-migration-salt-1234") + monkeypatch.setattr(proxy_server, "general_settings", {}) + return "sk-migration-salt-1234" + + +def _legacy_ct(value: str, monkeypatch) -> str: + """Produce a legacy (nacl) ciphertext with the AES gate off.""" + monkeypatch.setattr(proxy_server, "general_settings", {}) + return encrypt_value_helper(value) + + +def _enable_aes(monkeypatch): + monkeypatch.setattr( + proxy_server, "general_settings", {"encryption_algorithm": "aes-256-gcm"} + ) + + +def _empty_covered_tables(client): + """Wire every rotation-covered table on `client` to return no rows. + + Lets a `check_encryption` / scanner test isolate the location under test + without the other covered tables raising on an unconfigured mock. + """ + for _, db_attr, _, _ in cm._COVERED_TABLE_SPECS: + getattr(client.db, db_attr).find_many = AsyncMock(return_value=[]) + + +# --------------------------- pure engine --------------------------- + + +def test_classify_value(salt_key, monkeypatch): + legacy = _legacy_ct("secret", monkeypatch) + _enable_aes(monkeypatch) + migrated = encrypt_value_helper("secret") + + assert cm.classify_value(legacy) == "legacy" + assert cm.classify_value(migrated) == "migrated" + assert cm.classify_value("just-plaintext") == "plaintext" + assert cm.classify_value("") == "plaintext" + assert cm.classify_value(123) == "not-a-string" + assert cm.classify_value(None) == "not-a-string" + + +def test_is_migrated(salt_key, monkeypatch): + _enable_aes(monkeypatch) + assert cm.is_migrated(encrypt_value_helper("x")) is True + assert cm.is_migrated("plaintext") is False + assert cm.is_migrated(5) is False + + +def test_reencrypt_value_legacy_to_v2(salt_key, monkeypatch): + legacy = _legacy_ct("secret", monkeypatch) + _enable_aes(monkeypatch) + + out = cm.reencrypt_value(legacy) + assert out != legacy + assert out.startswith(_V2_GCM_PREFIX) + + +def test_reencrypt_value_is_idempotent(salt_key, monkeypatch): + _enable_aes(monkeypatch) + v2 = encrypt_value_helper("secret") + # Already v2 -> returned byte-for-byte unchanged (no re-wrap). + assert cm.reencrypt_value(v2) == v2 + + +def test_reencrypt_value_preserves_non_string_and_empty(salt_key, monkeypatch): + _enable_aes(monkeypatch) + assert cm.reencrypt_value(42) == 42 + assert cm.reencrypt_value("") == "" + assert cm.reencrypt_value(None) is None + + +def test_reencrypt_value_skips_undecryptable(salt_key, monkeypatch): + """A value that does not decrypt (legacy plaintext or corrupt) is preserved.""" + _enable_aes(monkeypatch) + plaintext = "not-actually-encrypted" + assert cm.reencrypt_value(plaintext) == plaintext + + +def test_reencrypt_selective_dict(salt_key, monkeypatch): + legacy_key = _legacy_ct("the-api-key", monkeypatch) + _enable_aes(monkeypatch) + + data = {"api_key": legacy_key, "base_url": "https://x", "integration_token": None} + out = cm.reencrypt_selective_dict(data, ["api_key", "integration_token"]) + + assert out["api_key"].startswith(_V2_GCM_PREFIX) + assert out["base_url"] == "https://x" # untouched non-sensitive + assert out["integration_token"] is None # null skipped + + +# --------------------------- gate enforcement --------------------------- + + +@pytest.mark.asyncio +async def test_migrate_requires_aes_gate(salt_key, monkeypatch): + monkeypatch.setattr(proxy_server, "general_settings", {}) # gate off + with pytest.raises(RuntimeError, match="encryption_algorithm"): + await cm.migrate_encryption( + prisma_client=MagicMock(), user_api_key_dict=MagicMock() + ) + + +# --------------------------- config-row walker --------------------------- + + +def _config_prisma(record): + """Build an AsyncMock prisma client whose litellm_config returns `record`.""" + client = MagicMock() + client.db.litellm_config.find_unique = AsyncMock(return_value=record) + client.db.litellm_config.update = AsyncMock() + return client + + +@pytest.mark.asyncio +async def test_vantage_walker_migrates_legacy_field(salt_key, monkeypatch): + legacy_api_key = _legacy_ct("vantage-secret", monkeypatch) + _enable_aes(monkeypatch) + record = SimpleNamespace( + param_value={ + "api_key": legacy_api_key, + "integration_token": None, + "base_url": "https://api.vantage.sh", + } + ) + client = _config_prisma(record) + + report = await cm._migrate_config_settings_row( + client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=False + ) + + assert report.migrated == 1 + assert report.legacy == 0 # migrated -> no longer residual legacy + client.db.litellm_config.update.assert_awaited_once() + written = json.loads( + client.db.litellm_config.update.call_args.kwargs["data"]["param_value"] + ) + assert written["api_key"].startswith(_V2_GCM_PREFIX) + assert written["base_url"] == "https://api.vantage.sh" # non-sensitive untouched + + +@pytest.mark.asyncio +async def test_vantage_walker_idempotent_no_write(salt_key, monkeypatch): + _enable_aes(monkeypatch) + record = SimpleNamespace( + param_value={"api_key": encrypt_value_helper("already-v2"), "base_url": "x"} + ) + client = _config_prisma(record) + + report = await cm._migrate_config_settings_row( + client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=False + ) + + assert report.already_v2 == 1 + assert report.migrated == 0 + client.db.litellm_config.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_config_walker_dry_run_does_not_write(salt_key, monkeypatch): + legacy_api_key = _legacy_ct("vantage-secret", monkeypatch) + _enable_aes(monkeypatch) + record = SimpleNamespace(param_value={"api_key": legacy_api_key}) + client = _config_prisma(record) + + report = await cm._migrate_config_settings_row( + client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=True + ) + + # A dry run reports residual legacy only; nothing is migrated (no write), so + # `migrated` and `residual_legacy` are never contradictory in --check output. + assert report.legacy == 1 + assert report.migrated == 0 + client.db.litellm_config.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_config_walker_handles_missing_row(salt_key, monkeypatch): + _enable_aes(monkeypatch) + client = _config_prisma(None) + report = await cm._migrate_config_settings_row( + client, "cloudzero_settings", cm._CLOUDZERO_SENSITIVE, dry_run=False + ) + assert report.scanned == 0 + client.db.litellm_config.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_sso_walker_real_run_migrates_and_clears_residual(salt_key, monkeypatch): + """SSO real run: a migrated field is counted as migrated, not residual legacy.""" + legacy = _legacy_ct("client-secret", monkeypatch) + _enable_aes(monkeypatch) + record = SimpleNamespace(sso_settings={"client_secret": legacy, "client_id": "id"}) + client = MagicMock() + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) + client.db.litellm_ssoconfig.update = AsyncMock() + + report = await cm._migrate_sso_config(client, dry_run=False) + + assert report.migrated == 1 + assert report.legacy == 0 # migrated -> no longer residual + client.db.litellm_ssoconfig.update.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_sso_walker_dry_run_reports_residual_not_migrated(salt_key, monkeypatch): + """SSO dry run: residual legacy only; migrated stays 0 (never contradictory).""" + legacy = _legacy_ct("client-secret", monkeypatch) + _enable_aes(monkeypatch) + record = SimpleNamespace(sso_settings={"client_secret": legacy}) + client = MagicMock() + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) + client.db.litellm_ssoconfig.update = AsyncMock() + + report = await cm._migrate_sso_config(client, dry_run=True) + + assert report.legacy == 1 + assert report.migrated == 0 + client.db.litellm_ssoconfig.update.assert_not_awaited() + + +# --------------------------- --check scanner --------------------------- + + +@pytest.mark.asyncio +async def test_check_reports_residual_legacy(salt_key, monkeypatch): + legacy_api_key = _legacy_ct("vantage-secret", monkeypatch) + _enable_aes(monkeypatch) + + client = MagicMock() + # Net-new walker tables: empty team / token / sso, one legacy vantage field. + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.update = AsyncMock() + _empty_covered_tables(client) + + def _find_unique(where): + if where.get("param_name") == "vantage_settings": + return SimpleNamespace(param_value={"api_key": legacy_api_key}) + return None + + client.db.litellm_config.find_unique = AsyncMock(side_effect=_find_unique) + + report = await cm.check_encryption(client) + + assert report.residual_legacy == 1 + client.db.litellm_config.update.assert_not_awaited() # read-only + + +@pytest.mark.asyncio +async def test_check_reports_zero_after_migration(salt_key, monkeypatch): + _enable_aes(monkeypatch) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.update = AsyncMock() + _empty_covered_tables(client) + + def _find_unique(where): + if where.get("param_name") == "vantage_settings": + return SimpleNamespace( + param_value={"api_key": encrypt_value_helper("already-v2")} + ) + return None + + client.db.litellm_config.find_unique = AsyncMock(side_effect=_find_unique) + + report = await cm.check_encryption(client) + assert report.residual_legacy == 0 + + +# --------------------------- callback_vars walker --------------------------- + + +@pytest.mark.asyncio +async def test_callback_vars_walker_migrates_team_metadata(salt_key, monkeypatch): + """A team row with a legacy-encrypted callback var is rewritten to v2.""" + from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars + + # Legacy-encrypt a callback var via the real callback path (gate off). + monkeypatch.setattr(proxy_server, "general_settings", {}) + legacy_meta = encrypt_callback_vars( + {"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]} + ) + _enable_aes(monkeypatch) + + team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm._migrate_callback_vars_table(client, "team", dry_run=False) + + assert report.migrated == 1 + assert report.scanned == 1 # one field examined, not "post-v2" count + client.db.litellm_teamtable.update.assert_awaited_once() + written = json.loads( + client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"] + ) + inner = written["logging"][0]["callback_vars"]["gcs_path_service_account"] + assert "v2:gcm:" in inner + + +@pytest.mark.asyncio +async def test_callback_vars_walker_dry_run_reports_legacy(salt_key, monkeypatch): + """In --check (dry-run) mode, a legacy callback var counts as residual legacy.""" + from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars + + monkeypatch.setattr(proxy_server, "general_settings", {}) + legacy_meta = encrypt_callback_vars( + {"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]} + ) + _enable_aes(monkeypatch) + + team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm._migrate_callback_vars_table(client, "team", dry_run=True) + + assert report.scanned == 1 + assert report.legacy == 1 # would-migrate -> residual legacy in attestation + assert report.migrated == 0 + client.db.litellm_teamtable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_callback_vars_walker_migrates_callback_settings_shape( + salt_key, monkeypatch +): + """Regression: credentials under ``metadata.callback_settings.callback_vars`` + with no top-level ``logging`` key must be migrated, not skipped. + + The walker previously early-continued on ``"logging" not in metadata``, so + this credential shape (which ``encrypt_callback_vars`` does encrypt) was left + in legacy format at rest while the migration still reported success. + """ + from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars + + monkeypatch.setattr(proxy_server, "general_settings", {}) + legacy_meta = encrypt_callback_vars( + { + "callback_settings": { + "callback_vars": {"gcs_path_service_account": "sa-secret"} + } + } + ) + _enable_aes(monkeypatch) + assert "logging" not in legacy_meta # the shape that used to be skipped + + team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm._migrate_callback_vars_table(client, "team", dry_run=False) + + assert report.migrated == 1 + assert report.scanned == 1 + client.db.litellm_teamtable.update.assert_awaited_once() + written = json.loads( + client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"] + ) + inner = written["callback_settings"]["callback_vars"]["gcs_path_service_account"] + assert "v2:gcm:" in inner + + +@pytest.mark.asyncio +async def test_check_reports_callback_var_legacy_with_gate_off(salt_key, monkeypatch): + """check_encryption must report residual legacy callback vars even when the + AES gate is OFF. + + Detection is decrypt-based, not a re-encrypt delta, so it does not depend on + the write gate. A heuristic that re-encrypts and counts new v2 values would + read zero here (gate off -> no v2 produced) and emit a false-clean + attestation -- exactly the compliance trap this guards against. + """ + from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars + + # Legacy-encrypt a callback var, and leave the gate OFF for the check itself. + monkeypatch.setattr(proxy_server, "general_settings", {}) + legacy_meta = encrypt_callback_vars( + {"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]} + ) + team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm.check_encryption(client) + + assert report.residual_legacy == 1 + assert report.as_dict()["locations"]["team.callback_vars"]["legacy"] == 1 + client.db.litellm_teamtable.update.assert_not_awaited() # read-only + + +# --------------------------- covered-tables scanner --------------------------- + + +@pytest.mark.asyncio +async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatch): + """The read-only scanner classifies the model and credentials tables.""" + legacy = _legacy_ct("model-secret", monkeypatch) + _enable_aes(monkeypatch) + v2 = encrypt_value_helper("cred-secret") + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[ + SimpleNamespace(litellm_params={"api_key": legacy, "model": "gpt-4"}) + ] + ) + client.db.litellm_credentialstable.find_many = AsyncMock( + return_value=[SimpleNamespace(credential_values={"api_key": v2})] + ) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + by_loc = {r.location: r for r in await cm._scan_covered_tables(client)} + + assert by_loc["model_table"].legacy == 1 + assert by_loc["model_table"].plaintext == 1 # "gpt-4" model name, not ciphertext + assert by_loc["credentials"].already_v2 == 1 + assert by_loc["credentials"].legacy == 0 + + +@pytest.mark.asyncio +async def test_check_counts_covered_table_residual(salt_key, monkeypatch): + """check_encryption now scans the rotation-covered tables (model table here), + so a legacy value there counts toward residual_legacy (the P1 attestation gap). + """ + legacy = _legacy_ct("model-secret", monkeypatch) + _enable_aes(monkeypatch) + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.update = AsyncMock() + client.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[SimpleNamespace(litellm_params={"api_key": legacy})] + ) + + report = await cm.check_encryption(client) + + assert report.residual_legacy == 1 + assert report.as_dict()["locations"]["model_table"]["legacy"] == 1 + client.db.litellm_config.update.assert_not_awaited() # read-only + + +@pytest.mark.asyncio +async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch): + """_migrate_covered_tables derives real per-table counts from pre/post scans, + instead of the always-zero report Greptile flagged (P1). + """ + legacy = _legacy_ct("model-secret", monkeypatch) + _enable_aes(monkeypatch) + v2 = encrypt_value_helper("model-secret") + + row = SimpleNamespace(litellm_params={"api_key": legacy}) + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[row]) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + async def fake_rotate(**kwargs): + # Stand in for _rotate_master_key: re-encrypt the model api_key in place. + row.litellm_params["api_key"] = v2 + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key", + fake_rotate, + ) + + by_loc = { + r.location: r for r in await cm._migrate_covered_tables(client, MagicMock()) + } + + assert by_loc["model_table"].migrated == 1 # was legacy pre, v2 post + assert by_loc["model_table"].legacy == 0 # residual zero after rotation + assert by_loc["model_table"].already_v2 == 1 diff --git a/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py new file mode 100644 index 00000000000..e92cb3b13a7 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py @@ -0,0 +1,100 @@ +"""Unit tests for the at-rest encryption-migration HTTP endpoints. + +The endpoint bodies are exercised directly with the migration engine mocked, so +the admin guard, db-not-connected guard, and success path are all covered +without touching a live DB. The live ASGI/auth contract is covered separately in +``tests/proxy_behavior/management/test_credential_migration_endpoint.py``. +""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles +from litellm.proxy.management_endpoints import credential_migration as cm +from litellm.proxy.management_endpoints.key_management_endpoints import ( + check_encryption_endpoint, + migrate_encryption_endpoint, +) + +ADMIN = SimpleNamespace(user_role=LitellmUserRoles.PROXY_ADMIN.value) +NONADMIN = SimpleNamespace(user_role=LitellmUserRoles.INTERNAL_USER.value) + + +def _sample_report() -> cm.MigrationReport: + report = cm.MigrationReport() + report.add( + cm.LocationReport(location="model_table", scanned=2, migrated=1, legacy=0) + ) + return report + + +# ------------------------------- check endpoint ------------------------------- + + +@pytest.mark.asyncio +async def test_check_endpoint_success(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr( + cm, "check_encryption", AsyncMock(return_value=_sample_report()) + ) + + out = await check_encryption_endpoint(user_api_key_dict=ADMIN) + + assert out["status"] == "success" + assert out["report"]["residual_legacy"] == 0 + assert out["report"]["locations"]["model_table"]["scanned"] == 2 + + +@pytest.mark.asyncio +async def test_check_endpoint_requires_admin(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + with pytest.raises(HTTPException) as exc: + await check_encryption_endpoint(user_api_key_dict=NONADMIN) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_check_endpoint_db_not_connected(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", None) + with pytest.raises(HTTPException) as exc: + await check_encryption_endpoint(user_api_key_dict=ADMIN) + assert exc.value.status_code == 500 + + +# ------------------------------ migrate endpoint ------------------------------ + + +@pytest.mark.asyncio +@pytest.mark.parametrize("dry_run", [False, True]) +async def test_migrate_endpoint_success(monkeypatch, dry_run): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + fake = AsyncMock(return_value=_sample_report()) + monkeypatch.setattr(cm, "migrate_encryption", fake) + + out = await migrate_encryption_endpoint(user_api_key_dict=ADMIN, dry_run=dry_run) + + assert out["status"] == "success" + assert out["dry_run"] is dry_run + assert out["report"]["locations"]["model_table"]["migrated"] == 1 + # dry_run is threaded through to the engine unchanged. + assert fake.await_args.kwargs["dry_run"] is dry_run + + +@pytest.mark.asyncio +async def test_migrate_endpoint_requires_admin(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + with pytest.raises(HTTPException) as exc: + await migrate_encryption_endpoint(user_api_key_dict=NONADMIN) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_migrate_endpoint_db_not_connected(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", None) + with pytest.raises(HTTPException) as exc: + await migrate_encryption_endpoint(user_api_key_dict=ADMIN) + assert exc.value.status_code == 500 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index dced1dd4eaa..b3c2b44ee4e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -2479,6 +2479,52 @@ export interface paths { patch?: never; trace?: never; }; + "/credentials/migrate-encryption": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Migrate Encryption Endpoint + * @description Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format. + * + * Admin only. Requires ``general_settings.encryption_algorithm: aes-256-gcm``. + * Idempotent and resumable — re-running skips already-migrated values. Pass + * ``dry_run=true`` for a non-mutating scan (equivalent to ``--check``). + */ + post: operations["migrate_encryption_endpoint_credentials_migrate_encryption_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/credentials/migrate-encryption/check": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Check Encryption Endpoint + * @description Read-only residual scan for compliance attestation. Reports how many at-rest + * values are still in the legacy format. ``residual_legacy == 0`` attests no + * legacy ciphertext remains. Admin only; performs no writes. + */ + get: operations["check_encryption_endpoint_credentials_migrate_encryption_check_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/credentials/{credential_name}": { parameters: { query?: never; @@ -37122,6 +37168,58 @@ export interface operations { }; }; }; + migrate_encryption_endpoint_credentials_migrate_encryption_post: { + parameters: { + query?: { + /** @description If true, scan and report without writing any changes. */ + dry_run?: boolean; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + check_encryption_endpoint_credentials_migrate_encryption_check_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; delete_credential_credentials__credential_name__delete: { parameters: { query?: never; From 0e5aee18383709f425c63e5222c22ee002096e3f Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 29 Jun 2026 12:44:31 -0700 Subject: [PATCH 13/37] fix(ui): keep virtual-keys filters across delete and refresh (LIT-4080) (#31533) * fix(ui): keep virtual-keys filters across delete and refresh (LIT-4080) Filtering virtual keys by User ID and then deleting a key reset the filter to show all keys, and re-clicking Fetch did not re-apply it. The page ran two competing fetch paths: useKeys (React Query) fetched the page unfiltered while a separate useFilterLogic hook held its own filteredKeys list and, on any refresh, only re-applied Team and Organization client-side, silently dropping the User ID and Key Alias filters. Delete refreshed through the unfiltered useKeys path, so the filtered view collapsed back to everything VirtualKeysTable now owns its filter state and feeds every filter (team, organization, key alias, user id, key hash) straight into the useKeys options, so the filters are part of the React Query key. Any refetch or invalidation re-runs the same filtered query, which makes the reset-on-delete bug structurally impossible. Free-text inputs are debounced with @tanstack/react-pacer, sorting and pagination are server-side, and changing a filter or sort resets to page 1 Delete now invalidates keyKeys.lists() from key_info_view, matching the create path, instead of prop-drilling a refetch; the window "storage" refetch effect is removed. The dual-path useFilterLogic hook (and its test) are deleted Regression coverage: VirtualKeysTable threads an active User ID filter into the useKeys query and clears it on reset, useKeys encodes filter options in its query key so a filter change refetches, and key_info_view invalidates the keys list on delete * refactor(ui): simplify virtual-keys table data flow VirtualKeysTable now fetches its own teams and organizations via useOrganizations and the existing all-teams query instead of taking them as props, so the prop-drill through UserDashboard and the two page callers (page.tsx, ApiKeysDashboard) is gone along with their redundant organization state and fetch Filter state collapses from a useState plus a useDebouncedState mirror into a single source whose debounced copy is derived with useDebouncedValue, and one typed toKeyListFilters adapter maps it to the key/list query options. Behavior is unchanged; same 300ms debounce and the same reset timing The unused onSortChange/currentSort props and their sync effect are removed since no caller passed them, leaving sorting fully internal Adds a created_by_user alias-over-email regression test that fails if the display precedence is swapped * test(ui): add required last_active to useKeys mock fixtures The KeyResponse type requires last_active, so the typed mockKeys fixtures were missing it. Add it so the file type-checks cleanly. * chore(ui): ratchet lint budgets after virtual-keys refactor Deleting filter_logic.tsx and simplifying VirtualKeysTable lowered the no-explicit-any (2026 to 2016) and complexity (128 to 127) counts, so the eslint-metrics.json baseline was stale and failed the frontend-lint budget gate. Regenerate it, and drop the now-dead filter_logic.tsx suppression entry for the file this PR removed. * fix(ui): show a loading state for data-backed filter dropdowns The Team ID and Organization ID filters source their options from async hooks (teams / organizations). While that data was still loading the dropdowns rendered 'No results found', so they looked empty rather than loading. Add an opt-in loading flag to FilterOption that the searchable select surfaces as a spinner and a 'Loading...' empty state, and wire it from the teams and organizations query loading states. While loading, the filter no longer caches an empty initial-options list, so the real options appear once the data arrives. --- ui/litellm-dashboard/eslint-metrics.json | 4 +- ui/litellm-dashboard/eslint-suppressions.json | 14 - .../api-keys/ApiKeysDashboard.test.tsx | 4 - .../(dashboard)/api-keys/ApiKeysDashboard.tsx | 7 - .../(dashboard)/hooks/keys/useKeys.test.ts | 27 + .../src/app/(dashboard)/page.tsx | 8 +- .../VirtualKeysPage/VirtualKeysTable.test.tsx | 893 +++--------------- .../VirtualKeysPage/VirtualKeysTable.tsx | 180 ++-- .../key_team_helpers/filter_logic.test.tsx | 181 ---- .../key_team_helpers/filter_logic.tsx | 188 ---- .../src/components/molecules/filter.test.tsx | 61 ++ .../src/components/molecules/filter.tsx | 8 +- .../KeyInfoView.handleKeyUpdate.test.tsx | 9 + .../templates/key_info_view.test.tsx | 37 +- .../components/templates/key_info_view.tsx | 4 + .../src/components/user_dashboard.test.tsx | 1 - .../src/components/user_dashboard.tsx | 4 +- 17 files changed, 367 insertions(+), 1263 deletions(-) delete mode 100644 ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.tsx diff --git a/ui/litellm-dashboard/eslint-metrics.json b/ui/litellm-dashboard/eslint-metrics.json index 01c5a241562..deef1136b43 100644 --- a/ui/litellm-dashboard/eslint-metrics.json +++ b/ui/litellm-dashboard/eslint-metrics.json @@ -1,5 +1,5 @@ { - "@typescript-eslint/no-explicit-any": 2026, - "complexity": 128, + "@typescript-eslint/no-explicit-any": 2016, + "complexity": 127, "max-depth": 61 } diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 770c953d3f3..5d3a54f9e81 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1133,20 +1133,6 @@ "count": 1 } }, - "src/components/key_team_helpers/filter_logic.tsx": { - "react-hooks/purity": { - "count": 1 - }, - "react-hooks/refs": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 3 - }, - "react-hooks/use-memo": { - "count": 1 - } - }, "src/components/key_team_helpers/key_list.tsx": { "react-hooks/set-state-in-effect": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx index 3a3251ea76a..13689afcb52 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx @@ -42,10 +42,6 @@ vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ teamListCall: vi.fn(() => new Promise(() => {})), })); -vi.mock("@/components/organizations", () => ({ - fetchOrganizations: vi.fn(), -})); - vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(""), })); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx index 54f4bf41a21..ae0c443910a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx @@ -3,9 +3,7 @@ import { teamListCall as v2TeamListCall } from "@/app/(dashboard)/hooks/teams/useTeams"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { KeyResponse, Team } from "@/components/key_team_helpers/key_list"; -import { Organization } from "@/components/networking"; import { CreateKeyPrefillData } from "@/components/organisms/create_key_button"; -import { fetchOrganizations } from "@/components/organizations"; import UserDashboard from "@/components/user_dashboard"; import { useAuth } from "@/contexts/AuthContext"; import { useSearchParams } from "next/navigation"; @@ -20,7 +18,6 @@ export default function ApiKeysDashboard() { const [teams, setTeams] = useState(null); const [keys, setKeys] = useState([]); - const [organizations, setOrganizations] = useState([]); const [createClicked, setCreateClicked] = useState(false); const autoOpenCreate = searchParams.get("create") === "true"; @@ -77,9 +74,6 @@ export default function ApiKeysDashboard() { .then((response) => setTeams(response.teams ?? [])) .catch(console.error); } - if (accessToken) { - fetchOrganizations(accessToken, setOrganizations); - } }, [accessToken, userID, userRole]); return ( @@ -94,7 +88,6 @@ export default function ApiKeysDashboard() { setUserEmail={setUserEmail} setTeams={setTeams} setKeys={setKeys} - organizations={organizations} addKey={addKey} createClicked={createClicked} autoOpenCreate={autoOpenCreate} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts index 164e393eb9b..8c9b33f2c3e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts @@ -69,6 +69,7 @@ const mockKeys: KeyResponse[] = [ organization_id: null, created_at: "2024-01-01T00:00:00Z", updated_at: "2024-01-01T00:00:00Z", + last_active: null, team_spend: 0, team_alias: "", team_tpm_limit: 0, @@ -125,6 +126,7 @@ const mockKeys: KeyResponse[] = [ organization_id: null, created_at: "2024-01-01T00:00:00Z", updated_at: "2024-01-01T00:00:00Z", + last_active: null, team_spend: 0, team_alias: "test-team", team_tpm_limit: 1000, @@ -460,6 +462,31 @@ describe("useKeys", () => { expect(callUrl).not.toContain("project_id"); }); + // LIT-4080 guard: filter options must be part of the query key, not just the + // queryFn closure. If they were dropped from the key, changing a filter would + // reuse the cached (unfiltered) result and never refetch — exactly the bug + // where deleting a key wiped the active User ID filter. + it("refetches with the new filter when a filter option changes (options are in the query key)", async () => { + mockFetch.mockResolvedValue({ + ok: true, + json: async () => mockKeysResponse, + }); + + const { result, rerender } = renderHook(({ userID }) => useKeys(1, 10, { userID }), { + wrapper, + initialProps: { userID: "user-1" }, + }); + + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(mockFetch).toHaveBeenCalledTimes(1); + expect(mockFetch.mock.calls[0][0]).toContain("user_id=user-1"); + + rerender({ userID: "user-2" }); + + await waitFor(() => expect(mockFetch).toHaveBeenCalledTimes(2)); + expect(mockFetch.mock.calls[1][0]).toContain("user_id=user-2"); + }); + it("should pass agentID filter to the API", async () => { mockFetch.mockResolvedValueOnce({ ok: true, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx index c5d28fab8a0..cb4a4a0de03 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx @@ -4,8 +4,7 @@ import ApiKeysDashboard from "@/app/(dashboard)/api-keys/ApiKeysDashboard"; import { teamListCall as v2TeamListCall } from "@/app/(dashboard)/hooks/teams/useTeams"; import LoadingScreen from "@/components/common_components/LoadingScreen"; import { Team } from "@/components/key_team_helpers/key_list"; -import { Organization, proxyBaseUrl } from "@/components/networking"; -import { fetchOrganizations } from "@/components/organizations"; +import { proxyBaseUrl } from "@/components/networking"; import UserDashboard from "@/components/user_dashboard"; import { useAuth } from "@/contexts/AuthContext"; import { @@ -25,7 +24,6 @@ function CreateKeyPageContent() { const [teams, setTeams] = useState(null); const [keys, setKeys] = useState([]); - const [organizations, setOrganizations] = useState([]); const router = useRouter(); const searchParams = useSearchParams()!; @@ -112,9 +110,6 @@ function CreateKeyPageContent() { .then((response) => setTeams(response.teams ?? [])) .catch(console.error); } - if (accessToken) { - fetchOrganizations(accessToken, setOrganizations); - } }, [accessToken, userID, userRole]); if (authLoading || redirectToLogin || isLegacyRedirect) { @@ -135,7 +130,6 @@ function CreateKeyPageContent() { setUserEmail={setUserEmail} setTeams={setTeams} setKeys={setKeys} - organizations={organizations} addKey={addKey} createClicked={createClicked} /> diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 4002bbd245f..6389b82c9fd 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -1,63 +1,46 @@ import { act, screen, waitFor, fireEvent } from "@testing-library/react"; -import { vi, it, expect, beforeEach, MockedFunction } from "vitest"; +import { vi, it, expect, beforeEach, describe, MockedFunction } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; import { VirtualKeysTable } from "./VirtualKeysTable"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; -import { Organization } from "../networking"; import { KeysResponse, useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; -import { useFilterLogic } from "../key_team_helpers/filter_logic"; import useTeams from "@/app/(dashboard)/hooks/useTeams"; -// Mock network calls -vi.mock("./networking", async (importOriginal) => { - const actual = await importOriginal(); +// Resolve debounced values synchronously so an applied filter lands in the useKeys query within the test tick. +vi.mock("@tanstack/react-pacer/debouncer", async () => { + const React = await vi.importActual("react"); return { - ...actual, - userListCall: vi.fn().mockResolvedValue({ - users: [ - { - user_id: "user-1", - user_email: "user@example.com", - user_role: "user", - }, - ], - }), - teamListCall: vi.fn().mockResolvedValue([]), + useDebouncedValue: (value: unknown) => [value, { cancel: vi.fn(), flush: vi.fn() }], + useDebouncedState: (initial: unknown) => { + const [value, setValue] = React.useState(initial); + return [value, setValue, { cancel: vi.fn(), flush: vi.fn() }]; + }, }; }); -// Mock filter helpers -vi.mock("./key_team_helpers/filter_helpers", () => ({ - fetchAllTeams: vi.fn().mockResolvedValue([ - { - team_id: "team-1", - team_alias: "Test Team", - }, - ]), - fetchAllOrganizations: vi.fn().mockResolvedValue([ - { - organization_id: "org-1", - organization_alias: "Test Organization", - }, - ]), +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: vi.fn(() => ({ + accessToken: "test-token", + userId: "test-user", + userRole: "Admin", + premiumUser: true, + token: "test-token", + })), +})); + +vi.mock("../key_team_helpers/filter_helpers", () => ({ + fetchAllTeams: vi.fn().mockResolvedValue([{ team_id: "team-1", team_alias: "Test Team" }]), })); -// Mock useKeys hook vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({ useKeys: vi.fn(), + keyKeys: { lists: () => ["keys", "list"] }, })); -// Mock useFilterLogic hook -vi.mock("../key_team_helpers/filter_logic", () => ({ - useFilterLogic: vi.fn(), -})); - -// Mock useTeams hook (used by KeyInfoView) vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({ default: vi.fn(), })); -// Mock useOrganizations hook vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({ useOrganizations: vi.fn().mockReturnValue({ data: [ @@ -69,15 +52,6 @@ vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({ }), })); -// Mock fetchTeams to prevent network calls -vi.mock("@/app/(dashboard)/networking", async (importOriginal) => { - const actual = await importOriginal(); - return { - ...actual, - fetchTeams: vi.fn().mockResolvedValue([]), - }; -}); - const mockKey: KeyResponse = { token: "sk-1234567890abcdef", token_id: "key-1", @@ -91,6 +65,7 @@ const mockKey: KeyResponse = { config: {}, user_id: "user-1", team_id: "team-1", + project_id: null, max_parallel_requests: 10, metadata: {}, tpm_limit: 1000, @@ -153,68 +128,33 @@ const mockTeam: Team = { created_at: "2024-10-01T10:00:00Z", keys: [], members_with_roles: [], + spend: 0, }; -const mockOrganization: Organization = { - organization_id: "org-1", - organization_alias: "Test Organization", - budget_id: "budget-1", - metadata: {}, - models: ["gpt-3.5-turbo", "gpt-4"], - spend: 100, - model_spend: { "gpt-3.5-turbo": 50, "gpt-4": 50 }, - created_at: "2024-10-01T10:00:00Z", - created_by: "user-1", - updated_at: "2024-11-01T10:00:00Z", - updated_by: "user-1", - litellm_budget_table: {}, - teams: [], - users: [], - members: [], -}; - -// Mock hook implementations const mockUseKeys = useKeys as MockedFunction; -const mockUseFilterLogic = useFilterLogic as MockedFunction; const mockUseTeams = useTeams as MockedFunction; -beforeEach(() => { - // Reset mocks before each test - vi.clearAllMocks(); - - // Setup default mock implementations - mockUseKeys.mockReturnValue({ +const keysResult = (keys: KeyResponse[], data: Partial = {}, extra: Record = {}) => + ({ data: { - keys: [mockKey], - total_count: 1, + keys, + total_count: keys.length, current_page: 1, total_pages: 1, + ...data, } as KeysResponse, isPending: false, isFetching: false, + isError: false, refetch: vi.fn(), - } as any); + ...extra, + }) as any; - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "team-1", - "Organization ID": "org-1", - "Key Alias": "Test Key Alias", - "User ID": "user-1", - "User Email": "user@example.com", - "User Role": "user", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [mockKey], - filteredTotalCount: null, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); +beforeEach(() => { + vi.clearAllMocks(); + + mockUseKeys.mockReturnValue(keysResult([mockKey])); - // Mock useTeams hook (used by KeyInfoView) mockUseTeams.mockReturnValue({ teams: [mockTeam], setTeams: vi.fn(), @@ -222,33 +162,12 @@ beforeEach(() => { }); it("should render VirtualKeysTable component", () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); - + renderWithProviders(); expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); it("should display key information correctly", async () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); @@ -258,17 +177,7 @@ it("should display key information correctly", async () => { }); it("should display user email correctly", async () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("user@example.com")).toBeInTheDocument(); @@ -276,121 +185,36 @@ it("should display user email correctly", async () => { }); it("should show loading message only on initial load (isPending)", () => { - // Mock initial loading state - mockUseKeys.mockReturnValue({ - data: null, - isPending: true, - isFetching: true, - refetch: vi.fn(), - } as any); + mockUseKeys.mockReturnValue(keysResult([], {}, { data: null, isPending: true, isFetching: true })); - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; + renderWithProviders(); - renderWithProviders(); - - // Check that loading message is shown on initial load expect(screen.getByText("🚅 Loading keys...")).toBeInTheDocument(); - - // Check that actual key data is not shown expect(screen.queryByText("Test Key Alias")).not.toBeInTheDocument(); expect(screen.queryByText("Test Team")).not.toBeInTheDocument(); }); -it("should show 'No keys found' message when filteredKeys is empty", () => { - // Mock empty filteredKeys - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); +it("should show 'No keys found' message when the key list is empty", () => { + mockUseKeys.mockReturnValue(keysResult([])); - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); expect(screen.getByText("No keys found")).toBeInTheDocument(); }); it("should handle models with more than 3 entries to trigger expansion UI", () => { - const keyWithManyModels = { - ...mockKey, - models: ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "claude-3", "claude-3-5-sonnet"], - }; + mockUseKeys.mockReturnValue( + keysResult([{ ...mockKey, models: ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "claude-3", "claude-3-5-sonnet"] }]), + ); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithManyModels], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); + renderWithProviders(); - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); - - // This test ensures the ChevronDownIcon import (line 6) is used - // by having a key with > 3 models which triggers the expansion logic - // that uses ChevronDownIcon and ChevronRightIcon expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); it("should render table headers correctly", () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; + renderWithProviders(); - renderWithProviders(); - - // Check that main headers are rendered (testing the header.isPlaceholder condition path) expect(screen.getByText("Key ID")).toBeInTheDocument(); expect(screen.getByText("Key Alias")).toBeInTheDocument(); expect(screen.getByText("Team")).toBeInTheDocument(); @@ -399,442 +223,164 @@ it("should render table headers correctly", () => { }); it("should handle column resizing hover events", () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; + renderWithProviders(); - renderWithProviders(); - - // Find a header cell with data-header-id attribute const headerCell = document.querySelector("[data-header-id]") as HTMLElement; - expect(headerCell).toBeInTheDocument(); - // Check that the resizer element exists within the header const resizer = headerCell?.querySelector(".resizer") as HTMLElement; expect(resizer).toBeInTheDocument(); - - // Initially, resizer should have opacity 0 expect(resizer.style.opacity).toBe("0"); - // Simulate mouse enter using fireEvent - should set opacity to 0.5 (lines 612-616) fireEvent.mouseEnter(headerCell); expect(resizer.style.opacity).toBe("0.5"); - // Simulate mouse leave using fireEvent - should set opacity back to 0 (lines 618-622) fireEvent.mouseLeave(headerCell); expect(resizer.style.opacity).toBe("0"); }); it("should open KeyInfoView when clicking on a key ID button", async () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; + renderWithProviders(); - renderWithProviders(); - - // Wait for the table to render await waitFor(() => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); - // Verify table is visible before clicking - check for table-specific text expect(screen.getByText(/Showing.*results/)).toBeInTheDocument(); - // Find the key ID button (it shows the full token value, truncation is CSS-only) const keyIdButton = screen.getByText("sk-1234567890abcdef"); - expect(keyIdButton).toBeInTheDocument(); - - // Click on the key ID button fireEvent.click(keyIdButton); - // Wait for KeyInfoView to appear - check for unique elements that only exist in KeyInfoView await waitFor(() => { expect(screen.getByText("Back to Keys")).toBeInTheDocument(); - // KeyInfoHeader shows "Created At" metadata label expect(screen.getByText("Created At")).toBeInTheDocument(); }); - // Verify that table-specific elements are no longer visible - // The "Showing X of Y results" text should not be visible when KeyInfoView is open expect(screen.queryByText(/Showing.*results/)).not.toBeInTheDocument(); }); it("should display 'Default Proxy Admin' for user_id when value is 'default_user_id'", async () => { - const keyWithDefaultUserId = { - ...mockKey, - user_id: "default_user_id", - user_email: "", - user: { user_id: "default_user_id", user_email: "", user_alias: null }, - }; + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + user_id: "default_user_id", + user_email: "", + user: { user_id: "default_user_id", user_email: "", user_alias: null }, + }, + ]), + ); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithDefaultUserId], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Default Proxy Admin")).toBeInTheDocument(); }); }); -it("should display 'Default Proxy Admin' for created_by when value is 'default_user_id'", async () => { - const keyWithDefaultCreatedBy = { - ...mockKey, - created_by: "default_user_id", - }; - - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithDefaultCreatedBy], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); - - await waitFor(() => { - // The created_by column should display "Default Proxy Admin" - const defaultProxyAdminElements = screen.getAllByText("Default Proxy Admin"); - expect(defaultProxyAdminElements.length).toBeGreaterThan(0); - }); -}); - it("should display created_by_user email in 'Created By' column when available", async () => { - const keyWithCreatedByUser = { - ...mockKey, - created_by: "some-uuid-1234", - created_by_user: { - user_id: "some-uuid-1234", - user_email: "creator@example.com", - user_alias: null, - }, - }; + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + created_by: "some-uuid-1234", + created_by_user: { user_id: "some-uuid-1234", user_email: "creator@example.com", user_alias: null }, + }, + ]), + ); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithCreatedByUser], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("creator@example.com")).toBeInTheDocument(); }); }); -it("should display created_by_user alias over email when both available", async () => { - const keyWithCreatedByUser = { - ...mockKey, - created_by: "some-uuid-1234", - created_by_user: { - user_id: "some-uuid-1234", - user_email: "creator@example.com", - user_alias: "The Creator", - }, - }; +it("should display created_by_user alias over email when both are available", async () => { + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + created_by: "some-uuid-1234", + created_by_user: { user_id: "some-uuid-1234", user_email: "creator@example.com", user_alias: "The Creator" }, + }, + ]), + ); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithCreatedByUser], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("The Creator")).toBeInTheDocument(); }); + expect(screen.queryByText("creator@example.com")).not.toBeInTheDocument(); }); it("should render table without crashing when models is null", async () => { - const keyWithNullModels = { - ...mockKey, - models: null as unknown as string[], - }; + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, models: null as unknown as string[] }])); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithNullModels], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - // This should not throw an error - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); }); -it("should render table without crashing when models is undefined", async () => { - const keyWithUndefinedModels = { - ...mockKey, - models: undefined as unknown as string[], - }; - - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithUndefinedModels], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - // This should not throw an error - renderWithProviders(); - - await waitFor(() => { - expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); - }); -}); - -it("should render Last Active column header with info icon", () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); - - expect(screen.getByText("Last Active")).toBeInTheDocument(); -}); - -it("should display formatted date for last_active when value exists", async () => { - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); - - await waitFor(() => { - const expectedDate = new Date("2024-11-20T14:30:00Z").toLocaleDateString(); - expect(screen.getByText(expectedDate)).toBeInTheDocument(); - }); -}); - it("should display 'Unknown' for last_active when value is null", async () => { - const keyWithNullLastActive = { - ...mockKey, - last_active: null, - }; + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, last_active: null }])); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [keyWithNullLastActive], - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - const mockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { - sortBy: "created_at", - sortOrder: "desc" as const, - }, - }; - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Unknown")).toBeInTheDocument(); }); }); -const defaultMockProps = { - teams: [mockTeam], - organizations: [mockOrganization], - onSortChange: vi.fn(), - currentSort: { sortBy: "created_at", sortOrder: "desc" as const }, -}; +describe("server-side filtering – the LIT-4080 regression guard", () => { + it("threads an active User ID filter into the useKeys query so any refetch keeps it", async () => { + renderWithProviders(); -describe("pagination display – total count and page count", () => { - it("should show total_count from useKeys when no filter is active (filteredTotalCount is null)", async () => { - mockUseKeys.mockReturnValue({ - data: { - keys: [mockKey], - total_count: 509, - current_page: 1, - total_pages: 11, - } as KeysResponse, - isPending: false, - isFetching: false, - refetch: vi.fn(), - } as any); + fireEvent.click(screen.getByRole("button", { name: "Filters" })); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [mockKey], - filteredTotalCount: null, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), + const userIdInput = await screen.findByPlaceholderText("Enter User ID..."); + fireEvent.change(userIdInput, { target: { value: "user-42" } }); + + await waitFor(() => { + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" })); + }); + }); + + it("does not send filter params to useKeys when no filter is active", () => { + renderWithProviders(); + + const lastCall = mockUseKeys.mock.calls[mockUseKeys.mock.calls.length - 1]; + expect(lastCall[2] ?? {}).toMatchObject({ userID: undefined, teamID: undefined, keyHash: undefined }); + }); + + it("drops the filter from the useKeys query when Reset Filters is clicked", async () => { + renderWithProviders(); + + fireEvent.click(screen.getByRole("button", { name: "Filters" })); + const userIdInput = await screen.findByPlaceholderText("Enter User ID..."); + fireEvent.change(userIdInput, { target: { value: "user-42" } }); + + await waitFor(() => { + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" })); }); - renderWithProviders(); + fireEvent.click(screen.getByRole("button", { name: "Reset Filters" })); + + await waitFor(() => { + const lastCall = mockUseKeys.mock.calls[mockUseKeys.mock.calls.length - 1]; + expect((lastCall[2] ?? {}).userID).toBeUndefined(); + }); + }); +}); + +describe("pagination display – total count comes from useKeys", () => { + it("shows total_count and page count from the useKeys response", async () => { + mockUseKeys.mockReturnValue(keysResult([mockKey], { total_count: 509, total_pages: 11 })); + + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Showing 1 - 50 of 509 results")).toBeInTheDocument(); @@ -842,86 +388,21 @@ describe("pagination display – total count and page count", () => { }); }); - it("should show filteredTotalCount in pagination text when a filter search returns results", async () => { - mockUseKeys.mockReturnValue({ - data: { - keys: [mockKey], - total_count: 509, - current_page: 1, - total_pages: 11, - } as KeysResponse, - isPending: false, - isFetching: false, - refetch: vi.fn(), - } as any); + it("reflects a narrowed total when a filtered fetch returns fewer results", async () => { + mockUseKeys.mockReturnValue(keysResult([mockKey], { total_count: 1, total_pages: 1 })); - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "aaaaa", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [mockKey], - filteredTotalCount: 1, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); expect(screen.getByText("Page 1 of 1")).toBeInTheDocument(); }); }); - - it("should not show stale unfiltered totals when filteredTotalCount is set", async () => { - mockUseKeys.mockReturnValue({ - data: { - keys: [mockKey], - total_count: 509, - current_page: 1, - total_pages: 11, - } as KeysResponse, - isPending: false, - isFetching: false, - refetch: vi.fn(), - } as any); - - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "aaaaa", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [mockKey], - filteredTotalCount: 1, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); - - renderWithProviders(); - - await waitFor(() => { - expect(screen.queryByText(/509 results/)).not.toBeInTheDocument(); - expect(screen.queryByText(/of 11/)).not.toBeInTheDocument(); - }); - }); }); describe("refetch button", () => { it("should show Fetch button in normal state", () => { - renderWithProviders(); + renderWithProviders(); const fetchButton = screen.getByTitle("Fetch data"); expect(fetchButton).toBeInTheDocument(); @@ -930,64 +411,31 @@ describe("refetch button", () => { }); it("should show Fetching state and keep table data visible during refetch", () => { - mockUseKeys.mockReturnValue({ - data: { - keys: [mockKey], - total_count: 1, - current_page: 1, - total_pages: 1, - } as KeysResponse, - isPending: false, - isFetching: true, - refetch: vi.fn(), - } as any); + mockUseKeys.mockReturnValue(keysResult([mockKey], {}, { isFetching: true })); - renderWithProviders(); + renderWithProviders(); - // Button should show "Fetching" and be disabled expect(screen.getByText("Fetching")).toBeInTheDocument(); - const fetchButton = screen.getByTitle("Fetch data"); - expect(fetchButton).toBeDisabled(); - - // Table data should still be visible (stale data) + expect(screen.getByTitle("Fetch data")).toBeDisabled(); expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); - - // "Loading keys..." should NOT appear during refetch expect(screen.queryByText("🚅 Loading keys...")).not.toBeInTheDocument(); }); it("should call refetch when Fetch button is clicked", () => { const mockRefetch = vi.fn(); - mockUseKeys.mockReturnValue({ - data: { - keys: [mockKey], - total_count: 1, - current_page: 1, - total_pages: 1, - } as KeysResponse, - isPending: false, - isFetching: false, - refetch: mockRefetch, - } as any); + mockUseKeys.mockReturnValue(keysResult([mockKey], {}, { refetch: mockRefetch })); - renderWithProviders(); + renderWithProviders(); - const fetchButton = screen.getByTitle("Fetch data"); - fireEvent.click(fetchButton); + fireEvent.click(screen.getByTitle("Fetch data")); expect(mockRefetch).toHaveBeenCalledTimes(1); }); it("should show Fetch button enabled on error so user can retry", () => { - mockUseKeys.mockReturnValue({ - data: null, - isPending: false, - isFetching: false, - isError: true, - refetch: vi.fn(), - } as any); + mockUseKeys.mockReturnValue(keysResult([], {}, { data: null, isError: true })); - renderWithProviders(); + renderWithProviders(); const fetchButton = screen.getByTitle("Fetch data"); expect(fetchButton).not.toBeDisabled(); @@ -997,24 +445,9 @@ describe("refetch button", () => { describe("Status column reflects key.blocked / scim_blocked metadata", () => { it("should render Active for a non-blocked key", async () => { - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [{ ...mockKey, blocked: false, metadata: {} }], - filteredTotalCount: null, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: false, metadata: {} }])); - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByTestId(`key-status-${mockKey.token_id}`)).toHaveTextContent("Active"); @@ -1022,24 +455,9 @@ describe("Status column reflects key.blocked / scim_blocked metadata", () => { }); it("should render Blocked when key.blocked is true", async () => { - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [{ ...mockKey, blocked: true, metadata: {} }], - filteredTotalCount: null, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: true, metadata: {} }])); - renderWithProviders(); + renderWithProviders(); await waitFor(() => { expect(screen.getByTestId(`key-status-${mockKey.token_id}`)).toHaveTextContent("Blocked"); @@ -1048,24 +466,9 @@ describe("Status column reflects key.blocked / scim_blocked metadata", () => { }); it("should mark a SCIM-blocked key with the SCIM tooltip reason", async () => { - mockUseFilterLogic.mockReturnValue({ - filters: { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }, - filteredKeys: [{ ...mockKey, blocked: true, metadata: { scim_blocked: true } }], - filteredTotalCount: null, - allTeams: [mockTeam], - allOrganizations: [mockOrganization], - handleFilterChange: vi.fn(), - handleFilterReset: vi.fn(), - }); + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: true, metadata: { scim_blocked: true } }])); - renderWithProviders(); + renderWithProviders(); const tag = await screen.findByTestId(`key-status-${mockKey.token_id}`); expect(tag).toHaveTextContent("Blocked"); diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index 0c95d5dcfb7..0c71c4526ad 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -1,14 +1,15 @@ "use client"; -import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; +import { useKeys, KeyListCallOptions } from "@/app/(dashboard)/hooks/keys/useKeys"; import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useQuery } from "@tanstack/react-query"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { ChevronDownIcon, ChevronRightIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; import { ColumnDef, flexRender, getCoreRowModel, - getPaginationRowModel, - getSortedRowModel, PaginationState, SortingState, useReactTable, @@ -27,57 +28,57 @@ import { } from "@tremor/react"; import { InfoCircleOutlined, SyncOutlined } from "@ant-design/icons"; import { Button as AntButton, Popover, Skeleton, Tag, Tooltip, Typography } from "antd"; -import React, { useEffect, useDeferredValue, useMemo, useState } from "react"; +import React, { useDeferredValue, useMemo, useState } from "react"; import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key"; -import { useFilterLogic } from "../key_team_helpers/filter_logic"; +import { fetchAllTeams } from "../key_team_helpers/filter_helpers"; import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; import FilterComponent, { FilterOption } from "../molecules/filter"; import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; -import { Organization } from "../networking"; import KeyInfoView from "../templates/key_info_view"; -interface VirtualKeysTableProps { - teams: Team[] | null; - organizations: Organization[] | null; - onSortChange?: (sortBy: string, sortOrder: "asc" | "desc") => void; - currentSort?: { - sortBy: string; - sortOrder: "asc" | "desc"; - }; -} +type KeyFilterState = { + "Team ID": string; + "Organization ID": string; + "Key Alias": string; + "User ID": string; + "Key Hash": string; +}; -/** - * VirtualKeysTable – a new table for keys that mimics the table styling used in view_logs. - * The team selector and filtering have been removed so that all keys are shown. - */ +const DEFAULT_KEY_FILTERS: KeyFilterState = { + "Team ID": "", + "Organization ID": "", + "Key Alias": "", + "User ID": "", + "Key Hash": "", +}; -export function VirtualKeysTable({ teams, organizations, onSortChange, currentSort }: VirtualKeysTableProps) { - const { data: fetchedOrganizations } = useOrganizations(); - const resolvedOrganizations = fetchedOrganizations ?? organizations ?? []; +type KeyListFilterOptions = Pick< + KeyListCallOptions, + "teamID" | "organizationID" | "selectedKeyAlias" | "userID" | "keyHash" +>; + +const toKeyListFilters = (filters: KeyFilterState): KeyListFilterOptions => ({ + teamID: filters["Team ID"].trim() || undefined, + organizationID: filters["Organization ID"].trim() || undefined, + selectedKeyAlias: filters["Key Alias"].trim() || undefined, + userID: filters["User ID"].trim() || undefined, + keyHash: filters["Key Hash"].trim() || undefined, +}); + +export function VirtualKeysTable() { + const { accessToken } = useAuthorized(); + const { data: fetchedOrganizations, isLoading: isOrgsLoading } = useOrganizations(); + const resolvedOrganizations = useMemo(() => fetchedOrganizations ?? [], [fetchedOrganizations]); const [selectedKey, setSelectedKey] = useState(null); - const [sorting, setSorting] = React.useState(() => { - if (currentSort) { - return [ - { - id: currentSort.sortBy, - desc: currentSort.sortOrder === "desc", - }, - ]; - } - return [ - { - id: "created_at", - desc: true, - }, - ]; - }); + const [sorting, setSorting] = React.useState([{ id: "created_at", desc: true }]); const [tablePagination, setTablePagination] = React.useState({ pageIndex: 0, pageSize: 50, }); + const [filters, setFilters] = useState(DEFAULT_KEY_FILTERS); + const [debouncedFilters] = useDebouncedValue(filters, { wait: 300 }); - // Extract sort parameters from sorting state const sortBy = sorting.length > 0 ? sorting[0].id : null; const sortOrder = sorting.length > 0 ? (sorting[0].desc ? "desc" : "asc") : null; @@ -88,29 +89,22 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo isError, refetch, } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize, { + ...toKeyListFilters(debouncedFilters), sortBy: sortBy || undefined, sortOrder: sortOrder || undefined, expand: "user", }); const [expandedAccordions, setExpandedAccordions] = useState>({}); - // Use the filter logic hook - const keyList = useMemo(() => keys?.keys ?? [], [keys]); - const { - filters, - filteredKeys, - filteredTotalCount, - allTeams, - allOrganizations, - handleFilterChange, - handleFilterReset, - } = useFilterLogic({ - keys: keyList, - teams, - organizations, + const { data: fetchedTeams, isLoading: isTeamsLoading } = useQuery({ + queryKey: ["allTeamsForKeyFilters", accessToken], + queryFn: async () => (accessToken ? await fetchAllTeams(accessToken) : []), + enabled: !!accessToken, + staleTime: 30000, }); + const allTeams = useMemo(() => fetchedTeams ?? [], [fetchedTeams]); // Defer the transition so the button stays in loading state until the table // has rendered with the new data (mirrors the spend-logs pattern) @@ -121,23 +115,23 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo refetch(); }; - const totalCount = filteredTotalCount ?? keys?.total_count ?? 0; + const handleFilterChange = (newFilters: Record) => { + setFilters({ + "Team ID": newFilters["Team ID"] || "", + "Organization ID": newFilters["Organization ID"] || "", + "Key Alias": newFilters["Key Alias"] || "", + "User ID": newFilters["User ID"] || "", + "Key Hash": newFilters["Key Hash"] || "", + }); + setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); + }; - // Add a useEffect to call refresh when a key is created - useEffect(() => { - if (refetch) { - const handleStorageChange = () => { - refetch(); - }; + const handleFilterReset = () => { + setFilters(DEFAULT_KEY_FILTERS); + setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); + }; - // Listen for storage events that might indicate a key was created - window.addEventListener("storage", handleStorageChange); - - return () => { - window.removeEventListener("storage", handleStorageChange); - }; - } - }, [refetch]); + const totalCount = keys?.total_count ?? 0; const columns: ColumnDef[] = useMemo( () => [ @@ -237,7 +231,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo cell: (info) => { const teamId = info.getValue() as string | null; if (!teamId) return "-"; - const team = teams?.find((t) => t.team_id === teamId); + const team = allTeams.find((t) => t.team_id === teamId); const displayValue = team?.team_alias || teamId; const width = info.cell.column.getSize(); return ( @@ -471,7 +465,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo return `$${formatNumberWithCommas(maxBudget)}`; } const teamId = info.row.original.team_id; - const team = teams?.find((t) => t.team_id === teamId); + const team = allTeams.find((t) => t.team_id === teamId); if (team?.max_budget != null) { return `$${formatNumberWithCommas(team.max_budget)} (Team)`; } @@ -591,7 +585,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo }, }, ], - [teams, resolvedOrganizations], + [allTeams, resolvedOrganizations], ); const filterOptions: FilterOption[] = [ @@ -599,6 +593,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo name: "Team ID", label: "Team ID", isSearchable: true, + loading: isTeamsLoading, searchFn: async (searchText: string) => { if (!allTeams || allTeams.length === 0) return []; @@ -618,10 +613,11 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo name: "Organization ID", label: "Organization ID", isSearchable: true, + loading: isOrgsLoading, searchFn: async (searchText: string) => { - if (!allOrganizations || allOrganizations.length === 0) return []; + if (!resolvedOrganizations || resolvedOrganizations.length === 0) return []; - const filteredOrgs = allOrganizations.filter( + const filteredOrgs = resolvedOrganizations.filter( (org) => org.organization_id?.toLowerCase().includes(searchText.toLowerCase()) ?? false, ); @@ -651,7 +647,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo ]; const table = useReactTable({ - data: filteredKeys, + data: keyList, columns: columns.filter((col) => col.id !== "expander"), columnResizeMode: "onChange", columnResizeDirection: "ltr", @@ -662,45 +658,16 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo onSortingChange: (updaterOrValue) => { const newSorting = typeof updaterOrValue === "function" ? updaterOrValue(sorting) : updaterOrValue; setSorting(newSorting); - if (newSorting && newSorting.length > 0) { - const sortState = newSorting[0]; - const sortBy = sortState.id; - const sortOrder = sortState.desc ? "desc" : "asc"; - // Update filters state without triggering debouncedSearch - // The useKeys hook will automatically refetch with the new sort parameters - handleFilterChange( - { - ...filters, - "Sort By": sortBy, - "Sort Order": sortOrder, - }, - true, // skipDebounce - let useKeys handle the API call with correct page size - ); - onSortChange?.(sortBy, sortOrder); - } + setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); }, onPaginationChange: setTablePagination, getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - getPaginationRowModel: getPaginationRowModel(), enableSorting: true, - manualSorting: false, + manualSorting: true, manualPagination: true, pageCount: Math.ceil(totalCount / tablePagination.pageSize), }); - // Update local sorting state when currentSort prop changes - React.useEffect(() => { - if (currentSort) { - setSorting([ - { - id: currentSort.sortBy, - desc: currentSort.sortOrder === "desc", - }, - ]); - } - }, [currentSort]); - const { pageIndex, pageSize } = table.getState().pagination; const start = pageIndex * pageSize + 1; const end = Math.min((pageIndex + 1) * pageSize, totalCount); @@ -713,7 +680,6 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo onClose={() => setSelectedKey(null)} keyData={selectedKey} teams={allTeams} - onDelete={refetch} /> ) : (
@@ -867,7 +833,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
- ) : filteredKeys.length > 0 ? ( + ) : keyList.length > 0 ? ( table.getRowModel().rows.map((row) => ( {row.getVisibleCells().map((cell) => ( diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.test.tsx b/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.test.tsx deleted file mode 100644 index 23259528687..00000000000 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.test.tsx +++ /dev/null @@ -1,181 +0,0 @@ -import { act, renderHook, waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { useFilterLogic } from "./filter_logic"; -import { keyListCall } from "../networking"; - -vi.mock("../networking", () => ({ - keyListCall: vi.fn(), -})); - -vi.mock("./filter_helpers", () => ({ - fetchAllTeams: vi.fn().mockResolvedValue([]), - fetchAllOrganizations: vi.fn().mockResolvedValue([]), -})); - -const mockKey = { - token: "abc123", - key_alias: "aaaaa", - team_id: null, - organization_id: null, -}; - -const defaultProps = { - keys: [mockKey] as any[], - teams: [], - organizations: [], -}; - -const makeApiResponse = (overrides: { keys?: any[]; total_count?: number; total_pages?: number } = {}) => ({ - keys: overrides.keys ?? [mockKey], - total_count: overrides.total_count ?? 1, - current_page: 1, - total_pages: overrides.total_pages ?? 1, -}); - -describe("useFilterLogic – filteredTotalCount", () => { - beforeEach(() => { - vi.clearAllMocks(); - vi.mocked(keyListCall).mockResolvedValue(makeApiResponse({ total_count: 509, total_pages: 11 })); - }); - - it("should expose filteredTotalCount as null before any filter search runs", () => { - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - expect(result.current.filteredTotalCount).toBeNull(); - }); - - it("should set filteredTotalCount to the API total_count after a Key Alias filter is applied", async () => { - vi.mocked(keyListCall).mockResolvedValue(makeApiResponse({ keys: [mockKey], total_count: 1, total_pages: 1 })); - - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "aaaaa" }); - }); - - await waitFor( - () => { - expect(result.current.filteredTotalCount).toBe(1); - }, - { timeout: 500 }, - ); - }); - - it("should reflect the filtered total_count even when it differs from the full key count", async () => { - vi.mocked(keyListCall).mockResolvedValue(makeApiResponse({ total_count: 7, total_pages: 1 })); - - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-x" }); - }); - - await waitFor( - () => { - expect(result.current.filteredTotalCount).toBe(7); - }, - { timeout: 500 }, - ); - }); - - it("should reset filteredTotalCount to null when handleFilterReset is called", async () => { - vi.mocked(keyListCall).mockResolvedValue(makeApiResponse({ total_count: 1 })); - - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "aaaaa" }); - }); - - await waitFor( - () => { - expect(result.current.filteredTotalCount).toBe(1); - }, - { timeout: 500 }, - ); - - act(() => { - result.current.handleFilterReset(); - }); - - // filteredTotalCount resets synchronously before the debounced reset search completes - expect(result.current.filteredTotalCount).toBeNull(); - }); - - it("should pass the Key Alias value to keyListCall", async () => { - vi.mocked(keyListCall).mockResolvedValue(makeApiResponse({ total_count: 2 })); - - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "my-alias" }); - }); - - await waitFor( - () => { - expect(keyListCall).toHaveBeenCalledWith( - expect.any(String), // accessToken - null, // organizationID (empty → null) - null, // teamID (empty → null) - "my-alias", // selectedKeyAlias ← the filter value - null, // userID - null, // keyHash - 1, // page (resets to 1 on filter change) - expect.any(Number), // pageSize (defaultPageSize) - expect.anything(), // sortBy - expect.anything(), // sortOrder - ); - }, - { timeout: 500 }, - ); - }); - - it("should not update filteredTotalCount when keyListCall throws", async () => { - vi.mocked(keyListCall).mockRejectedValue(new Error("Network error")); - - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "bad-alias" }); - }); - - await waitFor( - () => { - expect(keyListCall).toHaveBeenCalled(); - }, - { timeout: 500 }, - ); - - expect(result.current.filteredTotalCount).toBeNull(); - }); - - it("should not enter an infinite update loop when keys is a fresh array reference on every render", () => { - const sourceKeys = [mockKey]; - let renderCount = 0; - - const { result } = renderHook(() => { - renderCount += 1; - const value = useFilterLogic({ keys: [...sourceKeys], teams: [], organizations: [] }); - if (renderCount > 25) { - throw new Error(`useFilterLogic re-rendered ${renderCount} times; setFilteredKeys is looping`); - } - return value; - }); - - expect(result.current.filteredKeys).toEqual([mockKey]); - expect(renderCount).toBeLessThanOrEqual(25); - }); - - it("should not trigger a debounced search when skipDebounce is true", async () => { - const { result } = renderHook(() => useFilterLogic(defaultProps)); - - act(() => { - result.current.handleFilterChange({ "Sort By": "spend", "Sort Order": "asc" }, true); - }); - - await new Promise((resolve) => setTimeout(resolve, 350)); - - expect(keyListCall).not.toHaveBeenCalled(); - expect(result.current.filteredTotalCount).toBeNull(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.tsx b/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.tsx deleted file mode 100644 index e31a6fbee38..00000000000 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_logic.tsx +++ /dev/null @@ -1,188 +0,0 @@ -import { useCallback, useEffect, useState, useRef } from "react"; -import { KeyResponse } from "../key_team_helpers/key_list"; -import { keyListCall, Organization } from "../networking"; -import { Team } from "../key_team_helpers/key_list"; -import { fetchAllOrganizations, fetchAllTeams } from "./filter_helpers"; -import { debounce } from "lodash"; -import { defaultPageSize } from "../constants"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; - -export interface FilterState { - "Team ID": string; - "Organization ID": string; - "Key Alias": string; - [key: string]: string; - "User ID": string; - "Sort By": string; - "Sort Order": string; -} - -export function useFilterLogic({ - keys, - teams, - organizations, -}: { - keys: KeyResponse[]; - teams: Team[] | null; - organizations: Organization[] | null; -}) { - const defaultFilters: FilterState = { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Sort By": "created_at", - "Sort Order": "desc", - }; - const { accessToken } = useAuthorized(); - const [filters, setFilters] = useState(defaultFilters); - const [allTeams, setAllTeams] = useState(teams || []); - const [allOrganizations, setAllOrganizations] = useState(organizations || []); - const [filteredKeys, setFilteredKeys] = useState(keys); - const [filteredTotalCount, setFilteredTotalCount] = useState(null); - const lastSearchTimestamp = useRef(0); - const debouncedSearch = useCallback( - debounce(async (filters: FilterState) => { - if (!accessToken) { - return; - } - - const currentTimestamp = Date.now(); - lastSearchTimestamp.current = currentTimestamp; - - try { - // Make the API call using userListCall with all filter parameters - const data = await keyListCall( - accessToken, - filters["Organization ID"] || null, - filters["Team ID"] || null, - filters["Key Alias"] || null, - filters["User ID"] || null, - filters["Key Hash"] || null, - 1, // Reset to first page when searching - defaultPageSize, - filters["Sort By"] || null, - filters["Sort Order"] || null, - ); - - // Only update state if this is the most recent search - if (currentTimestamp === lastSearchTimestamp.current) { - if (data) { - setFilteredKeys(data.keys); - setFilteredTotalCount(data.total_count ?? null); - console.log("called from debouncedSearch filters:", JSON.stringify(filters)); - console.log("called from debouncedSearch data:", JSON.stringify(data)); - } - } - } catch (error) { - console.error("Error searching users:", error); - } - }, 300), - [accessToken], - ); - // Apply filters to keys whenever keys or filters change - useEffect(() => { - if (!keys) { - setFilteredKeys([]); - return; - } - - let result = [...keys]; - - // Apply Team ID filter - if (filters["Team ID"]) { - result = result.filter((key) => key.team_id === filters["Team ID"]); - } - - // Apply Organization ID filter - if (filters["Organization ID"]) { - result = result.filter((key) => (key.organization_id ?? key.org_id) === filters["Organization ID"]); - } - - setFilteredKeys((prev) => - prev.length === result.length && prev.every((key, index) => key === result[index]) ? prev : result, - ); - }, [keys, filters]); - - // Fetch all data for filters when component mounts - useEffect(() => { - const loadAllFilterData = async () => { - // Load all teams - no organization filter needed here - const teamsData = await fetchAllTeams(accessToken); - if (teamsData.length > 0) { - setAllTeams(teamsData); - } - - // Load all organizations - const orgsData = await fetchAllOrganizations(accessToken); - if (orgsData.length > 0) { - setAllOrganizations(orgsData); - } - }; - - if (accessToken) { - loadAllFilterData(); - } - }, [accessToken]); - - // Update teams and organizations when props change - useEffect(() => { - if (teams && teams.length > 0) { - setAllTeams((prevTeams) => { - // Only update if we don't already have a larger set of teams - return prevTeams.length < teams.length ? teams : prevTeams; - }); - } - }, [teams]); - - useEffect(() => { - if (organizations && organizations.length > 0) { - setAllOrganizations((prevOrgs) => { - // Only update if we don't already have a larger set of organizations - return prevOrgs.length < organizations.length ? organizations : prevOrgs; - }); - } - }, [organizations]); - - const handleFilterChange = (newFilters: Record, skipDebounce: boolean = false) => { - // Update filters state - setFilters({ - "Team ID": newFilters["Team ID"] || "", - "Organization ID": newFilters["Organization ID"] || "", - "Key Alias": newFilters["Key Alias"] || "", - "User ID": newFilters["User ID"] || "", - "Sort By": newFilters["Sort By"] || "created_at", - "Sort Order": newFilters["Sort Order"] || "desc", - }); - - // Only trigger debouncedSearch if skipDebounce is false - // This allows sorting to be handled by the parent component's useKeys hook - if (!skipDebounce) { - // Fetch keys based on new filters - const updatedFilters = { - ...filters, - ...newFilters, - }; - debouncedSearch(updatedFilters); - } - }; - - const handleFilterReset = () => { - // Reset filters state - setFilters(defaultFilters); - setFilteredTotalCount(null); - - // Reset selections - debouncedSearch(defaultFilters); - }; - - return { - filters, - filteredKeys, - filteredTotalCount, - allTeams, - allOrganizations, - handleFilterChange, - handleFilterReset, - }; -} diff --git a/ui/litellm-dashboard/src/components/molecules/filter.test.tsx b/ui/litellm-dashboard/src/components/molecules/filter.test.tsx index 3a15c5c84f1..d956cd93168 100644 --- a/ui/litellm-dashboard/src/components/molecules/filter.test.tsx +++ b/ui/litellm-dashboard/src/components/molecules/filter.test.tsx @@ -326,6 +326,67 @@ describe("FilterComponent", () => { }); }); + it("shows a loading state (not an empty list) while a searchable filter's data is still loading", async () => { + const user = userEvent.setup({ delay: null }); + const mockSearchFn = vi.fn().mockResolvedValue([]); + + const options: FilterOption[] = [ + { + name: "model", + label: "Model", + isSearchable: true, + loading: true, + searchFn: mockSearchFn, + }, + ]; + + renderWithProviders( + , + ); + + await user.click(screen.getByRole("button", { name: "Filters" })); + + const modelLabel = screen.getByText("Model"); + const modelSelect = within(modelLabel.closest("div")!).getByRole("combobox"); + await user.click(modelSelect); + + await waitFor(() => { + expect(screen.getByText("Loading...")).toBeInTheDocument(); + }); + expect(screen.queryByText("No results found")).not.toBeInTheDocument(); + // It must not cache an empty initial-options list while the source is still loading. + expect(mockSearchFn).not.toHaveBeenCalled(); + }); + + it("loads initial options once a searchable filter's data finishes loading", async () => { + const user = userEvent.setup({ delay: null }); + const mockSearchFn = vi.fn().mockResolvedValue([{ label: "Team A", value: "team-a" }]); + const baseOption: FilterOption = { name: "model", label: "Model", isSearchable: true, searchFn: mockSearchFn }; + + const { rerender } = renderWithProviders( + , + ); + + await user.click(screen.getByRole("button", { name: "Filters" })); + expect(mockSearchFn).not.toHaveBeenCalled(); + + rerender( + , + ); + + await waitFor(() => { + expect(mockSearchFn).toHaveBeenCalledWith(""); + }); + }); + it("should handle search errors gracefully", async () => { const user = userEvent.setup({ delay: null }); const consoleErrorSpy = vi.spyOn(console, "error").mockImplementation(() => {}); diff --git a/ui/litellm-dashboard/src/components/molecules/filter.tsx b/ui/litellm-dashboard/src/components/molecules/filter.tsx index 7086892c32d..45ad3a6b9ca 100644 --- a/ui/litellm-dashboard/src/components/molecules/filter.tsx +++ b/ui/litellm-dashboard/src/components/molecules/filter.tsx @@ -17,6 +17,7 @@ export interface FilterOption { searchFn?: (searchText: string) => Promise>; options?: Array<{ label: string; value: string }>; customComponent?: React.ComponentType; + loading?: boolean; } interface FilterValues { @@ -74,7 +75,7 @@ const FilterComponent: React.FC = ({ // Load initial options for searchable filters const loadInitialOptions = useCallback( async (option: FilterOption) => { - if (!option.isSearchable || !option.searchFn || initialOptionsLoaded[option.name]) return; + if (!option.isSearchable || !option.searchFn || option.loading || initialOptionsLoaded[option.name]) return; setSearchLoadingMap((prev) => ({ ...prev, [option.name]: true })); setInitialOptionsLoaded((prev) => ({ ...prev, [option.name]: true })); @@ -145,6 +146,7 @@ const FilterComponent: React.FC = ({ {showFilters && (
{options.map((option) => { + const isOptionLoading = searchLoadingMap[option.name] || option.loading; return (
@@ -166,10 +168,10 @@ const FilterComponent: React.FC = ({ } }} filterOption={false} - loading={searchLoadingMap[option.name]} + loading={isOptionLoading} options={searchOptionsMap[option.name] || []} allowClear - notFoundContent={searchLoadingMap[option.name] ? "Loading..." : "No results found"} + notFoundContent={isOptionLoading ? "Loading..." : "No results found"} /> ) : option.options ? (