mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/suspicious-jennings-5b6ef7
This commit is contained in:
commit
91676c424b
30 changed files with 2278 additions and 190 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
56
tests/e2e/access_control/access_control_client.py
Normal file
56
tests/e2e/access_control/access_control_client.py
Normal 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())
|
||||
10
tests/e2e/access_control/conftest.py
Normal file
10
tests/e2e/access_control/conftest.py
Normal 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()
|
||||
83
tests/e2e/access_control/test_access_control_e2e.py
Normal file
83
tests/e2e/access_control/test_access_control_e2e.py
Normal 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]}"
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
219
tests/e2e/llm_translation/endpoints_client.py
Normal file
219
tests/e2e/llm_translation/endpoints_client.py
Normal 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())
|
||||
40
tests/e2e/llm_translation/test_audio_speech_e2e.py
Normal file
40
tests/e2e/llm_translation/test_audio_speech_e2e.py
Normal 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"
|
||||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
133
tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py
Normal file
133
tests/e2e/llm_translation/test_deepseek_reasoning_e2e.py
Normal 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}"
|
||||
)
|
||||
42
tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py
Normal file
42
tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py
Normal 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]}"
|
||||
)
|
||||
42
tests/e2e/llm_translation/test_image_generation_e2e.py
Normal file
42
tests/e2e/llm_translation/test_image_generation_e2e.py
Normal 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]}"
|
||||
)
|
||||
39
tests/e2e/llm_translation/test_messages_e2e.py
Normal file
39
tests/e2e/llm_translation/test_messages_e2e.py
Normal 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]}"
|
||||
148
tests/e2e/llm_translation/test_provider_features_e2e.py
Normal file
148
tests/e2e/llm_translation/test_provider_features_e2e.py
Normal 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})"
|
||||
)
|
||||
49
tests/e2e/llm_translation/test_rerank_e2e.py
Normal file
49
tests/e2e/llm_translation/test_rerank_e2e.py
Normal 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]}"
|
||||
)
|
||||
36
tests/e2e/llm_translation/test_responses_e2e.py
Normal file
36
tests/e2e/llm_translation/test_responses_e2e.py
Normal 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]}"
|
||||
37
tests/e2e/logging/conftest.py
Normal file
37
tests/e2e/logging/conftest.py
Normal 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")
|
||||
48
tests/e2e/logging/logging_client.py
Normal file
48
tests/e2e/logging/logging_client.py
Normal 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())
|
||||
70
tests/e2e/logging/test_prometheus_cardinality_e2e.py
Normal file
70
tests/e2e/logging/test_prometheus_cardinality_e2e.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue