Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/suspicious-jennings-5b6ef7

This commit is contained in:
Yuneng Jiang 2026-07-04 19:15:01 -07:00
commit 91676c424b
No known key found for this signature in database
30 changed files with 2278 additions and 190 deletions

View file

@ -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

View file

@ -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:

View file

@ -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)

View file

@ -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,

View file

@ -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]

View file

@ -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

View file

@ -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())

View file

@ -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()

View file

@ -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]}"

View file

@ -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)

View file

@ -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

View file

@ -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()

View file

@ -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())

View file

@ -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"

View file

@ -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}"
)

View file

@ -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"
)

View file

@ -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}"
)

View file

@ -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]}"
)

View file

@ -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]}"
)

View file

@ -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]}"

View file

@ -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})"
)

View file

@ -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]}"
)

View file

@ -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]}"

View file

@ -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")

View file

@ -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())

View file

@ -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"
)

View file

@ -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

View file

@ -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")

View file

@ -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."""

View file

@ -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.