diff --git a/CLAUDE.md b/CLAUDE.md index 651a88bf9aa..683993c9476 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -51,7 +51,7 @@ Every lint or type suppression must name the exact rule inside brackets and carr Commit and push your work when you're done without asking -When you must use real LLM models to, for example, write e2e tests, write a QA runbook, etc., make sure to use the latest models (doesn't have to be smartest, can also be a modern small, fast one. No strong preference for smart vs fast here, just use something modern) as of the year and month of the current date. Do a web search as necessary to figure that out +When referencing or running models (coding, QA'ing, writing docs, writing tests, etc.), use the latest model in that model family unless otherwise specified; treat your training knowledge, memories, configs, and tests as stale, and determine the family's latest with model_prices_and_context_window.json or the web If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index cb37725d79c..4bf36a0d5c6 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -193,6 +193,14 @@ class AgenticAnthropicStreamingIterator: raise StopAsyncIteration + async def aclose(self) -> None: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + aclose_if_supported, + ) + + await aclose_if_supported(self._inner) + await aclose_if_supported(self._follow_up_iterator) + async def _process_agentic_hooks(self) -> None: """Rebuild the Anthropic response from collected SSE bytes and call hooks.""" if self._hook_processing_done: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 2357960f716..5f2b23d7eca 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -1,8 +1,13 @@ import asyncio import json from datetime import datetime -from typing import Any, AsyncIterator, List, Union +from typing import Any, AsyncIterator, List, Protocol, Union, runtime_checkable +import httpx +from pydantic import TypeAdapter +from typing_extensions import TypedDict + +from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, @@ -12,6 +17,93 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ = PassThroughEndpointLogging() +INCOMPLETE_STREAM_ERROR_MESSAGE = ( + "Provider stream ended before emitting a message_stop event; " + "the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated." +) + + +def _is_message_stop_chunk(chunk: object) -> bool: + if isinstance(chunk, dict): + return chunk.get("type") == "message_stop" + if isinstance(chunk, (bytes, bytearray)): + return any(line == b"event: message_stop" for line in chunk.splitlines()) + return False + + +def _is_provider_error_chunk(chunk: object) -> bool: + if isinstance(chunk, dict): + return chunk.get("type") == "error" + if isinstance(chunk, (bytes, bytearray)): + return any(line == b"event: error" for line in chunk.splitlines()) + return False + + +def _is_terminal_stream_chunk(chunk: object) -> bool: + return _is_message_stop_chunk(chunk) or _is_provider_error_chunk(chunk) + + +def _incomplete_stream_error_sse_event() -> bytes: + payload = json.dumps( + { + "type": "error", + "error": {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE}, + } + ) + return f"event: error\ndata: {payload}\n\n".encode() + + +class AnthropicMessagesStreamHiddenParams(TypedDict): + additional_headers: dict[str, str] + + +@runtime_checkable +class SupportsAclose(Protocol): + async def aclose(self) -> None: ... + + +async def aclose_if_supported(stream: object) -> None: + if isinstance(stream, SupportsAclose): + await stream.aclose() + + +_RESPONSE_HEADERS_ADAPTER: TypeAdapter[dict[str, str]] = TypeAdapter(dict[str, str]) + + +def anthropic_messages_stream_hidden_params( + response_headers: httpx.Headers, +) -> AnthropicMessagesStreamHiddenParams: + return AnthropicMessagesStreamHiddenParams( + additional_headers=_RESPONSE_HEADERS_ADAPTER.validate_python(process_response_headers(response_headers)) + ) + + +class AnthropicMessagesStreamingResponse: + """ + Wraps the /v1/messages SSE byte stream so upstream provider response + headers (e.g. Bedrock's x-amzn-requestid / x-amzn-trace-id) survive as + ``_hidden_params["additional_headers"]``, which the proxy forwards to + clients as ``llm_provider-*`` response headers. Bare async generators + cannot carry attributes, so header context was previously dropped. + """ + + def __init__( + self, + completion_stream: AsyncIterator[bytes], + hidden_params: AnthropicMessagesStreamHiddenParams, + ) -> None: + self.completion_stream = completion_stream + self._hidden_params = hidden_params + + def __aiter__(self) -> "AnthropicMessagesStreamingResponse": + return self + + async def __anext__(self) -> bytes: + return await self.completion_stream.__anext__() + + async def aclose(self) -> None: + await aclose_if_supported(self.completion_stream) + class BaseAnthropicMessagesStreamingIterator: """ @@ -102,13 +194,18 @@ class BaseAnthropicMessagesStreamingIterator: This method provides the common logic for both Anthropic and Bedrock implementations. """ collected_chunks = [] + saw_terminal_event = False async for chunk in completion_stream: if self.completion_start_time is None: self.completion_start_time = datetime.now() + saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk) encoded_chunk = self._convert_chunk_to_sse_format(chunk) collected_chunks.append(encoded_chunk) yield encoded_chunk + if not saw_terminal_event: + yield _incomplete_stream_error_sse_event() + # Handle logging after all chunks are processed await self._handle_streaming_logging(collected_chunks) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 3c10239f868..05401488f27 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2084,12 +2084,18 @@ class BaseLLMHTTPHandler: initial_response: Union[AsyncIterator, AnthropicMessagesResponse] if stream: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + AnthropicMessagesStreamingResponse, + anthropic_messages_stream_hidden_params, + ) + completion_stream = anthropic_messages_provider_config.get_async_streaming_response_iterator( model=model, httpx_response=response, request_body=request_body, litellm_logging_obj=logging_obj, ) + stream_hidden_params = anthropic_messages_stream_hidden_params(response.headers) if not self._has_agentic_completion_hook(logging_obj): # No callback overrides async_should_run_agentic_loop, so the @@ -2097,7 +2103,10 @@ class BaseLLMHTTPHandler: # and rebuilding the response from SSE at end-of-stream to call # hooks that all return (False, {}). Stream through directly and # skip that per-chunk + end-of-stream overhead. - return completion_stream + return AnthropicMessagesStreamingResponse( + completion_stream=completion_stream, + hidden_params=stream_hidden_params, + ) from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( AgenticAnthropicStreamingIterator, @@ -2114,7 +2123,10 @@ class BaseLLMHTTPHandler: custom_llm_provider=custom_llm_provider, kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, ) - return initial_response + return AnthropicMessagesStreamingResponse( + completion_stream=initial_response, + hidden_params=stream_hidden_params, + ) else: initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response( model=model, diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 37863e0e356..84d03fe7144 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -282,7 +282,7 @@ class HeadroomGuardrail(CustomGuardrail): self, messages: list[dict[str, object]], model: str | None, - ) -> tuple[list[dict[str, object]], bool]: + ) -> tuple[list[dict[str, object]], bool, dict[str, object]]: payload: dict[str, object] = {"messages": messages} if model: payload["model"] = model @@ -294,62 +294,94 @@ class HeadroomGuardrail(CustomGuardrail): headers=self._request_headers(), ) except httpx.HTTPStatusError as e: - return self._handle_compress_failure( - messages, - "Headroom compression service returned an error", - {"status_code": e.response.status_code, "body": e.response.text}, - ), False - except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e: - return self._handle_compress_failure( - messages, - "Headroom compression service unreachable", - {"detail": str(e)}, - ), False - if raw_response is None: - return self._handle_compress_failure( - messages, - "Headroom compression service returned no response", + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned an error", + {"status_code": e.response.status_code, "body": e.response.text}, + ), + False, {}, - ), False + ) + except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError, litellm.Timeout) as e: + return ( + self._handle_compress_failure( + messages, + "Headroom compression service unreachable", + {"detail": str(e)}, + ), + False, + {}, + ) + if raw_response is None: + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned no response", + {}, + ), + False, + {}, + ) response: HttpxResponse = raw_response if response.status_code != 200: - return self._handle_compress_failure( - messages, - "Headroom compression service returned an error", - {"status_code": response.status_code, "body": response.text}, - ), False + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned an error", + {"status_code": response.status_code, "body": response.text}, + ), + False, + {}, + ) try: body: object = response.json() except ValueError: - return self._handle_compress_failure( - messages, - "Headroom compression service returned non-JSON response", - {"body": response.text[:500]}, - ), False + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned non-JSON response", + {"body": response.text[:500]}, + ), + False, + {}, + ) if not _is_str_object_dict(body): - return self._handle_compress_failure( - messages, - "Headroom compression service returned unexpected response shape", - {"body": response.text[:500]}, - ), False + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned unexpected response shape", + {"body": response.text[:500]}, + ), + False, + {}, + ) compressed_messages = body.get("messages") if not _is_object_list(compressed_messages): - return self._handle_compress_failure( - messages, - "Headroom compression service response missing 'messages'", - {"body": response.text}, - ), False + return ( + self._handle_compress_failure( + messages, + "Headroom compression service response missing 'messages'", + {"body": response.text}, + ), + False, + {}, + ) filtered = [item for item in compressed_messages if _is_str_object_dict(item)] if not filtered: - return self._handle_compress_failure( - messages, - "Headroom compression service returned empty message list", - {"body": response.text}, - ), False + return ( + self._handle_compress_failure( + messages, + "Headroom compression service returned empty message list", + {"body": response.text}, + ), + False, + {}, + ) verbose_proxy_logger.debug( "Headroom: compressed %s tokens -> %s tokens (ratio %.2f)", @@ -357,7 +389,19 @@ class HeadroomGuardrail(CustomGuardrail): body.get("tokens_after", "?"), body.get("compression_ratio", 0), ) - return filtered, True + + stats = { + key: body[key] + for key in ( + "tokens_before", + "tokens_after", + "tokens_saved", + "compression_ratio", + "transforms_applied", + ) + if key in body + } + return filtered, True, stats async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str: params: dict[str, str] = {} @@ -421,14 +465,26 @@ class HeadroomGuardrail(CustomGuardrail): return inputs model = self.headroom_model or request_data.get("model") - compressed, compression_succeeded = await self._call_compress( + start_time = time.time() + compressed, compression_succeeded, stats = await self._call_compress( messages=messages, model=model if isinstance(model, str) else None, ) + end_time = time.time() if not compression_succeeded: return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType] + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=stats, + request_data=request_data, + guardrail_status="success", + guardrail_provider="headroom", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + hashes = extract_hashes_from_messages(compressed) if not hashes: return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType] diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index b155a1a7024..ae2bf3ce754 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -6,7 +6,8 @@ Code-style rules for writing tests under `tests/e2e/`. The harness already encod 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 +- `llm_translation/` - LLM endpoint and provider-translation behavior: passthrough, custom pricing, OCR, and the non-chat inference endpoints (`/v1/responses`, `/v1/messages`, `/embeddings`, `/v1/rerank`, `/v1/audio/speech`, `/v1/images/generations`), each against a deployment the test creates via `/model/new` and deletes on teardown +- `access_control/` - the gateway's authorization and error-shape contract: per-key model allow-lists, route-group permissions (`allowed_routes`), and unknown-model validation - `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 diff --git a/tests/e2e/access_control/access_control_client.py b/tests/e2e/access_control/access_control_client.py new file mode 100644 index 00000000000..d7bc9c280aa --- /dev/null +++ b/tests/e2e/access_control/access_control_client.py @@ -0,0 +1,56 @@ +"""Client for the access-control e2e suite.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from e2e_gateway import Gateway, build_gateway +from e2e_http import StreamingResponse +from models import ( + ChatBody, + ChatMessage, + KeyGenerateBody, + LiteLLMParamsBody, + ModelInfoBody, + ModelNewBody, +) + +MODEL_ACCESS_DENIED_MARKER = "key_model_access_denied" +ROUTE_NOT_ALLOWED_MARKER = "not allowed to call this route" + + +@dataclass(frozen=True, slots=True) +class AccessControlClient: + gateway: Gateway + + def llm_only_key(self) -> str: + return self.gateway.generate_key( + KeyGenerateBody(models=[], allowed_routes=["llm_api_routes"]) + ) + + def delete_key(self, key: str) -> None: + self.gateway.delete_key(key) + + def chat_status(self, key: str, model: str, content: str) -> StreamingResponse: + return self.gateway.transport.send( + "/chat/completions", + headers=self.gateway.transport.bearer(key), + json=ChatBody( + model=model, messages=[ChatMessage(role="user", content=content)] + ), + ) + + def create_model_status(self, key: str, model_name: str) -> StreamingResponse: + return self.gateway.transport.send( + "/model/new", + headers=self.gateway.transport.bearer(key), + json=ModelNewBody( + model_name=model_name, + litellm_params=LiteLLMParamsBody(model="openai/gpt-4o-mini"), + model_info=ModelInfoBody(id=model_name), + ), + ) + + +def build_client() -> AccessControlClient: + return AccessControlClient(gateway=build_gateway()) diff --git a/tests/e2e/access_control/conftest.py b/tests/e2e/access_control/conftest.py new file mode 100644 index 00000000000..9f4a00fe06f --- /dev/null +++ b/tests/e2e/access_control/conftest.py @@ -0,0 +1,10 @@ +"""Access-control suite client fixture; lifecycle/skip/marker live in the parent conftest.""" + +import pytest + +from access_control_client import AccessControlClient, build_client + + +@pytest.fixture(scope="session") +def client() -> AccessControlClient: + return build_client() diff --git a/tests/e2e/access_control/test_access_control_e2e.py b/tests/e2e/access_control/test_access_control_e2e.py new file mode 100644 index 00000000000..ce649fa2400 --- /dev/null +++ b/tests/e2e/access_control/test_access_control_e2e.py @@ -0,0 +1,83 @@ +"""Live e2e: the gateway's authorization and error-shape contract. + +A virtual key may only call models in its allow-list and route groups in its +allowed_routes; both denials are a 403 raised before any provider is touched. A +syntactically valid request naming a non-existent model is a 400 with a JSON body, +never forwarded and never a 5xx. Migrated from +litellm-regression-tests/tests/test_access_control.py: the source asserted 401 for +the disallowed-model case against an older proxy, but the current contract +(auth_checks.py) is a 403 key_model_access_denied, and the unknown-route check is +replaced by a stronger route-permission check (an llm-only key rejected from a +management route). +""" + +from __future__ import annotations + +import json + +import pytest + +from access_control_client import ( + AccessControlClient, + MODEL_ACCESS_DENIED_MARKER, + ROUTE_NOT_ALLOWED_MARKER, +) +from e2e_config import unique_marker +from lifecycle import ResourceManager + +pytestmark = pytest.mark.e2e + +ALLOWED_MODEL = "gemini-2.5-flash" +DISALLOWED_MODEL = "gpt-5.5" + + +def _is_json(body: str) -> bool: + try: + json.loads(body) + return True + except ValueError: + return False + + +class TestAccessControl: + def test_disallowed_model_is_denied_403( + self, client: AccessControlClient, resources: ResourceManager + ) -> None: + key = resources.key(models=[ALLOWED_MODEL]) + result = client.chat_status( + key, DISALLOWED_MODEL, f"capital of France? {unique_marker()}" + ) + assert result.status_code == 403, ( + f"key limited to {ALLOWED_MODEL!r} calling {DISALLOWED_MODEL!r} must be " + f"denied 403, got {result.status_code}: {result.body[:300]}" + ) + assert MODEL_ACCESS_DENIED_MARKER in result.body, ( + f"403 body must be a model-access denial, got: {result.body[:300]}" + ) + + def test_llm_only_key_forbidden_from_management_route_403( + self, client: AccessControlClient, resources: ResourceManager + ) -> None: + key = client.llm_only_key() + resources.defer(lambda: client.delete_key(key)) + result = client.create_model_status(key, f"e2e-forbidden-{unique_marker()}") + assert result.status_code == 403, ( + f"llm-only key calling a management route must be denied 403, got " + f"{result.status_code}: {result.body[:300]}" + ) + assert ROUTE_NOT_ALLOWED_MARKER in result.body, ( + f"403 body must be a route-permission denial, got: {result.body[:300]}" + ) + + def test_unknown_model_returns_400( + self, client: AccessControlClient, resources: ResourceManager + ) -> None: + key = resources.key() + result = client.chat_status( + key, f"nonexistent-model-{unique_marker()}", "hi this is a test" + ) + assert result.status_code == 400, ( + f"unknown model must be rejected 400 before forwarding, got " + f"{result.status_code}: {result.body[:300]}" + ) + assert _is_json(result.body), f"400 body must be valid JSON: {result.body[:300]}" diff --git a/tests/e2e/batches/conftest.py b/tests/e2e/batches/conftest.py index 92905b33eee..2c6070c437a 100644 --- a/tests/e2e/batches/conftest.py +++ b/tests/e2e/batches/conftest.py @@ -4,13 +4,49 @@ The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker live in the parent tests/e2e/conftest.py. BatchClient holds the shared Gateway, so the `resources` fixture cleans up keys through it; tests register file deletes and batch cancels via `resources.defer(...)`. + +Batch deployments (openai-batch, azure-batch, vertex-batch, ...) are registered +once per session via /model/new and deleted on teardown so they need not live in +the proxy config. """ +from __future__ import annotations + +from typing import Iterator + import pytest from batch_client import BatchClient, build_client +from capabilities import PROVIDERS +from e2e_http import NoBody + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "covers: registry cell a test covers, e.g. llm.batches.openai.basic.nonstream.works", + ) @pytest.fixture(scope="session") def client() -> BatchClient: return build_client() + + +@pytest.fixture(scope="session") +def batch_deployments(client: BatchClient) -> Iterator[None]: + probe = client.gateway.probe("/health/liveliness", params=NoBody()) + if not probe.healthy: + yield + return + + registered: list[str] = [] + try: + for provider in PROVIDERS: + registered.append( + client.create_model(provider.model, provider.litellm_params()) + ) + yield + finally: + for model_id in registered: + client.delete_model(model_id) diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index 1141e2f6a2f..a998f962c04 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -16,12 +16,13 @@ misroute to the wrong provider fails the create. from __future__ import annotations import json -import os import time from typing import Callable import pytest +from e2e_config import unique_marker + from batch_client import ( BatchClient, BatchCreateBody, @@ -42,16 +43,36 @@ from e2e_http import ( FileUploadForm, Result, StreamingResponse, + Success, + UnknownApiError, require_successful_call, unwrap, ) from lifecycle import ResourceManager +from models import KeyGenerateBody, SpendLogRow, SpendLogsParams pytestmark = pytest.mark.e2e CREATED_BATCH_STATUSES = {"validating", "in_progress", "finalizing"} BATCH_CANCEL_DELAY_SECONDS = 2 BATCH_TERMINAL_BEFORE_CANCEL = {"failed", "cancelled", "expired"} +BATCH_CANCEL_RETRIES = 3 + + +def cancel_batch( + client: BatchClient, batch_id: str, *, key: str, provider: str | None +) -> BatchObject: + last = client.cancel_batch(batch_id, key=key, provider=provider) + for _ in range(BATCH_CANCEL_RETRIES - 1): + match last: + case Success(data=data): + return data + case UnknownApiError(status_code=500): + time.sleep(1) + last = client.cancel_batch(batch_id, key=key, provider=provider) + case _: + break + return unwrap(last) def render_jsonl(model: str) -> bytes: @@ -146,7 +167,10 @@ def assert_batch_object(batch: BatchObject) -> None: @pytest.mark.parametrize("cap", CAPABILITIES, ids=[c.id for c in CAPABILITIES]) def test_batch_lifecycle( - cap: Capability, client: BatchClient, resources: ResourceManager + cap: Capability, + client: BatchClient, + resources: ResourceManager, + batch_deployments: None, ) -> None: key = resources.key() provider = op_provider(cap) @@ -199,11 +223,9 @@ def test_batch_lifecycle( ) if pre_cancel.status == "completed": return - cancelled = unwrap(client.cancel_batch(batch.id, key=key, provider=provider)) + cancelled = cancel_batch(client, batch.id, key=key, provider=provider) assert cancelled.id == batch.id assert cancelled.object == "batch" - # Vertex cancel is async: the job may still show its pre-cancel status - # briefly before transitioning to cancelling/cancelled. valid_post_cancel = {"cancelling", "cancelled"} if cap.provider == "vertex_ai": valid_post_cancel |= CREATED_BATCH_STATUSES @@ -213,7 +235,6 @@ def test_batch_lifecycle( if cap.can_list: listed = unwrap(client.list_batches(key=key, provider=provider)) - # OpenAI includes object="list"; Azure provider list often omits the envelope field. if listed.object is not None: assert listed.object == "list", f"list envelope object={listed.object!r}" match = next((b for b in listed.data if b.id == batch.id), None) @@ -222,7 +243,7 @@ def test_batch_lifecycle( def test_batch_key_model_access_denied( - client: BatchClient, resources: ResourceManager + client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: key = resources.key(models=["openai-batch"]) @@ -257,7 +278,7 @@ def test_batch_key_model_access_denied( def test_file_upload_and_delete_outputs( - client: BatchClient, resources: ResourceManager + client: BatchClient, resources: ResourceManager, batch_deployments: None ) -> None: key = resources.key() file = unwrap( @@ -276,14 +297,69 @@ def test_file_upload_and_delete_outputs( assert deleted.deleted is True, "file was not reported deleted" -def test_anthropic_batch_retrieve(client: BatchClient, scoped_key: str) -> None: - batch_id = os.environ.get("ANTHROPIC_BATCH_ID") - if not batch_id: - pytest.skip( - "set ANTHROPIC_BATCH_ID to a real anthropic batch id to exercise retrieve" - ) - fetched = unwrap( - client.retrieve_batch(batch_id, key=scoped_key, provider="anthropic") +def unattributed_rows(rows: list[SpendLogRow]) -> list[SpendLogRow]: + """Spend rows that carry no caller identity (empty api_key). + + Every request the proxy bills is stamped with the calling key. A row with no + api_key is one the proxy could not attribute; LIT-3266 is exactly this: the + batch rate limiter's internal input-file read ran without the batch's auth + metadata, landing a spend row with empty api_key/user. The symptom is not + tied to a single call_type, so this catches any unattributed row rather than + only a named file-content one. + """ + return [row for row in rows if not row.api_key] + + +def test_rate_limited_batch_create_leaves_no_unattributed_spend_row( + client: BatchClient, resources: ResourceManager, batch_deployments: None +) -> None: + """LIT-3266: creating a batch on a rate-limited key runs the batch rate + limiter, which reads the input file to count tokens (the limiter only reads + the file when the key has applicable rpm/tpm limits, so an unlimited key + hides the path). That internal read must carry the batch's auth metadata; + the reported gap was that it did not, spawning a spend-log row with empty + api_key/user. Create returning 200 is not a reliable signal (the read error + is swallowed), so this asserts the hygiene contract instead: the operation + introduces no new unattributed spend row. + + The key sets generous rpm/tpm limits (not a restrictive model allowlist) so + the file-read path fires while the batch itself is not blocked. + ``resources.key()`` cannot set limits, so the key is minted on the gateway + directly and its delete deferred. + """ + user_id = f"e2e-batch-rl-{unique_marker()}" + key = client.gateway.generate_key( + KeyGenerateBody(models=[], tpm_limit=1_000_000, rpm_limit=1_000, user_id=user_id) + ) + resources.defer(lambda: client.gateway.delete_key(key)) + + before = frozenset( + row.request_id for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams())) + ) + + file = unwrap( + client.upload_file( + content=render_jsonl("gpt-4o-mini"), + form=FileUploadForm(purpose="batch"), + model="openai-batch", + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + require_successful_call(created) + batch = BatchObject.model_validate_json(created.body) + resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key))) + + _ = client.gateway.poll_logs_for_key(key, min_rows=1) + + new_orphans = [ + row + for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams())) + if row.request_id not in before + ] + assert not new_orphans, ( + "batch create on a rate-limited key left an unattributed spend row " + f"(LIT-3266); rows={[(r.request_id, r.call_type, r.model) for r in new_orphans]}" ) - assert fetched.id == batch_id - assert fetched.status diff --git a/tests/e2e/llm_translation/conftest.py b/tests/e2e/llm_translation/conftest.py index fbf008cf085..2a87ef7259d 100644 --- a/tests/e2e/llm_translation/conftest.py +++ b/tests/e2e/llm_translation/conftest.py @@ -7,9 +7,22 @@ Gateway, so the `resources` fixture cleans up keys this suite creates. import pytest +from endpoints_client import EndpointsClient, build_endpoints_client from passthrough_client import PassthroughClient, build_client +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "covers: registry cell a test covers, e.g. llm.chat_completions.provider.basic.nonstream.works", + ) + + @pytest.fixture(scope="session") def client() -> PassthroughClient: return build_client() + + +@pytest.fixture(scope="session") +def endpoints_client() -> EndpointsClient: + return build_endpoints_client() diff --git a/tests/e2e/llm_translation/endpoints_client.py b/tests/e2e/llm_translation/endpoints_client.py new file mode 100644 index 00000000000..508aa3e9fc6 --- /dev/null +++ b/tests/e2e/llm_translation/endpoints_client.py @@ -0,0 +1,219 @@ +"""Client for the non-chat inference endpoints (responses, messages, rerank, +embeddings, audio speech, image generation). + +Each test registers the deployment it needs through /model/new (deleted on +teardown), so nothing is hardcoded into the gateway config, then drives the +endpoint with `send` and parses the provider-native body with a suite-local model +so the assertion is on real content, not just a 200. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from pydantic import BaseModel + +from e2e_gateway import Gateway, build_gateway +from e2e_http import NoBody, StreamingResponse, is_ok, unwrap +from models import ( + ChatMessage, + LiteLLMParamsBody, + ModelDeleteBody, + ModelInfoBody, + ModelNewBody, + ModelNewResponse, +) + + +class ResponsesRequest(BaseModel): + model: str + input: str + instructions: str | None = None + + +class MessagesRequest(BaseModel): + model: str + max_tokens: int + messages: list[ChatMessage] + + +class EmbeddingsRequest(BaseModel): + model: str + input: str + + +class RerankRequest(BaseModel): + model: str + query: str + documents: list[str] + top_n: int + + +class SpeechRequest(BaseModel): + model: str + input: str + voice: str + + +class ImageRequest(BaseModel): + model: str + prompt: str + n: int = 1 + size: str = "1024x1024" + + +class ResponsesOutputContent(BaseModel): + type: str | None = None + text: str | None = None + + +class ResponsesOutputItem(BaseModel): + type: str | None = None + content: list[ResponsesOutputContent] = [] + + +class ResponsesResult(BaseModel): + id: str | None = None + status: str | None = None + model: str | None = None + output: list[ResponsesOutputItem] = [] + + @property + def text(self) -> str: + return "".join( + content.text or "" for item in self.output for content in item.content + ) + + +class AnthropicContentBlock(BaseModel): + type: str | None = None + text: str | None = None + + +class MessagesResult(BaseModel): + id: str | None = None + role: str | None = None + model: str | None = None + content: list[AnthropicContentBlock] = [] + + @property + def text(self) -> str: + return "".join(block.text or "" for block in self.content) + + +class EmbeddingItem(BaseModel): + embedding: list[float] = [] + + +class EmbeddingsResult(BaseModel): + data: list[EmbeddingItem] = [] + + @property + def first_vector(self) -> tuple[float, ...]: + return tuple(self.data[0].embedding) if self.data else () + + +class RerankItem(BaseModel): + index: int | None = None + relevance_score: float | None = None + + +class RerankResult(BaseModel): + results: list[RerankItem] = [] + + +class ImageItem(BaseModel): + url: str | None = None + b64_json: str | None = None + + +class ImagesResult(BaseModel): + data: list[ImageItem] = [] + + +@dataclass(frozen=True, slots=True) +class EndpointsClient: + gateway: Gateway + + def create_model(self, model_name: str, litellm_params: LiteLLMParamsBody) -> str: + """Register a deployment under `model_name` (id == model_name) and return the + model_id. add_deployment runs synchronously in /model/new, so the model is + callable as soon as this returns.""" + return unwrap( + self.gateway.transport.post( + "/model/new", + headers=self.gateway.transport.master, + json=ModelNewBody( + model_name=model_name, + litellm_params=litellm_params, + model_info=ModelInfoBody(id=model_name), + ), + response_type=ModelNewResponse, + ) + ).model_id + + def delete_model(self, model_id: str) -> None: + result = self.gateway.transport.post( + "/model/delete", + headers=self.gateway.transport.master, + json=ModelDeleteBody(id=model_id), + response_type=NoBody, + ) + if not is_ok(result): + import warnings + warnings.warn(f"delete_model({model_id!r}) failed: {result}", stacklevel=2) + + def _send(self, path: str, key: str, body: BaseModel) -> StreamingResponse: + return self.gateway.transport.send( + path, headers=self.gateway.transport.bearer(key), json=body + ) + + def responses(self, key: str, model: str, text: str) -> StreamingResponse: + return self._send( + "/v1/responses", + key, + ResponsesRequest( + model=model, input=text, instructions="You are a helpful assistant" + ), + ) + + def messages( + self, key: str, model: str, text: str, *, max_tokens: int = 64 + ) -> StreamingResponse: + return self._send( + "/v1/messages", + key, + MessagesRequest( + model=model, + max_tokens=max_tokens, + messages=[ChatMessage(role="user", content=text)], + ), + ) + + def embeddings(self, key: str, model: str, text: str) -> StreamingResponse: + return self._send("/embeddings", key, EmbeddingsRequest(model=model, input=text)) + + def rerank( + self, key: str, model: str, query: str, documents: list[str], top_n: int + ) -> StreamingResponse: + return self._send( + "/v1/rerank", + key, + RerankRequest(model=model, query=query, documents=documents, top_n=top_n), + ) + + def audio_speech( + self, key: str, model: str, text: str, *, voice: str = "alloy" + ) -> StreamingResponse: + return self._send( + "/v1/audio/speech", key, SpeechRequest(model=model, input=text, voice=voice) + ) + + def images(self, key: str, model: str, prompt: str) -> StreamingResponse: + return self._send( + "/v1/images/generations", key, ImageRequest(model=model, prompt=prompt) + ) + + +def build_endpoints_client() -> EndpointsClient: + return EndpointsClient(gateway=build_gateway()) diff --git a/tests/e2e/llm_translation/test_audio_speech_e2e.py b/tests/e2e/llm_translation/test_audio_speech_e2e.py new file mode 100644 index 00000000000..f7a04d94cb3 --- /dev/null +++ b/tests/e2e/llm_translation/test_audio_speech_e2e.py @@ -0,0 +1,40 @@ +"""Live e2e: POST /v1/audio/speech returns audio. + +Registers an OpenAI text-to-speech deployment at runtime and asserts the response +is an audio body (binary, not JSON). Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class TestAudioSpeech: + def test_audio_speech_returns_audio( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-speech-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-4o-mini-tts", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.audio_speech(key, model, "Hello!") + require_successful_call(result) + assert "audio" in (result.content_type or ""), ( + f"/audio/speech content-type is not audio: {result.content_type!r}" + ) + assert result.body, "/audio/speech returned an empty body" diff --git a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py new file mode 100644 index 00000000000..269cb5d6d22 --- /dev/null +++ b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py @@ -0,0 +1,58 @@ +"""Live regression net for /chat/completions across the configured providers. + +GH #28991 broke /chat/completions (and /responses) for most models on some +releases: a clean 200 came back but with no real completion. A status check +alone would not have caught it, so each case here asserts the product promise - +a non-empty assistant message and a real model name in the body - across the +three providers wired into the gateway config (OpenAI, Anthropic, Gemini). A +regression that empties the completion for any provider fails that provider's +row here. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import unwrap +from models import ChatBody, ChatMessage +from passthrough_client import PassthroughClient + +pytestmark = pytest.mark.e2e + +CHAT_MODELS: tuple[tuple[str, str], ...] = ( + ("gpt-5.5", "openai"), + ("claude-haiku-4-5", "anthropic"), + ("gemini-2.5-flash", "gemini"), +) + + +class TestChatCompletionsRegression: + @pytest.mark.parametrize( + ("model", "route"), + CHAT_MODELS, + ids=[f"{model}-{route}" for model, route in CHAT_MODELS], + ) + @pytest.mark.covers("llm.chat_completions.provider.basic.nonstream.works", exercised_on=[]) + def test_chat_returns_real_completion( + self, client: PassthroughClient, scoped_key: str, model: str, route: str + ) -> None: + response = unwrap( + client.gateway.chat( + scoped_key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content=f"reply with one word {unique_marker()}") + ], + max_tokens=512, + ), + ) + ) + + assert response.model, f"{model} ({route}): response carried no model name: {response}" + assert response.choices, f"{model} ({route}): response had no choices: {response}" + message = response.choices[0].message + assert message is not None and message.content and message.content.strip(), ( + f"{model} ({route}): 200 with an empty completion (#28991): {response}" + ) diff --git a/tests/e2e/llm_translation/test_custom_pricing_e2e.py b/tests/e2e/llm_translation/test_custom_pricing_e2e.py index 58faab61aac..7894b447be9 100644 --- a/tests/e2e/llm_translation/test_custom_pricing_e2e.py +++ b/tests/e2e/llm_translation/test_custom_pricing_e2e.py @@ -1,52 +1,48 @@ -"""Live e2e: a model's custom per-token pricing is loaded, billed, and isolated. +"""Live e2e: a deployment's custom per-token pricing is loaded, billed, and isolated. -The gateway config declares ``custom-priced-flash`` (gemini-2.5-flash underneath) -with input/output rates deliberately far above the canonical gemini price, read -back here from the same config file. Three behaviors are checked independently: +Each test registers the deployment(s) it needs through /model/new (deleted on +teardown) instead of relying on a statically configured model, so the check is +self-contained and never inherits pricing another suite or a stale config left on +the shared proxy. custom-priced-flash sets input/output rates deliberately far +above the canonical gemini price; the isolation sibling shares the same +gemini/gemini-2.5-flash backend but sets no override. Three behaviors are checked +independently: - billing: a real call's logged cost breakdown charges input and output tokens at the custom rates, each component checked separately (a base-rate bill lands ~100x lower; a swapped input/output rate passes a total-only check but not this) -- reporting: /model/info surfaces those rates for the model -- isolation: gemini-2.5-flash shares the same underlying gemini/gemini-2.5-flash - but sets no override, so it must keep its own price; an override that leaks into - the shared cost map misprices it. A regression that reintroduces that leak makes - the sibling's rate match the custom one and fails the isolation check. +- reporting: /model/info surfaces those rates for the deployment +- isolation: the sibling keeps its own price; an override that leaks into the + shared backend cost map (LIT-3897) misprices it, making the sibling's rate match + the custom one and failing the isolation check """ import time -from dataclasses import dataclass -from pathlib import Path import pytest -import yaml from pydantic import BaseModel, RootModel from e2e_config import unique_marker +from e2e_gateway import Gateway from e2e_http import Success, unwrap -from models import ChatBody, ChatMessage, CustomPricing, ModelInfoEntry, SpendLogsParams -from passthrough_client import PassthroughClient +from endpoints_client import EndpointsClient +from lifecycle import ResourceManager +from models import ( + ChatBody, + ChatMessage, + LiteLLMParamsBody, + ModelInfoEntry, + SpendLogsParams, +) pytestmark = pytest.mark.e2e -CUSTOM_MODEL = "custom-priced-flash" -BASE_MODEL = "gemini-2.5-flash" -CONFIG_PATH = Path(__file__).resolve().parents[1] / "gateway" / "litellm-config.yml" - - -@dataclass(frozen=True, slots=True) -class _Rates: - input_per_token: float - output_per_token: float - - -class _ConfiguredModel(BaseModel): - model_name: str - litellm_params: CustomPricing - - -class _GatewayConfig(BaseModel): - model_list: list[_ConfiguredModel] +BACKEND_MODEL = "gemini/gemini-2.5-flash" +GEMINI_API_KEY = "os.environ/GEMINI_API_KEY" +# Deliberately ~100x above canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6) +# so an override that is ignored or under-applied bills at the base rate and fails. +CUSTOM_INPUT_RATE = 5e-05 +CUSTOM_OUTPUT_RATE = 1e-04 class _CostBreakdown(BaseModel): @@ -74,40 +70,60 @@ def _approx_equal(actual: float, expected: float) -> bool: return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2) -def _configured_pricing(model_name: str) -> _Rates: - """The custom rates declared for `model_name` in the gateway config the proxy - runs with - the source of truth the billed and reported prices are checked - against.""" - config = _GatewayConfig.model_validate(yaml.safe_load(CONFIG_PATH.read_text())) - for entry in config.model_list: - if entry.model_name == model_name: - pricing = entry.litellm_params - assert pricing.input_cost_per_token and pricing.output_cost_per_token, ( - f"{model_name} declares no custom per-token rates in {CONFIG_PATH.name}" - ) - return _Rates(pricing.input_cost_per_token, pricing.output_cost_per_token) - pytest.fail(f"{model_name} not found in {CONFIG_PATH.name}") +def _provision( + endpoints_client: EndpointsClient, + resources: ResourceManager, + prefix: str, + *, + input_cost_per_token: float | None, + output_cost_per_token: float | None, +) -> str: + """Register a fresh gemini/gemini-2.5-flash deployment (deleted on teardown) and + return its model name. With the cost fields set the deployment carries a custom + pricing override; with them None it is a plain sibling on the same backend. The + marker keeps the name unique so concurrent runs on the shared proxy never + collide.""" + model_name = f"{prefix}-{unique_marker()}" + model_id = endpoints_client.create_model( + model_name, + LiteLLMParamsBody( + model=BACKEND_MODEL, + api_key=GEMINI_API_KEY, + input_cost_per_token=input_cost_per_token, + output_cost_per_token=output_cost_per_token, + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + return model_name -def _model_info_entry( - entries: list[ModelInfoEntry], model_name: str -) -> ModelInfoEntry: +def _provision_custom_priced( + endpoints_client: EndpointsClient, resources: ResourceManager +) -> str: + return _provision( + endpoints_client, + resources, + "custom-priced-flash", + input_cost_per_token=CUSTOM_INPUT_RATE, + output_cost_per_token=CUSTOM_OUTPUT_RATE, + ) + + +def _model_info_entry(entries: list[ModelInfoEntry], model_name: str) -> ModelInfoEntry: for entry in entries: if entry.model_name == model_name: return entry pytest.fail(f"{model_name} absent from /model/info; the override did not load") -def _poll_breakdown_row( - client: PassthroughClient, key: str, response_id: str | None -) -> _SpendRow: +def _poll_breakdown_row(gateway: Gateway, key: str, response_id: str | None) -> _SpendRow: """Poll /spend/logs until the call's row lands with a cost breakdown (rows flush ~60s behind the call via proxy_batch_write_at).""" - deadline = time.monotonic() + client.gateway.poll_timeout + deadline = time.monotonic() + gateway.poll_timeout while time.monotonic() < deadline: - result = client.gateway.transport.get( + result = gateway.transport.get( "/spend/logs", - headers=client.gateway.transport.master, + headers=gateway.transport.master, params=SpendLogsParams(api_key=key), response_type=_SpendRows, ) @@ -128,85 +144,99 @@ def _poll_breakdown_row( return row if priced and response_id is None: return priced[0] - time.sleep(client.gateway.poll_interval) + time.sleep(gateway.poll_interval) pytest.fail("no spend row with a cost breakdown landed before the deadline") -def test_custom_pricing_is_billed_at_configured_rate( - client: PassthroughClient, scoped_key: str -) -> None: - rates = _configured_pricing(CUSTOM_MODEL) +class TestCustomPricing: + def test_custom_pricing_is_billed_at_configured_rate( + self, + endpoints_client: EndpointsClient, + resources: ResourceManager, + scoped_key: str, + ) -> None: + model = _provision_custom_priced(endpoints_client, resources) - chat = unwrap( - client.gateway.chat( - scoped_key, - ChatBody( - model=CUSTOM_MODEL, - messages=[ - ChatMessage( - role="user", content=f"reply with one word {unique_marker()}" - ) - ], - max_tokens=16, - ), + chat = unwrap( + endpoints_client.gateway.chat( + scoped_key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", content=f"reply with one word {unique_marker()}" + ) + ], + max_tokens=16, + ), + ) ) - ) - row = _poll_breakdown_row(client, scoped_key, chat.id) - assert row.metadata and row.metadata.cost_breakdown # guaranteed by the poll - breakdown = row.metadata.cost_breakdown + row = _poll_breakdown_row(endpoints_client.gateway, scoped_key, chat.id) + assert row.metadata and row.metadata.cost_breakdown # guaranteed by the poll + breakdown = row.metadata.cost_breakdown - prompt = row.prompt_tokens or 0 - completion = row.completion_tokens or 0 - assert prompt > 0 and completion > 0, f"call tokens not logged on the row: {row}" + prompt = row.prompt_tokens or 0 + completion = row.completion_tokens or 0 + assert prompt > 0 and completion > 0, f"call tokens not logged on the row: {row}" - input_cost = breakdown.input_cost - output_cost = breakdown.output_cost - assert input_cost is not None and output_cost is not None, ( - f"row cost breakdown missing input/output cost: {breakdown}" - ) - assert _approx_equal(input_cost, prompt * rates.input_per_token), ( - f"input_cost {input_cost} != {prompt} tokens * {rates.input_per_token} " - f"= {prompt * rates.input_per_token}" - ) - assert _approx_equal(output_cost, completion * rates.output_per_token), ( - f"output_cost {output_cost} != {completion} tokens * {rates.output_per_token} " - f"= {completion * rates.output_per_token}" - ) + input_cost = breakdown.input_cost + output_cost = breakdown.output_cost + assert input_cost is not None and output_cost is not None, ( + f"row cost breakdown missing input/output cost: {breakdown}" + ) + assert _approx_equal(input_cost, prompt * CUSTOM_INPUT_RATE), ( + f"input_cost {input_cost} != {prompt} tokens * {CUSTOM_INPUT_RATE} " + f"= {prompt * CUSTOM_INPUT_RATE}" + ) + assert _approx_equal(output_cost, completion * CUSTOM_OUTPUT_RATE), ( + f"output_cost {output_cost} != {completion} tokens * {CUSTOM_OUTPUT_RATE} " + f"= {completion * CUSTOM_OUTPUT_RATE}" + ) + def test_model_info_reports_custom_pricing( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = _provision_custom_priced(endpoints_client, resources) + entry = _model_info_entry(endpoints_client.gateway.model_info(), model) -def test_model_info_reports_custom_pricing(client: PassthroughClient) -> None: - rates = _configured_pricing(CUSTOM_MODEL) - entry = _model_info_entry(client.gateway.model_info(), CUSTOM_MODEL) + assert entry.litellm_params.input_cost_per_token == CUSTOM_INPUT_RATE, ( + f"/model/info litellm_params input rate " + f"{entry.litellm_params.input_cost_per_token} != configured {CUSTOM_INPUT_RATE}" + ) + assert entry.litellm_params.output_cost_per_token == CUSTOM_OUTPUT_RATE, ( + f"/model/info litellm_params output rate " + f"{entry.litellm_params.output_cost_per_token} != configured {CUSTOM_OUTPUT_RATE}" + ) - assert entry.litellm_params.input_cost_per_token == rates.input_per_token, ( - f"/model/info litellm_params input rate " - f"{entry.litellm_params.input_cost_per_token} != configured " - f"{rates.input_per_token}" - ) - assert entry.litellm_params.output_cost_per_token == rates.output_per_token, ( - f"/model/info litellm_params output rate " - f"{entry.litellm_params.output_cost_per_token} != configured " - f"{rates.output_per_token}" - ) + def test_custom_pricing_is_isolated_from_sibling_deployment( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + # Register the override first so its rate is in the backend cost map before + # the sibling resolves; a leak (LIT-3897) would then poison the sibling. + custom = _provision_custom_priced(endpoints_client, resources) + sibling = _provision( + endpoints_client, + resources, + "base-flash", + input_cost_per_token=None, + output_cost_per_token=None, + ) + entries = {entry.model_name: entry for entry in endpoints_client.gateway.model_info()} + custom_entry = entries.get(custom) + sibling_entry = entries.get(sibling) + assert custom_entry is not None, f"{custom} absent from /model/info" + assert sibling_entry is not None, f"{sibling} absent from /model/info" -def test_custom_pricing_is_isolated_from_sibling_deployment( - client: PassthroughClient, -) -> None: - entries = {entry.model_name: entry for entry in client.gateway.model_info()} - custom = entries.get(CUSTOM_MODEL) - base = entries.get(BASE_MODEL) - assert custom is not None, f"{CUSTOM_MODEL} absent from /model/info" - assert base is not None, f"{BASE_MODEL} absent from /model/info" - - # custom-priced-flash overrides pricing; gemini-2.5-flash shares the same - # underlying gemini/gemini-2.5-flash but sets no override, so it must keep its - # own price. Equal rates mean the override leaked into the shared cost map. - assert ( - base.model_info.input_cost_per_token != custom.model_info.input_cost_per_token - ), ( - f"{BASE_MODEL} input rate {base.model_info.input_cost_per_token} matches " - f"{CUSTOM_MODEL}'s override {custom.model_info.input_cost_per_token}; " - f"per-deployment custom pricing is not isolated" - ) + # custom-priced-flash overrides pricing; the sibling shares the same + # gemini/gemini-2.5-flash backend but sets no override, so it must keep its + # own price. Equal rates mean the override leaked into the shared cost map. + assert ( + sibling_entry.model_info.input_cost_per_token + != custom_entry.model_info.input_cost_per_token + ), ( + f"{sibling} input rate {sibling_entry.model_info.input_cost_per_token} matches " + f"{custom}'s override {custom_entry.model_info.input_cost_per_token}; " + f"per-deployment custom pricing is not isolated" + ) diff --git a/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py b/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py new file mode 100644 index 00000000000..f8f229aa2a7 --- /dev/null +++ b/tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py @@ -0,0 +1,133 @@ +"""Live e2e: DeepSeek reasoner honors a request to turn reasoning OFF. + +DeepSeek's reasoner defaults thinking ON and surfaces the chain as +``message.reasoning_content``. Two documented ways to disable it are +``reasoning_effort="none"`` and ``thinking={"type": "disabled"}``. Today the +DeepSeek param mapper (``litellm/llms/deepseek/chat/transformation.py`` +``map_openai_params``) drops both without forwarding any disable signal, so the +outbound body carries no ``thinking`` key and DeepSeek keeps thinking on; the +response still comes back with ``reasoning_content``. That is the product gap +tracked by LIT-3686 / GH #27453. + +The control case proves the model and path work (reasoning is returned when +nothing asks to disable it), so the two disable assertions are meaningful. Those +two are marked xfail(strict) until the mapper forwards a real disable signal; an +xpass then alerts that the fix landed. + +Requires DEEPSEEK_API_KEY on the proxy (tests/e2e/.env). No skip gate: once the +proxy is up, a failure here is real, per the suite's hard-fail contract. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody, ThinkingParam +from passthrough_client import PassthroughClient + +pytestmark = pytest.mark.e2e + +REASONER = "deepseek/deepseek-reasoner" +PROMPT = "What is 17 + 26? Answer with just the number." + + +def _register_reasoner(client: PassthroughClient, resources: ResourceManager) -> str: + model = f"e2e-deepseek-reasoner-{unique_marker()}" + model_id = client.gateway.create_model( + model, + LiteLLMParamsBody(model=REASONER, api_key="os.environ/DEEPSEEK_API_KEY"), + ) + resources.defer(lambda: client.gateway.delete_model(model_id)) + return model + + +def _reasoning_content(response: ChatResponse) -> str | None: + assert response.choices, f"reasoner returned no choices: {response}" + message = response.choices[0].message + assert message is not None, f"reasoner choice has no message: {response}" + return message.reasoning_content + + +class TestDeepSeekReasoningDisable: + def test_reasoner_returns_reasoning_by_default( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_reasoner(client, resources) + key = resources.key() + + response = unwrap( + client.gateway.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=PROMPT)], + max_tokens=64, + ), + ) + ) + reasoning = _reasoning_content(response) + assert reasoning, ( + "control case: deepseek-reasoner returned no reasoning_content with no " + f"disable param, so the disable assertions below can't be trusted: {response}" + ) + + @pytest.mark.xfail( + strict=True, + reason=( + "LIT-3686 / GH #27453: DeepSeek reasoning_effort='none' and " + "thinking type='disabled' are silently dropped; reasoning not disabled" + ), + ) + def test_reasoning_effort_none_disables_reasoning( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_reasoner(client, resources) + key = resources.key() + + response = unwrap( + client.gateway.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=PROMPT)], + max_tokens=64, + reasoning_effort="none", + ), + ) + ) + assert not _reasoning_content(response), ( + "reasoning_effort='none' must disable reasoning, but reasoning_content " + f"is still present: {response}" + ) + + @pytest.mark.xfail( + strict=True, + reason=( + "LIT-3686 / GH #27453: DeepSeek reasoning_effort='none' and " + "thinking type='disabled' are silently dropped; reasoning not disabled" + ), + ) + def test_thinking_disabled_disables_reasoning( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_reasoner(client, resources) + key = resources.key() + + response = unwrap( + client.gateway.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=PROMPT)], + max_tokens=64, + thinking=ThinkingParam(type="disabled"), + ), + ) + ) + assert not _reasoning_content(response), ( + "thinking={'type': 'disabled'} must disable reasoning, but " + f"reasoning_content is still present: {response}" + ) diff --git a/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py new file mode 100644 index 00000000000..56f2de8bd4f --- /dev/null +++ b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py @@ -0,0 +1,42 @@ +"""Live e2e: POST /embeddings returns a real vector. + +Registers an OpenAI embedding deployment at runtime and asserts a non-empty, +non-zero vector came back. Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py; the LIT-3167 guard in +tests/e2e/embeddings/ covers the Gemini embedding path. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EmbeddingsResult, EndpointsClient +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class TestEmbeddingsEndpoint: + def test_embeddings_returns_vector( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-embeddings-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/text-embedding-3-small", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.embeddings(key, model, "Say this is a test!") + require_successful_call(result) + parsed = EmbeddingsResult.model_validate_json(result.body) + assert parsed.first_vector, f"/embeddings returned no vector: {result.body[:300]}" + assert any(component != 0.0 for component in parsed.first_vector), ( + f"embedding vector is all zeros: {result.body[:300]}" + ) diff --git a/tests/e2e/llm_translation/test_image_generation_e2e.py b/tests/e2e/llm_translation/test_image_generation_e2e.py new file mode 100644 index 00000000000..4d2211f3be4 --- /dev/null +++ b/tests/e2e/llm_translation/test_image_generation_e2e.py @@ -0,0 +1,42 @@ +"""Live e2e: POST /v1/images/generations returns an image. + +Registers an OpenAI image deployment at runtime and asserts the response carries a +generated image (url or base64). Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient, ImagesResult +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class TestImageGeneration: + def test_image_generation_returns_image( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-image-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-image-1-mini", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.images(key, model, "Draw a cute cat") + require_successful_call(result) + parsed = ImagesResult.model_validate_json(result.body) + assert parsed.data, f"/images/generations returned no data: {result.body[:300]}" + first = parsed.data[0] + assert first.b64_json or first.url, ( + f"generated image has neither b64_json nor url: {result.body[:300]}" + ) diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py new file mode 100644 index 00000000000..b0a48f22118 --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -0,0 +1,39 @@ +"""Live e2e: POST /v1/messages (Anthropic Messages API) returns a real completion. + +Registers an Anthropic deployment at runtime, drives the Messages endpoint through +the gateway, and asserts an assistant message with text came back. Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient, MessagesResult +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class TestAnthropicMessages: + def test_messages_returns_completion( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-messages-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="anthropic/claude-haiku-4-5", api_key="os.environ/ANTHROPIC_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.messages(key, model, "reply with one word") + require_successful_call(result) + parsed = MessagesResult.model_validate_json(result.body) + assert parsed.role == "assistant", f"unexpected role: {result.body[:300]}" + assert parsed.text.strip(), f"/v1/messages returned no text: {result.body[:300]}" diff --git a/tests/e2e/llm_translation/test_provider_features_e2e.py b/tests/e2e/llm_translation/test_provider_features_e2e.py new file mode 100644 index 00000000000..9c99c1be161 --- /dev/null +++ b/tests/e2e/llm_translation/test_provider_features_e2e.py @@ -0,0 +1,148 @@ +"""Live e2e for model-specific request features: service_tier and prompt caching. + +Each case asserts the feature took effect, not just a 200. + +service_tier is an OpenAI concept. The proxy forwards it and the provider echoes +the tier back on the response, so sending a non-default tier ("flex") and reading +it back off ``service_tier`` proves the param was honored end to end; litellm's own +default injection would report "default", so a "flex" echo can only come from the +request being forwarded. Bedrock and Vertex do not accept service_tier, so that +cell is OpenAI-only by design. + +Prompt caching is asserted through provider prompt-cache usage tokens. The +deterministic path is explicit ``cache_control`` on an Anthropic-family model +(here Bedrock's Claude): a large cacheable prefix is sent twice and the second +call must report ``cache_read_input_tokens > 0``. OpenAI and Gemini only offer +implicit automatic caching, which does not deterministically produce a cache read +within a test window (verified: repeated >3k-token prompts kept +``prompt_tokens_details.cached_tokens`` at 0), so those caching cells are out of +scope here and covered only by the explicit-cache-control Bedrock case. +""" + +from __future__ import annotations + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody +from passthrough_client import PassthroughClient + +pytestmark = pytest.mark.e2e + +SERVICE_TIER = "flex" +CACHE_MIN_READ_TOKENS = 1 + + +class CacheControl(BaseModel): + type: str = "ephemeral" + + +class CacheTextBlock(BaseModel): + type: str = "text" + text: str + cache_control: CacheControl | None = None + + +class RichMessage(BaseModel): + role: str + content: list[CacheTextBlock] + + +class CacheChatBody(BaseModel): + model: str + messages: list[RichMessage] + max_tokens: int + + +def cacheable_prefix() -> str: + return ( + "You are a policy compliance auditor. The following corpus is the immutable " + "reference the assistant must consult on every turn. " + ) + ("Clause: obey all safety, formatting, and citation rules exactly. " * 400) + + +def post_chat(client: PassthroughClient, key: str, body: BaseModel) -> ChatResponse: + return unwrap( + client.gateway.transport.post( + "/chat/completions", + headers=client.gateway.transport.bearer(key), + json=body, + response_type=ChatResponse, + ) + ) + + +class TestServiceTier: + @pytest.mark.covers("llm.chat_completions.openai.service_tier.works", exercised_on=[]) + def test_openai_service_tier_is_echoed( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = f"e2e-service-tier-{unique_marker()}" + model_id = client.gateway.create_model( + model, LiteLLMParamsBody(model="openai/gpt-5.5", api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.gateway.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.gateway.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content="reply with one word")], + max_tokens=64, + service_tier=SERVICE_TIER, + ), + ) + ) + assert response.service_tier == SERVICE_TIER, ( + f"service_tier not honored: sent {SERVICE_TIER!r}, response reported " + f"{response.service_tier!r} ({response})" + ) + + +class TestPromptCaching: + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.cache_hit", exercised_on=[] + ) + def test_bedrock_cache_control_produces_cache_read( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = f"e2e-bedrock-cache-{unique_marker()}" + model_id = client.gateway.create_model( + model, + LiteLLMParamsBody( + model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + aws_region_name="us-east-1", + ), + ) + resources.defer(lambda: client.gateway.delete_model(model_id)) + key = resources.key() + + body = CacheChatBody( + model=model, + max_tokens=32, + messages=[ + RichMessage( + role="user", + content=[ + CacheTextBlock(text=cacheable_prefix(), cache_control=CacheControl()), + CacheTextBlock(text="Answer in one word: acknowledged?"), + ], + ) + ], + ) + + first = post_chat(client, key, body) + assert first.usage is not None, f"first call reported no usage: {first}" + + second = post_chat(client, key, body) + assert second.usage is not None, f"second call reported no usage: {second}" + cache_read = second.usage.cache_read_input_tokens + assert cache_read is not None and cache_read >= CACHE_MIN_READ_TOKENS, ( + "second identical request did not read the prompt cache: " + f"cache_read_input_tokens={cache_read!r} (usage={second.usage})" + ) diff --git a/tests/e2e/llm_translation/test_rerank_e2e.py b/tests/e2e/llm_translation/test_rerank_e2e.py new file mode 100644 index 00000000000..4b30ac1ea5c --- /dev/null +++ b/tests/e2e/llm_translation/test_rerank_e2e.py @@ -0,0 +1,49 @@ +"""Live e2e: POST /v1/rerank ranks documents by relevance. + +Registers a Cohere rerank deployment at runtime and asserts the endpoint returns +scored results within the requested top_n. Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient, RerankResult +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + +DOCUMENTS = [ + "Carson City is the capital city of the American state of Nevada.", + "The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean.", + "Washington, D.C. is the capital of the United States.", + "Capital punishment has existed in the United States since before it was a country.", +] + + +class TestRerank: + def test_rerank_scores_top_n( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-rerank-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody(model="cohere/rerank-v3.5", api_key="os.environ/COHERE_API_KEY"), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.rerank( + key, model, "What is the capital of the United States?", DOCUMENTS, top_n=3 + ) + require_successful_call(result) + parsed = RerankResult.model_validate_json(result.body) + assert parsed.results, f"/rerank returned no results: {result.body[:300]}" + assert len(parsed.results) <= 3, f"top_n=3 not honored: {result.body[:300]}" + assert parsed.results[0].relevance_score is not None, ( + f"top rerank result has no relevance_score: {result.body[:300]}" + ) diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py new file mode 100644 index 00000000000..743de79880f --- /dev/null +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -0,0 +1,36 @@ +"""Live e2e: POST /v1/responses returns a real completion. + +Registers an OpenAI deployment at runtime, drives the Responses API through the +gateway, and asserts output text came back. Migrated from +litellm-regression-tests/tests/test_inference_endpoints.py. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient, ResponsesResult +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class TestResponses: + def test_responses_returns_completion( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-responses-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY"), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.responses(key, model, "reply with one word") + require_successful_call(result) + parsed = ResponsesResult.model_validate_json(result.body) + assert parsed.text.strip(), f"/responses returned no output text: {result.body[:300]}" diff --git a/tests/e2e/logging/conftest.py b/tests/e2e/logging/conftest.py new file mode 100644 index 00000000000..40c19aefca7 --- /dev/null +++ b/tests/e2e/logging/conftest.py @@ -0,0 +1,37 @@ +"""Fixtures for the Datadog logging suite. + +These tests drive the Datadog batch-send path (#25663) directly against the real +Datadog logs intake with synthetic events - no LLM calls, no proxy, no log +read-back - so they need only the shipping credentials DD_API_KEY + DD_SITE +(DD_SERVICE is an optional tag). No Datadog Application key is required, and they +skip when the shipping credentials are absent from the environment. +""" + +import os + +import pytest + +from logging_client import LoggingClient, build_logging_client + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "covers: registry cell a test covers, e.g. logging.datadog.success.writes_object", + ) + + +@pytest.fixture(scope="session") +def client() -> LoggingClient: + """The logging suite's client: holds the shared Gateway so `resources` / + `scoped_key` clean up keys, and adds `/metrics` scraping.""" + return build_logging_client() + + +@pytest.fixture +def datadog_creds() -> None: + """Gate the suite on the Datadog shipping credentials. The DataDogLogger is built + inside each async test, not here, because its __init__ schedules a periodic-flush + task via asyncio.create_task and so needs a running event loop.""" + if not (os.getenv("DD_API_KEY") and os.getenv("DD_SITE")): + pytest.skip("set DD_API_KEY and DD_SITE to run the Datadog logging suite") diff --git a/tests/e2e/logging/logging_client.py b/tests/e2e/logging/logging_client.py new file mode 100644 index 00000000000..a3213fbdb00 --- /dev/null +++ b/tests/e2e/logging/logging_client.py @@ -0,0 +1,48 @@ +"""Client for the logging e2e suite: drive traffic and scrape the proxy's +Prometheus ``/metrics`` endpoint. + +Holds the shared Gateway so the ``resources`` fixture cleans up keys it creates. +``/metrics`` is exposed as plaintext (not a typed JSON body), so scraping goes +through ``transport.probe`` and returns the raw exposition text for a Prometheus +parser to read. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from e2e_gateway import Gateway, build_gateway +from e2e_http import NoBody, unwrap +from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody + + +@dataclass(frozen=True, slots=True) +class LoggingClient: + gateway: Gateway + + def key_with_alias(self, alias: str, *, models: list[str]) -> str: + return self.gateway.generate_key( + KeyGenerateBody(key_alias=alias, models=models, user_id=f"e2e-{alias}") + ) + + def delete_key(self, key: str) -> None: + self.gateway.delete_key(key) + + def chat(self, key: str, model: str, text: str) -> ChatResponse: + return unwrap( + self.gateway.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=text)], + max_tokens=64, + ), + ) + ) + + def scrape_metrics(self) -> str: + return self.gateway.probe("/metrics", params=NoBody()).body + + +def build_logging_client() -> LoggingClient: + return LoggingClient(gateway=build_gateway()) diff --git a/tests/e2e/logging/test_prometheus_cardinality_e2e.py b/tests/e2e/logging/test_prometheus_cardinality_e2e.py new file mode 100644 index 00000000000..163293a3009 --- /dev/null +++ b/tests/e2e/logging/test_prometheus_cardinality_e2e.py @@ -0,0 +1,70 @@ +"""Live e2e: Prometheus request metrics grow one series per virtual key. + +The proxy exposes ``/metrics`` (prometheus is in the callbacks and +``require_auth_for_metrics_endpoint`` is off in the e2e config). The counter +``litellm_requests_metric_total`` carries an ``api_key_alias`` label, so driving +traffic through keys with distinct aliases must produce a distinct labeled series +per alias. This is the per-key cardinality contract: a regression that stops +stamping ``api_key_alias`` (or collapses every key onto one series) would drop +the aliases and fail here. + +Scraping goes through ``transport.probe`` (raw text) and is parsed with +prometheus_client; the metric is eventually consistent (it increments on the +success-logging callback), so the scrape polls to a deadline. +""" + +from __future__ import annotations + +import time + +import pytest +from prometheus_client.parser import text_string_to_metric_families + +from e2e_config import unique_marker +from lifecycle import ResourceManager +from logging_client import LoggingClient + +pytestmark = pytest.mark.e2e + +DRIVER_MODEL = "gemini-2.5-flash" +REQUESTS_METRIC = "litellm_requests_metric_total" +ALIAS_LABEL = "api_key_alias" +DISTINCT_KEYS = 3 + + +def _aliases_in_metric(exposition: str, metric: str, label: str) -> frozenset[str]: + """The set of ``label`` values present on ``metric`` samples in a scrape.""" + return frozenset( + sample.labels[label] + for family in text_string_to_metric_families(exposition) + for sample in family.samples + if sample.name == metric and label in sample.labels + ) + + +class TestPrometheusPerKeyCardinality: + @pytest.mark.covers("logging.prometheus.success.exports_metric", exercised_on=[]) + def test_distinct_key_aliases_produce_distinct_series( + self, client: LoggingClient, resources: ResourceManager + ) -> None: + aliases = tuple(f"e2e-prom-{unique_marker()}" for _ in range(DISTINCT_KEYS)) + for alias in aliases: + key = client.key_with_alias(alias, models=[DRIVER_MODEL]) + resources.defer(lambda k=key: client.delete_key(k)) + response = client.chat(key, DRIVER_MODEL, f"reply with one word {alias}") + assert response.model, f"driver call for {alias} returned no model: {response}" + + wanted = frozenset(aliases) + deadline = time.monotonic() + client.gateway.poll_timeout + seen: frozenset[str] = frozenset() + while time.monotonic() < deadline: + seen = _aliases_in_metric(client.scrape_metrics(), REQUESTS_METRIC, ALIAS_LABEL) + if wanted <= seen: + break + time.sleep(client.gateway.poll_interval) + + missing = wanted - seen + assert not missing, ( + f"{REQUESTS_METRIC} is missing a per-key series for aliases {sorted(missing)}; " + f"each distinct {ALIAS_LABEL} must grow its own series" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 972f816905e..075880eb126 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -6,6 +6,8 @@ response validates without mirroring every proxy field. No untyped dicts. from __future__ import annotations +from typing import Literal + from pydantic import BaseModel, ConfigDict, RootModel # ---------- keys ---------- @@ -30,11 +32,13 @@ class KeyGenerateBody(BaseModel): user_id: str | None = None team_id: str | None = None budget_id: str | None = None + key_alias: str | None = None model_max_budget: dict[str, ModelBudgetEntry] | None = None budget_fallbacks: dict[str, list[str]] | None = None budget_limits: list[BudgetWindow] | None = None tpm_limit: int | None = None rpm_limit: int | None = None + allowed_routes: list[str] | None = None class KeyGenerateResponse(BaseModel): @@ -87,6 +91,16 @@ class ChatMessage(BaseModel): content: str +class ThinkingParam(BaseModel): + """Extended-thinking control shared by Anthropic and DeepSeek reasoner models. + DeepSeek accepts only ``type`` (enabled/disabled) and ignores budget_tokens; + Anthropic also honors budget_tokens. Sending ``type="disabled"`` is the + product-facing way a caller turns reasoning off (LIT-3686 / GH #27453).""" + + type: Literal["enabled", "disabled"] + budget_tokens: int | None = None + + class ChatBody(BaseModel): model: str messages: list[ChatMessage] @@ -94,6 +108,9 @@ class ChatBody(BaseModel): max_tokens: int | None = None user: str | None = None metadata: ChatMetadata | None = None + reasoning_effort: str | None = None + thinking: ThinkingParam | None = None + service_tier: str | None = None class AnthropicMessagesBody(BaseModel): @@ -104,16 +121,24 @@ class AnthropicMessagesBody(BaseModel): class OutMessage(BaseModel): content: str | None = None + reasoning_content: str | None = None class ChatChoice(BaseModel): message: OutMessage | None = None +class PromptTokensDetails(BaseModel): + cached_tokens: int | None = None + + class Usage(BaseModel): prompt_tokens: int | None = None completion_tokens: int | None = None total_tokens: int | None = None + cache_read_input_tokens: int | None = None + cache_creation_input_tokens: int | None = None + prompt_tokens_details: PromptTokensDetails | None = None class ChatResponse(BaseModel): @@ -121,6 +146,7 @@ class ChatResponse(BaseModel): model: str | None = None choices: list[ChatChoice] = [] usage: Usage | None = None + service_tier: str | None = None class EmbedBody(BaseModel): @@ -165,6 +191,7 @@ class OcrResponse(BaseModel): class SpendLogRow(BaseModel): request_id: str | None = None + api_key: str | None = None model: str | None = None spend: float | None = None status: str | None = None @@ -284,3 +311,78 @@ class ModelInfoEntry(BaseModel): class ModelInfoResponse(BaseModel): data: list[ModelInfoEntry] = [] + + +class FileEntry(BaseModel): + id: str + + +class FileListResponse(BaseModel): + """GET /files answer. `data` is required on purpose: a 200 whose body lacks + the OpenAI-format file list must fail validation, not pass vacuously.""" + + data: list[FileEntry] + + +class FineTuningJobsParams(BaseModel): + custom_llm_provider: Literal["openai", "azure"] + + +class FineTuningJobEntry(BaseModel): + id: str + + +class FineTuningJobsResponse(BaseModel): + """GET /fine_tuning/jobs answer; `data` required for the same reason as + FileListResponse.""" + + data: list[FineTuningJobEntry] + + +# ---------- model management ---------- + + +class LiteLLMParamsBody(BaseModel): + """POST /model/new litellm_params: `model` is the only required field; `api_key` + et al may be an `os.environ/FOO` reference the proxy resolves at call time. + `input_cost_per_token`/`output_cost_per_token` register a per-deployment custom + pricing override; left None (and dropped from the body) the deployment keeps the + backend's canonical rate.""" + + model: str + api_key: str | None = None + api_base: str | None = None + api_version: str | None = None + aws_region_name: str | None = None + vertex_project: str | None = None + vertex_location: str | None = None + vertex_credentials: str | None = None + bucket_name: str | None = None + s3_bucket_name: str | None = None + s3_region_name: str | None = None + s3_access_key_id: str | None = None + s3_secret_access_key: str | None = None + aws_batch_role_arn: str | None = None + input_cost_per_token: float | None = None + output_cost_per_token: float | None = None + + +class ModelInfoBody(BaseModel): + id: str + mode: Literal["batch", "realtime", "image_generation"] | None = None + + +class ModelNewBody(BaseModel): + model_config = ConfigDict(protected_namespaces=()) + model_name: str + litellm_params: LiteLLMParamsBody + model_info: ModelInfoBody + + +class ModelNewResponse(BaseModel): + model_config = ConfigDict(protected_namespaces=()) + model_id: str + + +class ModelDeleteBody(BaseModel): + id: str diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py new file mode 100644 index 00000000000..6ea9098c228 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py @@ -0,0 +1,245 @@ +import json +import os +import sys +from datetime import datetime + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + INCOMPLETE_STREAM_ERROR_MESSAGE, + BaseAnthropicMessagesStreamingIterator, + _incomplete_stream_error_sse_event, + _is_message_stop_chunk, +) + + +class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator): + def __init__(self, litellm_logging_obj: LiteLLMLoggingObj, request_body: dict): + super().__init__(litellm_logging_obj=litellm_logging_obj, request_body=request_body) + self.logged_chunks: list = [] + + async def _handle_streaming_logging(self, collected_chunks): + self.logged_chunks = list(collected_chunks) + + +def _make_logging_obj(test_name: str) -> LiteLLMLoggingObj: + return LiteLLMLoggingObj( + model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="chat", + start_time=datetime.now(), + litellm_call_id=test_name, + function_id=test_name, + ) + + +def _make_iterator(test_name: str) -> BaseAnthropicMessagesStreamingIterator: + return BaseAnthropicMessagesStreamingIterator( + litellm_logging_obj=_make_logging_obj(test_name), + request_body={}, + ) + + +async def _collect(iterator, stream): + return [chunk async for chunk in iterator.async_sse_wrapper(stream)] + + +TRUNCATED_TOOL_USE_EVENTS = ( + {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 10, "output_tokens": 1}}}, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "tool_use", "id": "tooluse_1", "name": "write", "input": {}}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "input_json_delta", "partial_json": '{"path": "/builder/docs/QUAL'}, + }, +) + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_emits_error_event_when_stream_ends_without_message_stop(): + """ + Regression test for LIT-3724: a Bedrock stream that goes silent + mid tool_use must not be passed through as a successful, complete + SSE stream. An `error` SSE event must be appended so strict clients + (Anthropic SDK, Claude Code) surface the truncation instead of + crashing on unterminated tool-call JSON. + """ + + async def _truncated_stream(): + for event in TRUNCATED_TOOL_USE_EVENTS: + yield event + + iterator = _make_iterator("test_truncated_stream_emits_error") + chunks = await _collect(iterator, _truncated_stream()) + + assert len(chunks) == len(TRUNCATED_TOOL_USE_EVENTS) + 1 + error_chunk = chunks[-1].decode() + assert error_chunk.startswith("event: error\n") + assert error_chunk.endswith("\n\n") + + error_payload = json.loads(error_chunk.split("data: ", 1)[1]) + assert error_payload["type"] == "error" + assert error_payload["error"]["type"] == "api_error" + assert error_payload["error"]["message"] == INCOMPLETE_STREAM_ERROR_MESSAGE + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_no_error_event_on_complete_stream(): + async def _complete_stream(): + for event in TRUNCATED_TOOL_USE_EVENTS: + yield event + yield {"type": "content_block_stop", "index": 0} + yield {"type": "message_delta", "delta": {"stop_reason": "tool_use"}, "usage": {"output_tokens": 5}} + yield {"type": "message_stop"} + + iterator = _make_iterator("test_complete_stream_no_error") + chunks = await _collect(iterator, _complete_stream()) + + assert len(chunks) == len(TRUNCATED_TOOL_USE_EVENTS) + 3 + decoded = [chunk.decode() for chunk in chunks] + assert decoded[-1].startswith("event: message_stop\n") + assert not any(chunk.startswith("event: error\n") for chunk in decoded) + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_emits_error_event_on_empty_stream(): + async def _empty_stream(): + return + yield + + iterator = _make_iterator("test_empty_stream_emits_error") + chunks = await _collect(iterator, _empty_stream()) + + assert len(chunks) == 1 + assert chunks[0].decode().startswith("event: error\n") + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_treats_message_stop_bytes_as_complete(): + async def _byte_stream(): + yield b'event: message_start\ndata: {"type": "message_start"}\n\n' + yield b'event: message_stop\ndata: {"type": "message_stop"}\n\n' + + iterator = _make_iterator("test_byte_stream_message_stop") + chunks = await _collect(iterator, _byte_stream()) + + assert len(chunks) == 2 + assert not any(chunk.startswith(b"event: error\n") for chunk in chunks) + + +def test_is_message_stop_chunk(): + assert _is_message_stop_chunk({"type": "message_stop"}) is True + assert _is_message_stop_chunk({"type": "message_delta"}) is False + assert _is_message_stop_chunk(b'event: message_stop\ndata: {}\n\n') is True + assert _is_message_stop_chunk(b"raw-bytes") is False + assert _is_message_stop_chunk("message_stop") is False + + +def test_is_message_stop_chunk_ignores_substring_in_payload(): + """ + Regression: a `content_block_delta` frame whose payload happens to contain + the literal string `message_stop` (e.g. inside a tool's partial_json) must + not be treated as a terminal stop event. + """ + delta_frame_with_substring = ( + b'event: content_block_delta\n' + b'data: {"type": "content_block_delta", "delta": ' + b'{"type": "input_json_delta", "partial_json": "\\"message_stop\\""}}\n\n' + ) + assert _is_message_stop_chunk(delta_frame_with_substring) is False + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_emits_error_when_bytes_stream_only_mentions_message_stop_in_payload(): + """ + Regression for the bytes-branch substring false positive: a stream whose + payload text contains `message_stop` (but never emits the actual + `event: message_stop` frame) must still be flagged as incomplete. + """ + async def _byte_stream(): + yield b'event: message_start\ndata: {"type": "message_start"}\n\n' + yield ( + b'event: content_block_delta\n' + b'data: {"type": "content_block_delta", "delta": ' + b'{"type": "input_json_delta", "partial_json": "\\"message_stop\\""}}\n\n' + ) + + iterator = _make_iterator("test_bytes_substring_does_not_mark_complete") + chunks = await _collect(iterator, _byte_stream()) + + assert len(chunks) == 3 + assert chunks[-1].decode().startswith("event: error\n") + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_does_not_double_error_on_provider_error_dict(): + """ + Regression: when the provider itself terminates the stream with an + `error` event (without a `message_stop`), the wrapper must forward that + error and not append a second synthetic incomplete-stream error. + """ + provider_error = {"type": "error", "error": {"type": "overloaded_error", "message": "boom"}} + + async def _error_terminated_stream(): + yield {"type": "message_start", "message": {"id": "msg_1"}} + yield provider_error + + iterator = _make_iterator("test_provider_error_terminal_dict") + chunks = await _collect(iterator, _error_terminated_stream()) + + assert len(chunks) == 2 + error_frames = [c for c in chunks if c.startswith(b"event: error\n")] + assert len(error_frames) == 1 + payload = json.loads(error_frames[0].decode().split("data: ", 1)[1]) + assert payload == provider_error + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_does_not_double_error_on_provider_error_bytes(): + async def _byte_stream(): + yield b'event: message_start\ndata: {"type": "message_start"}\n\n' + yield b'event: error\ndata: {"type": "error", "error": {"type": "overloaded_error"}}\n\n' + + iterator = _make_iterator("test_provider_error_terminal_bytes") + chunks = await _collect(iterator, _byte_stream()) + + assert len(chunks) == 2 + error_frames = [c for c in chunks if c.startswith(b"event: error\n")] + assert len(error_frames) == 1 + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_excludes_synthetic_error_event_from_logged_chunks(): + async def _truncated_stream(): + for event in TRUNCATED_TOOL_USE_EVENTS: + yield event + + iterator = _RecordingLoggingIterator( + litellm_logging_obj=_make_logging_obj("test_synthetic_error_not_logged"), + request_body={}, + ) + chunks = await _collect(iterator, _truncated_stream()) + + assert chunks[-1].startswith(b"event: error\n") + assert iterator.logged_chunks == chunks[:-1] + assert not any(chunk.startswith(b"event: error\n") for chunk in iterator.logged_chunks) + + +def test_incomplete_stream_error_sse_event_is_valid_anthropic_error(): + event = _incomplete_stream_error_sse_event().decode() + lines = event.split("\n") + assert lines[0] == "event: error" + payload = json.loads(lines[1].removeprefix("data: ")) + assert payload == { + "type": "error", + "error": {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE}, + } + assert event.endswith("\n\n") diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index 03d0d87a58c..d7a62aae38b 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -90,6 +90,87 @@ async def test_bedrock_sse_wrapper_encodes_dict_chunks(): assert collected[1] == b"raw-bytes" +@pytest.mark.asyncio +async def test_bedrock_sse_wrapper_appends_error_event_when_stream_truncates_mid_tool_use(): + """ + Regression test for LIT-3724: Bedrock invoke streams that go silent + mid tool_use (no content_block_stop / message_delta / message_stop) + used to be closed as a successful SSE stream, handing clients + unterminated tool-call JSON with HTTP 200. The stream must now end + with an Anthropic-protocol `error` SSE event. + """ + cfg = AmazonAnthropicClaudeMessagesConfig() + + async def _truncated_stream(): + yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 3, "output_tokens": 1}}} + yield { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "tool_use", "id": "tooluse_1", "name": "write", "input": {}}, + } + yield { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "input_json_delta", "partial_json": '{"path": "/builder/docs/QUAL'}, + } + + collected: list[bytes] = [] + async for chunk in cfg.bedrock_sse_wrapper( + _truncated_stream(), + litellm_logging_obj=LiteLLMLoggingObj( + model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "write the file"}], + stream=True, + call_type="chat", + start_time=datetime.now(), + litellm_call_id="test_bedrock_sse_wrapper_truncated_tool_use", + function_id="test_bedrock_sse_wrapper_truncated_tool_use", + ), + request_body={}, + ): + collected.append(chunk) + + assert len(collected) == 4 + error_event = collected[-1].decode() + assert error_event.startswith("event: error\n") + error_payload = json.loads(error_event.split("data: ", 1)[1]) + assert error_payload["type"] == "error" + assert error_payload["error"]["type"] == "api_error" + + +@pytest.mark.asyncio +async def test_bedrock_sse_wrapper_no_error_event_when_stream_ends_with_message_stop(): + cfg = AmazonAnthropicClaudeMessagesConfig() + + async def _complete_stream(): + yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 3, "output_tokens": 1}}} + yield {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}} + yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}} + yield {"type": "content_block_stop", "index": 0} + yield {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}} + yield {"type": "message_stop"} + + collected: list[bytes] = [] + async for chunk in cfg.bedrock_sse_wrapper( + _complete_stream(), + litellm_logging_obj=LiteLLMLoggingObj( + model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="chat", + start_time=datetime.now(), + litellm_call_id="test_bedrock_sse_wrapper_complete_stream", + function_id="test_bedrock_sse_wrapper_complete_stream", + ), + request_body={}, + ): + collected.append(chunk) + + assert len(collected) == 6 + assert collected[-1].startswith(b"event: message_stop\n") + assert not any(chunk.startswith(b"event: error\n") for chunk in collected) + + @pytest.mark.asyncio async def test_bedrock_sse_wrapper_keeps_usage_in_message_start_and_message_delta(): """Regression test: usage should be available on both message_start and message_delta SSE events.""" diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index b18af060a20..0b4187d1bcf 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -435,6 +435,227 @@ async def test_async_anthropic_messages_handler_extra_headers(): assert captured_headers["X-Auth-Token"] == "token123" +@pytest.mark.asyncio +async def test_async_anthropic_messages_handler_streaming_forwards_provider_response_headers(): + """ + Regression test for LIT-3724 (issue 2): streaming /v1/messages responses + dropped the upstream provider's HTTP response headers, so Bedrock's + x-amzn-requestid / x-amzn-trace-id never reached clients even with + `return_response_headers: true`. The returned stream object must carry + them in `_hidden_params["additional_headers"]` (llm_provider-* prefixed), + which the proxy merges into the client-facing response headers. + """ + from collections.abc import AsyncIterator as ABCAsyncIterator + + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + handler = BaseLLMHTTPHandler() + + sse_body = ( + b'event: message_start\ndata: {"type": "message_start"}\n\n' + b'event: message_stop\ndata: {"type": "message_stop"}\n\n' + ) + upstream_response = httpx.Response( + 200, + headers={ + "x-amzn-requestid": "amzn-req-123", + "x-amzn-trace-id": "Root=1-abc-def", + }, + content=sse_body, + request=httpx.Request("POST", "https://api.anthropic.com/v1/messages"), + ) + mock_client = AsyncMock(spec=AsyncHTTPHandler) + mock_client.post = AsyncMock(return_value=upstream_response) + + mock_logging_obj = Mock() + mock_logging_obj.model_call_details = {} + + result = await handler.async_anthropic_messages_handler( + model="claude-sonnet-4-20250514", + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_provider_config=AnthropicMessagesConfig(), + anthropic_messages_optional_request_params={"max_tokens": 32}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=mock_logging_obj, + client=mock_client, + api_key="sk-test", + stream=True, + kwargs={}, + ) + + assert isinstance(result, ABCAsyncIterator) + + additional_headers = result._hidden_params["additional_headers"] + assert additional_headers["llm_provider-x-amzn-requestid"] == "amzn-req-123" + assert additional_headers["llm_provider-x-amzn-trace-id"] == "Root=1-abc-def" + + collected = b"".join([chunk async for chunk in result]) + assert b"message_start" in collected + assert b"message_stop" in collected + + +@pytest.mark.asyncio +async def test_async_anthropic_messages_handler_agentic_streaming_forwards_provider_response_headers(): + """ + Companion to the test above for the agentic branch: when a callback + overrides async_should_run_agentic_loop, the handler wraps + AgenticAnthropicStreamingIterator in AnthropicMessagesStreamingResponse. + That wrapping must still expose the provider headers and delegate + iteration through the two-phase agentic iterator unchanged. + """ + from collections.abc import AsyncIterator as ABCAsyncIterator + + from litellm.integrations.custom_logger import CustomLogger + from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + class NoOpAgenticCallback(CustomLogger): + async def async_should_run_agentic_loop( + self, + response, + model, + messages, + tools, + stream, + custom_llm_provider, + kwargs, + ): + return False, {} + + handler = BaseLLMHTTPHandler() + + sse_body = ( + b'event: message_start\ndata: {"type": "message_start"}\n\n' + b'event: message_stop\ndata: {"type": "message_stop"}\n\n' + ) + upstream_response = httpx.Response( + 200, + headers={"x-amzn-requestid": "amzn-req-456"}, + content=sse_body, + request=httpx.Request("POST", "https://api.anthropic.com/v1/messages"), + ) + mock_client = AsyncMock(spec=AsyncHTTPHandler) + mock_client.post = AsyncMock(return_value=upstream_response) + + mock_logging_obj = Mock() + mock_logging_obj.model_call_details = {} + mock_logging_obj.dynamic_success_callbacks = [NoOpAgenticCallback()] + + result = await handler.async_anthropic_messages_handler( + model="claude-sonnet-4-20250514", + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_provider_config=AnthropicMessagesConfig(), + anthropic_messages_optional_request_params={"max_tokens": 32}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=mock_logging_obj, + client=mock_client, + api_key="sk-test", + stream=True, + kwargs={}, + ) + + assert isinstance(result, ABCAsyncIterator) + assert isinstance(result.completion_stream, AgenticAnthropicStreamingIterator) + assert result._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "amzn-req-456" + + collected = b"".join([chunk async for chunk in result]) + assert b"message_start" in collected + assert b"message_stop" in collected + + +@pytest.mark.asyncio +async def test_anthropic_messages_streaming_response_aclose_closes_upstream_stream(): + """ + Regression test: the proxy's streaming cleanup calls aclose on the + handler's return value (see _finalize_streaming_generator_cleanup's + hasattr(response, "aclose") check). The wrapper must forward aclose to + the upstream stream so provider connections are released on client + disconnect instead of lingering until garbage collection. + """ + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + AnthropicMessagesStreamingResponse, + ) + + class UpstreamTracker: + def __init__(self): + self.closed = False + + tracker = UpstreamTracker() + + async def upstream(): + try: + yield b'data: {"type": "message_start"}\n\n' + yield b'data: {"type": "message_stop"}\n\n' + finally: + tracker.closed = True + + stream = AnthropicMessagesStreamingResponse( + completion_stream=upstream(), + hidden_params={"additional_headers": {}}, + ) + + first_chunk = await stream.__anext__() + assert b"message_start" in first_chunk + assert tracker.closed is False + + await stream.aclose() + assert tracker.closed is True + + +@pytest.mark.asyncio +async def test_anthropic_messages_streaming_response_aclose_closes_agentic_upstream_stream(): + from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + AnthropicMessagesStreamingResponse, + ) + + class UpstreamTracker: + def __init__(self): + self.closed = False + + tracker = UpstreamTracker() + + async def upstream(): + try: + yield b'data: {"type": "message_start"}\n\n' + yield b'data: {"type": "message_stop"}\n\n' + finally: + tracker.closed = True + + agentic_iterator = AgenticAnthropicStreamingIterator( + completion_stream=upstream(), + http_handler=Mock(), + model="claude-sonnet-4-20250514", + messages=[{"role": "user", "content": "Hello"}], + anthropic_messages_provider_config=Mock(), + anthropic_messages_optional_request_params={}, + logging_obj=Mock(), + custom_llm_provider="anthropic", + kwargs={}, + ) + stream = AnthropicMessagesStreamingResponse( + completion_stream=agentic_iterator, + hidden_params={"additional_headers": {}}, + ) + + first_chunk = await stream.__anext__() + assert b"message_start" in first_chunk + assert tracker.closed is False + + await stream.aclose() + assert tracker.closed is True + + @pytest.mark.asyncio async def test_async_anthropic_messages_handler_passes_litellm_metadata(): """Ensure litellm_metadata from kwargs is forwarded via update_from_kwargs.