mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_dashscope_rerank_endpoint
This commit is contained in:
commit
c614d68d11
8 changed files with 600 additions and 37 deletions
|
|
@ -26,7 +26,7 @@ If `db.useStackgresOperator` is used (not yet implemented):
|
|||
| `replicaCount` | The number of LiteLLM Proxy pods to be deployed | `1` |
|
||||
| `masterkeySecretName` | The name of the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use the generated secret name. | N/A |
|
||||
| `masterkeySecretKey` | The key within the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use `masterkey` as the key. | N/A |
|
||||
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated. | N/A |
|
||||
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated on first install and reused on upgrades. | N/A |
|
||||
| `environmentSecrets` | An optional array of Secret object names. The keys and values in these secrets will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
|
||||
| `environmentConfigMaps` | An optional array of ConfigMap object names. The keys and values in these configmaps will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
|
||||
| `image.repository` | LiteLLM Proxy image repository | `ghcr.io/berriai/litellm` |
|
||||
|
|
@ -212,6 +212,8 @@ service, the **Proxy Endpoint** should be set to `http://<RELEASE>-litellm:4000`
|
|||
The **Proxy Key** is the value specified for `masterkey` or, if a `masterkey`
|
||||
was not provided to the helm command line, the `masterkey` is a randomly
|
||||
generated string in the `sk-...` format stored in the `<RELEASE>-litellm-masterkey` Kubernetes Secret.
|
||||
The key is generated once on the first install; later `helm upgrade` runs reuse the
|
||||
value already in that Secret, so upgrading never rotates the master key.
|
||||
|
||||
```bash
|
||||
kubectl -n litellm get secret <RELEASE>-litellm-masterkey -o jsonpath="{.data.masterkey}"
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
{{- if not .Values.masterkeySecretName }}
|
||||
{{ $masterkey := (.Values.masterkey | default (printf "sk-%s" (randAlphaNum 18))) }}
|
||||
{{- $secretName := printf "%s-masterkey" (include "litellm.fullname" .) }}
|
||||
{{- $existing := lookup "v1" "Secret" .Release.Namespace $secretName }}
|
||||
{{- $masterkey := .Values.masterkey | default (dig "data" "masterkey" "" $existing | b64dec) | default (printf "sk-%s" (randAlphaNum 18)) }}
|
||||
apiVersion: v1
|
||||
kind: Secret
|
||||
metadata:
|
||||
name: {{ include "litellm.fullname" . }}-masterkey
|
||||
name: {{ $secretName }}
|
||||
data:
|
||||
masterkey: {{ $masterkey | b64enc }}
|
||||
type: Opaque
|
||||
|
|
|
|||
|
|
@ -15,6 +15,53 @@ tests:
|
|||
# Note: The masterkey is generated as "sk-<18-random-chars>" in plain text,
|
||||
# but stored as base64 encoded in Kubernetes secret (requirement).
|
||||
# "sk-" base64 encodes to "c2st", so we check for "^c2st" pattern.
|
||||
- it: should reuse the master key already stored in the cluster instead of generating a new one on upgrade
|
||||
template: secret-masterkey.yaml
|
||||
set:
|
||||
masterkeySecretName: ""
|
||||
kubernetesProvider:
|
||||
scheme:
|
||||
"v1/Secret":
|
||||
gvr:
|
||||
version: "v1"
|
||||
resource: "secrets"
|
||||
namespaced: true
|
||||
objects:
|
||||
- kind: Secret
|
||||
apiVersion: v1
|
||||
metadata:
|
||||
name: RELEASE-NAME-litellm-masterkey
|
||||
namespace: NAMESPACE
|
||||
data:
|
||||
masterkey: c2stZXhpc3Rpbmcta2V5
|
||||
asserts:
|
||||
- equal:
|
||||
path: data.masterkey
|
||||
value: c2stZXhpc3Rpbmcta2V5
|
||||
- it: should let an explicit masterkey value override the one already stored in the cluster
|
||||
template: secret-masterkey.yaml
|
||||
set:
|
||||
masterkeySecretName: ""
|
||||
masterkey: sk-explicit
|
||||
kubernetesProvider:
|
||||
scheme:
|
||||
"v1/Secret":
|
||||
gvr:
|
||||
version: "v1"
|
||||
resource: "secrets"
|
||||
namespaced: true
|
||||
objects:
|
||||
- kind: Secret
|
||||
apiVersion: v1
|
||||
metadata:
|
||||
name: RELEASE-NAME-litellm-masterkey
|
||||
namespace: NAMESPACE
|
||||
data:
|
||||
masterkey: c2stZXhpc3Rpbmcta2V5
|
||||
asserts:
|
||||
- equal:
|
||||
path: data.masterkey
|
||||
value: c2stZXhwbGljaXQ=
|
||||
- it: should not create a secret if masterkeySecretName is set
|
||||
template: secret-masterkey.yaml
|
||||
set:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import anyio
|
||||
|
|
@ -20,8 +21,11 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
|
|||
MCPUpstreamAuthError,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
||||
ServerListOk,
|
||||
ServerOutcome,
|
||||
classify_list_exception,
|
||||
list_fault_http_status,
|
||||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
acting_user_auth,
|
||||
|
|
@ -99,6 +103,7 @@ if MCP_AVAILABLE:
|
|||
ListMCPToolsRestAPIResponseObject,
|
||||
MCPInfo,
|
||||
MCPServer,
|
||||
_aggregate_server_key, # pyright: ignore[reportPrivateUsage] # same per-server key as the tools/list _meta outcomes
|
||||
_apply_toolset_scope,
|
||||
_fire_mcp_tool_call_logging,
|
||||
execute_mcp_tool,
|
||||
|
|
@ -803,9 +808,6 @@ if MCP_AVAILABLE:
|
|||
list(allowed_server_ids_set), _rest_client_ip
|
||||
)
|
||||
|
||||
list_tools_result: Final = []
|
||||
error_message = None
|
||||
|
||||
# If server_id is specified, only query that specific server
|
||||
if server_id:
|
||||
return await _list_tools_for_single_server(
|
||||
|
|
@ -849,22 +851,19 @@ if MCP_AVAILABLE:
|
|||
else {}
|
||||
)
|
||||
|
||||
# Query all servers the user has access to
|
||||
errors: Final = []
|
||||
for allowed_server_id in allowed_server_ids:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
|
||||
if server is None:
|
||||
continue
|
||||
|
||||
server_auth_header = _get_server_auth_header(server, mcp_server_auth_headers, mcp_auth_header)
|
||||
user_oauth_extra_headers = await _get_user_oauth_extra_headers(
|
||||
async def list_server(
|
||||
server: MCPServer,
|
||||
) -> tuple[Sequence[ListMCPToolsRestAPIResponseObject], ServerOutcome]:
|
||||
server_auth_header: Final = _get_server_auth_header(
|
||||
server, mcp_server_auth_headers, mcp_auth_header
|
||||
)
|
||||
user_oauth_extra_headers: Final = await _get_user_oauth_extra_headers(
|
||||
server,
|
||||
user_api_key_dict,
|
||||
prefetched_creds=prefetched_oauth_creds,
|
||||
)
|
||||
|
||||
try:
|
||||
tools_result = await _get_tools_for_single_server(
|
||||
tools_result: Final = await _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers_from_request,
|
||||
|
|
@ -872,24 +871,36 @@ if MCP_AVAILABLE:
|
|||
extra_headers=user_oauth_extra_headers,
|
||||
apply_tool_filters=apply_tool_filters,
|
||||
)
|
||||
list_tools_result.extend(tools_result)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error getting tools from %s: %s", server.name, e)
|
||||
errors.append(
|
||||
f"{get_server_prefix(server)}: {classify_list_exception(e).tag}"
|
||||
if isinstance(e, (MCPServerListError, MCPUpstreamAuthError))
|
||||
else f"{get_server_prefix(server)}: {e}"
|
||||
)
|
||||
continue
|
||||
return (), classify_list_exception(e)
|
||||
return tools_result, ServerListOk(tool_count=len(tools_result))
|
||||
|
||||
if errors and not list_tools_result:
|
||||
error_message = "Failed to get tools from servers: " + "; ".join(errors)
|
||||
|
||||
return {
|
||||
"tools": list_tools_result,
|
||||
"error": "partial_failure" if error_message else None,
|
||||
"message": (error_message if error_message else "Successfully retrieved tools"),
|
||||
}
|
||||
# Query all servers the user has access to
|
||||
queried_servers: Final = tuple(
|
||||
server
|
||||
for server in map(global_mcp_server_manager.get_mcp_server_by_id, allowed_server_ids)
|
||||
if server is not None
|
||||
)
|
||||
listings: Final = tuple([await list_server(server) for server in queried_servers])
|
||||
list_tools_result: Final = [tool for tools, _ in listings for tool in tools]
|
||||
server_outcomes: Final = MappingProxyType(
|
||||
{_aggregate_server_key(server): outcome for server, (_, outcome) in zip(queried_servers, listings)}
|
||||
)
|
||||
errors: Final = tuple(
|
||||
f"{key}: {outcome.tag}" for key, outcome in server_outcomes.items() if outcome.tag != "ok"
|
||||
)
|
||||
error_message: Final = (
|
||||
"Failed to get tools from servers: " + "; ".join(errors)
|
||||
if errors and not list_tools_result
|
||||
else None
|
||||
)
|
||||
return {
|
||||
"tools": list_tools_result,
|
||||
"error": "partial_failure" if error_message else None,
|
||||
"message": (error_message if error_message else "Successfully retrieved tools"),
|
||||
"server_outcomes": {key: outcome_wire_value(outcome) for key, outcome in server_outcomes.items()},
|
||||
}
|
||||
|
||||
except MCPUpstreamAuthError as e:
|
||||
# Surface upstream pass-through 401/403 challenges to the client so
|
||||
|
|
|
|||
174
litellm/proxy/anthropic_endpoints/streaming_model_restamp.py
Normal file
174
litellm/proxy/anthropic_endpoints/streaming_model_restamp.py
Normal file
|
|
@ -0,0 +1,174 @@
|
|||
"""
|
||||
Restamp the public ``model`` on the Anthropic Messages ``message_start`` event, the only
|
||||
stream event carrying a model, so streamed responses report the requested model like
|
||||
non-streaming ones do.
|
||||
|
||||
Chunks reach the serializer either as already-encoded SSE frames (``bytes``/``str``, the
|
||||
provider passthrough path) or as event dicts (fake-stream and agentic paths).
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
_MESSAGE_START_EVENT: Final = "message_start"
|
||||
_MESSAGE_START_MARKER: Final = b"message_start"
|
||||
_SSE_DATA_FIELD: Final = "data:"
|
||||
_SSE_FRAME_END_PATTERN: Final = re.compile(rb"\r\n\r\n|\r\r|\n\n")
|
||||
_MAX_HELD_BYTES: Final = 65536
|
||||
_PING_MARKERS: Final = (b"event: ping", b'"type": "ping"', b'"type":"ping"')
|
||||
|
||||
_EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _restamped_event(event: Mapping[str, object], requested_model: str) -> Mapping[str, object] | None:
|
||||
message: Final = event.get("message")
|
||||
if event.get("type") != _MESSAGE_START_EVENT or not isinstance(message, dict):
|
||||
return None
|
||||
if message.get("model") == requested_model:
|
||||
return None
|
||||
return {**event, "message": {**message, "model": requested_model}} # mutable-ok: SSE payload, re-serialized as is
|
||||
|
||||
|
||||
def _restamped_data_line(line: str, requested_model: str) -> str | None:
|
||||
stripped: Final = line.strip()
|
||||
if not stripped.startswith(_SSE_DATA_FIELD):
|
||||
return None
|
||||
payload: Final = stripped[len(_SSE_DATA_FIELD) :].strip()
|
||||
if not payload or payload == "[DONE]":
|
||||
return None
|
||||
try:
|
||||
event: Final = _EVENT_ADAPTER.validate_json(payload)
|
||||
except ValidationError:
|
||||
return None
|
||||
restamped: Final = _restamped_event(event, requested_model)
|
||||
if restamped is None:
|
||||
return None
|
||||
terminator: Final = line[len(line.rstrip("\r\n")) :]
|
||||
return f"data: {json.dumps(restamped, separators=(',', ':'))}{terminator}"
|
||||
|
||||
|
||||
def _restamped_frame(frame: str, requested_model: str) -> str | None:
|
||||
lines: Final = frame.splitlines(keepends=True)
|
||||
restamped: Final = tuple(_restamped_data_line(line, requested_model) for line in lines)
|
||||
if all(line is None for line in restamped):
|
||||
return None
|
||||
return "".join(new if new is not None else old for new, old in zip(restamped, lines))
|
||||
|
||||
|
||||
def restamp_anthropic_stream_chunk_model(chunk: object, requested_model: str) -> object:
|
||||
"""
|
||||
Return ``chunk`` with the ``message_start`` model replaced by ``requested_model``.
|
||||
|
||||
Chunks that carry no model are returned unchanged.
|
||||
"""
|
||||
if isinstance(chunk, dict):
|
||||
try:
|
||||
event: Final = _EVENT_ADAPTER.validate_python(chunk)
|
||||
except ValidationError:
|
||||
return chunk
|
||||
return _restamped_event(event, requested_model) or chunk
|
||||
|
||||
if isinstance(chunk, (bytes, bytearray)):
|
||||
if _MESSAGE_START_EVENT.encode() not in chunk:
|
||||
return chunk
|
||||
restamped_bytes: Final = _restamped_frame(chunk.decode("utf-8", errors="ignore"), requested_model)
|
||||
return chunk if restamped_bytes is None else restamped_bytes.encode("utf-8")
|
||||
|
||||
if isinstance(chunk, str):
|
||||
if _MESSAGE_START_EVENT not in chunk:
|
||||
return chunk
|
||||
restamped_text: Final = _restamped_frame(chunk, requested_model)
|
||||
return chunk if restamped_text is None else restamped_text
|
||||
|
||||
return chunk
|
||||
|
||||
|
||||
def _is_ping_frame(frame: bytes) -> bool:
|
||||
return any(marker in frame for marker in _PING_MARKERS)
|
||||
|
||||
|
||||
class AnthropicStreamModelRestamper:
|
||||
"""
|
||||
Per-stream restamper for the encoded passthrough path, where chunks are raw
|
||||
transport reads: the ``message_start`` SSE frame can arrive split across
|
||||
chunks or coalesced with later frames. Complete frames (``\\n\\n``,
|
||||
``\\r\\n\\r\\n``, or ``\\r\\r`` terminated) are emitted as their terminator
|
||||
closes them and an incomplete tail is held until it completes, so the
|
||||
restamp never misses a torn frame; ``flush`` returns whatever is still held
|
||||
when the stream ends so no bytes are swallowed. Once ``message_start`` has
|
||||
been handled, or the first real event proves the stream carries none, every
|
||||
later chunk passes through untouched.
|
||||
"""
|
||||
|
||||
def __init__(self, requested_model: str) -> None:
|
||||
self._requested_model: Final = requested_model
|
||||
self._held = b""
|
||||
self._armed = True
|
||||
|
||||
def process(self, chunk: object) -> object:
|
||||
if not self._armed:
|
||||
return chunk
|
||||
if isinstance(chunk, (bytes, bytearray)):
|
||||
return self._process_encoded(bytes(chunk))
|
||||
if isinstance(chunk, str):
|
||||
return self._process_encoded(chunk.encode("utf-8"))
|
||||
restamped: Final = restamp_anthropic_stream_chunk_model(chunk, self._requested_model)
|
||||
if isinstance(chunk, dict) and chunk.get("type") not in (None, "ping"):
|
||||
self._armed = False
|
||||
return restamped
|
||||
|
||||
def flush(self) -> bytes:
|
||||
held: Final = self._held
|
||||
self._held = b""
|
||||
self._armed = False
|
||||
if not held:
|
||||
return b""
|
||||
restamped: Final = restamp_anthropic_stream_chunk_model(held, self._requested_model)
|
||||
return restamped if isinstance(restamped, bytes) else held
|
||||
|
||||
def _process_encoded(self, data: bytes) -> bytes:
|
||||
combined: Final = self._held + data
|
||||
boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(combined))
|
||||
if not boundaries:
|
||||
if len(combined) > _MAX_HELD_BYTES:
|
||||
self._held = b""
|
||||
self._armed = False
|
||||
return combined
|
||||
self._held = combined
|
||||
return b""
|
||||
emitted: Final = self._restamped_closed_block(combined[: boundaries[-1]])
|
||||
tail: Final = combined[boundaries[-1] :]
|
||||
if not self._armed:
|
||||
self._held = b""
|
||||
return emitted + tail
|
||||
self._held = tail
|
||||
return emitted
|
||||
|
||||
def _restamped_closed_block(self, closed: bytes) -> bytes:
|
||||
boundaries: Final = tuple(match.end() for match in _SSE_FRAME_END_PATTERN.finditer(closed))
|
||||
frames: Final = tuple(closed[start:end] for start, end in zip((0, *boundaries[:-1]), boundaries))
|
||||
decider: Final = next(
|
||||
(
|
||||
index
|
||||
for index, frame in enumerate(frames)
|
||||
if _MESSAGE_START_MARKER in frame or (b"data:" in frame and not _is_ping_frame(frame))
|
||||
),
|
||||
None,
|
||||
)
|
||||
if decider is None:
|
||||
return closed
|
||||
self._armed = False
|
||||
if _MESSAGE_START_MARKER not in frames[decider]:
|
||||
return closed
|
||||
restamped_text: Final = _restamped_frame(
|
||||
frames[decider].decode("utf-8", errors="ignore"), self._requested_model
|
||||
)
|
||||
if restamped_text is None:
|
||||
return closed
|
||||
return b"".join(
|
||||
restamped_text.encode("utf-8") if index == decider else frame for index, frame in enumerate(frames)
|
||||
)
|
||||
|
|
@ -176,6 +176,9 @@ if TYPE_CHECKING:
|
|||
ProxyConfig = _ProxyConfig
|
||||
else:
|
||||
ProxyConfig = Any
|
||||
from litellm.proxy.anthropic_endpoints.streaming_model_restamp import (
|
||||
AnthropicStreamModelRestamper,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
add_litellm_data_to_request,
|
||||
refresh_proxy_server_request_body_snapshot,
|
||||
|
|
@ -2490,6 +2493,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request_data=self.data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
request=request,
|
||||
restamp_model=(
|
||||
None if _should_return_raw_model_name(self.data) else requested_model_from_client
|
||||
),
|
||||
)
|
||||
return await create_response(
|
||||
generator=wrap_sse_stream_with_keepalive_pings(
|
||||
|
|
@ -3442,6 +3448,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
else:
|
||||
return chunk
|
||||
|
||||
@staticmethod
|
||||
def _sse_chunk_serializer(restamper: AnthropicStreamModelRestamper | None) -> StreamChunkSerializer:
|
||||
if restamper is None:
|
||||
return ProxyBaseLLMRequestProcessing.return_sse_chunk
|
||||
|
||||
def serialize(chunk: object) -> str:
|
||||
return ProxyBaseLLMRequestProcessing.return_sse_chunk(restamper.process(chunk))
|
||||
|
||||
return serialize
|
||||
|
||||
@staticmethod
|
||||
async def _finalize_streaming_generator_cleanup(
|
||||
request: Request | None,
|
||||
|
|
@ -3502,11 +3518,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
serialize_chunk: StreamChunkSerializer,
|
||||
serialize_error: StreamErrorSerializer,
|
||||
request: Request | None = None,
|
||||
flush_tail: Callable[[], bytes] | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
|
||||
cost injection, then yields chunks via serialize_chunk; on exception runs
|
||||
failure hook and yields via serialize_error. Use for SSE or NDJSON.
|
||||
|
||||
``flush_tail`` runs once after the upstream iterator completes cleanly and
|
||||
its non-empty result is yielded, so a serializer that buffers bytes across
|
||||
chunks can emit anything still held at end of stream.
|
||||
"""
|
||||
verbose_proxy_logger.debug("inside generator")
|
||||
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
|
||||
|
|
@ -3569,6 +3590,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# so it must not suppress that refund.
|
||||
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
yield serialize_chunk(chunk)
|
||||
held_tail: Final = flush_tail() if flush_tail is not None else b""
|
||||
if held_tail:
|
||||
yield serialize_chunk(held_tail)
|
||||
stream_completed = True
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
# Client disconnected mid-stream. CancelledError / GeneratorExit
|
||||
|
|
@ -3579,8 +3603,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# billing and release exactly once. This is the outermost generator
|
||||
# Starlette closes on disconnect, so the nested iterator hook (which
|
||||
# only sees GeneratorExit on GC) cannot own the refund.
|
||||
if not stream_completed:
|
||||
client_disconnected = True
|
||||
client_disconnected = not stream_completed
|
||||
if not delivered_chunk and not _withheld_provider_output(response):
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
release_budget_reservation_on_cancel,
|
||||
|
|
@ -3634,6 +3657,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request_data: dict,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
request: Request | None = None,
|
||||
restamp_model: str | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Anthropic /messages and Google /generateContent streaming data generator require SSE events.
|
||||
|
|
@ -3642,17 +3666,23 @@ class ProxyBaseLLMRequestProcessing:
|
|||
SSE serializers directly (rather than re-wrapping it in another
|
||||
``async for: yield`` trampoline), so a streamed chunk traverses one
|
||||
fewer async-generator layer / coroutine resume on the hot path.
|
||||
|
||||
``restamp_model`` publishes that name on the Anthropic ``message_start``
|
||||
event in place of the provider's model, matching what the non-streaming
|
||||
response reports.
|
||||
"""
|
||||
restamper: Final = AnthropicStreamModelRestamper(restamp_model) if restamp_model else None
|
||||
return ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
serialize_chunk=ProxyBaseLLMRequestProcessing.return_sse_chunk,
|
||||
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper),
|
||||
serialize_error=lambda proxy_exc: (
|
||||
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
|
||||
),
|
||||
request=request,
|
||||
flush_tail=None if restamper is None else restamper.flush,
|
||||
)
|
||||
|
||||
@overload
|
||||
|
|
|
|||
|
|
@ -1431,7 +1431,11 @@ class TestListToolsRestAPI:
|
|||
async def test_aggregate_list_absorbs_one_server_auth_failure(self, monkeypatch):
|
||||
"""The multi-server aggregate listing degrades a server whose upstream
|
||||
rejects auth to an empty contribution and still returns the healthy
|
||||
server's tools with a 200, rather than surfacing a 401."""
|
||||
server's tools with a 200, rather than surfacing a 401. The absorbed
|
||||
server must still show up as a classified per-server outcome so a REST
|
||||
caller can tell "needs upstream auth" apart from "has no tools"."""
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
|
|
@ -1497,6 +1501,11 @@ class TestListToolsRestAPI:
|
|||
|
||||
assert result["tools"] == ["good-tool"]
|
||||
assert result["error"] is None
|
||||
wire_body = json.loads(TypeAdapter(dict).dump_json(result))
|
||||
assert wire_body["server_outcomes"] == {
|
||||
"good": {"status": "ok", "tool_count": 1},
|
||||
"bad": {"status": "auth_required", "http_status": 401},
|
||||
}
|
||||
|
||||
async def test_name_resolution_finds_server_by_uuid(self, monkeypatch):
|
||||
"""When server_id is a name string, it should be resolved to its UUID
|
||||
|
|
|
|||
|
|
@ -0,0 +1,288 @@
|
|||
"""
|
||||
Tests for restamping the public model on Anthropic Messages streaming chunks.
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.anthropic_endpoints.streaming_model_restamp import (
|
||||
AnthropicStreamModelRestamper,
|
||||
restamp_anthropic_stream_chunk_model,
|
||||
)
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
|
||||
def _message_start_frame(model: str, line_end: str = "\n") -> bytes:
|
||||
payload = {
|
||||
"type": "message_start",
|
||||
"message": {"id": "msg_1", "type": "message", "role": "assistant", "model": model, "content": []},
|
||||
}
|
||||
return f"event: message_start{line_end}data: {json.dumps(payload)}{line_end}{line_end}".encode()
|
||||
|
||||
|
||||
def _proxy_logging_obj_streaming(frames: list[bytes]) -> MagicMock:
|
||||
async def _iterator_hook(**_kwargs):
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.async_post_call_streaming_iterator_hook = _iterator_hook
|
||||
proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["response"])
|
||||
return proxy_logging_obj
|
||||
|
||||
|
||||
def _model_from_frame(frame: bytes | str) -> str:
|
||||
text = frame.decode("utf-8") if isinstance(frame, bytes) else frame
|
||||
data_line = next(line for line in text.split("\n") if line.startswith("data:"))
|
||||
return json.loads(data_line[len("data:") :])["message"]["model"]
|
||||
|
||||
|
||||
def test_restamps_sse_bytes_frame():
|
||||
restamped = restamp_anthropic_stream_chunk_model(
|
||||
_message_start_frame("claude-haiku-4-5-20251001"), "claude-auto-1"
|
||||
)
|
||||
|
||||
assert isinstance(restamped, bytes)
|
||||
assert _model_from_frame(restamped) == "claude-auto-1"
|
||||
assert b"event: message_start" in restamped
|
||||
|
||||
|
||||
def test_restamps_event_dict():
|
||||
chunk = {"type": "message_start", "message": {"id": "msg_1", "model": "claude-sonnet-4-6"}}
|
||||
|
||||
restamped = restamp_anthropic_stream_chunk_model(chunk, "claude-auto-2")
|
||||
|
||||
assert restamped == {"type": "message_start", "message": {"id": "msg_1", "model": "claude-auto-2"}}
|
||||
assert chunk["message"]["model"] == "claude-sonnet-4-6"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"chunk",
|
||||
[
|
||||
b'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\n\n',
|
||||
{"type": "content_block_delta", "delta": {"text": "hi"}},
|
||||
{"type": "message_start", "message": "not-a-dict"},
|
||||
b"event: message_start\ndata: not-json\n\n",
|
||||
b"data: [DONE]\n\n",
|
||||
],
|
||||
)
|
||||
def test_leaves_chunks_without_a_model_untouched(chunk):
|
||||
assert restamp_anthropic_stream_chunk_model(chunk, "claude-auto-1") == chunk
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_generator_publishes_requested_model_on_message_start():
|
||||
"""The message_start event reports the requested model, not the provider's."""
|
||||
delta_frame = b'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\n\n'
|
||||
proxy_logging_obj = _proxy_logging_obj_streaming([_message_start_frame("claude-haiku-4-5-20251001"), delta_frame])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_data={"model": "claude-auto-1"},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
restamp_model="claude-auto-1",
|
||||
)
|
||||
]
|
||||
|
||||
assert _model_from_frame(chunks[0]) == "claude-auto-1"
|
||||
assert chunks[1] == delta_frame
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_generator_keeps_provider_model_when_restamping_is_off():
|
||||
proxy_logging_obj = _proxy_logging_obj_streaming([_message_start_frame("claude-haiku-4-5-20251001")])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_data={"model": "claude-auto-1"},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
]
|
||||
|
||||
assert _model_from_frame(chunks[0]) == "claude-haiku-4-5-20251001"
|
||||
|
||||
|
||||
def test_restamps_message_start_split_across_transport_chunks():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
held = restamper.process(frame[:25])
|
||||
emitted = restamper.process(frame[25:])
|
||||
|
||||
assert held == b""
|
||||
assert isinstance(emitted, bytes)
|
||||
assert _model_from_frame(emitted) == "claude-auto-1"
|
||||
|
||||
|
||||
def test_emits_coalesced_frames_with_only_message_start_rewritten():
|
||||
delta = b'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\n\n'
|
||||
combined = _message_start_frame("claude-haiku-4-5-20251001") + delta
|
||||
|
||||
emitted = restamper_output = AnthropicStreamModelRestamper("claude-auto-1").process(combined)
|
||||
|
||||
assert isinstance(restamper_output, bytes)
|
||||
assert _model_from_frame(emitted) == "claude-auto-1"
|
||||
assert emitted.endswith(delta)
|
||||
|
||||
|
||||
def test_ping_frames_keep_the_restamper_armed():
|
||||
ping = b'event: ping\ndata: {"type": "ping"}\n\n'
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
assert restamper.process(ping) == ping
|
||||
reassembled = restamper.process(frame[:10])
|
||||
reassembled += restamper.process(frame[10:])
|
||||
|
||||
assert _model_from_frame(reassembled) == "claude-auto-1"
|
||||
|
||||
|
||||
def test_first_non_ping_event_disarms_the_restamper():
|
||||
delta = b'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\n\n'
|
||||
late_message_start = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
assert restamper.process(delta) == delta
|
||||
assert restamper.process(late_message_start) == late_message_start
|
||||
|
||||
|
||||
def test_oversized_unterminated_chunk_flushes_unmodified():
|
||||
blob = b"data: " + b"x" * 70000
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
assert restamper.process(blob) == blob
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
assert restamper.process(frame) == frame
|
||||
|
||||
|
||||
def test_dict_message_start_disarms_after_restamp():
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
first = restamper.process({"type": "message_start", "message": {"id": "msg_1", "model": "claude-sonnet-4-6"}})
|
||||
second = {"type": "message_start", "message": {"id": "msg_2", "model": "claude-sonnet-4-6"}}
|
||||
|
||||
assert first == {"type": "message_start", "message": {"id": "msg_1", "model": "claude-auto-1"}}
|
||||
assert restamper.process(second) == second
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_generator_restamps_message_start_split_across_chunks():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
proxy_logging_obj = _proxy_logging_obj_streaming([frame[:30], frame[30:]])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_data={"model": "claude-auto-1"},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
restamp_model="claude-auto-1",
|
||||
)
|
||||
]
|
||||
|
||||
joined = b"".join(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in chunks)
|
||||
assert _model_from_frame(joined) == "claude-auto-1"
|
||||
|
||||
|
||||
def test_restamps_crlf_terminated_message_start_frame():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
|
||||
delta = b'event: content_block_delta\r\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\r\n\r\n'
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
emitted = restamper.process(frame)
|
||||
|
||||
assert isinstance(emitted, bytes)
|
||||
assert _model_from_frame(emitted) == "claude-auto-1"
|
||||
assert emitted.endswith(b"\r\n\r\n")
|
||||
assert restamper.process(delta) == delta
|
||||
|
||||
|
||||
def test_restamps_cr_terminated_message_start_frame():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r")
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
emitted = restamper.process(frame)
|
||||
|
||||
assert isinstance(emitted, bytes)
|
||||
assert b'"model":"claude-auto-1"' in emitted
|
||||
assert emitted.endswith(b"\r\r")
|
||||
|
||||
|
||||
def test_restamps_crlf_message_start_split_across_transport_chunks():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
held = restamper.process(frame[:25])
|
||||
emitted = restamper.process(frame[25:])
|
||||
|
||||
assert held == b""
|
||||
assert isinstance(emitted, bytes)
|
||||
assert _model_from_frame(emitted) == "claude-auto-1"
|
||||
|
||||
|
||||
def test_flush_returns_restamped_held_tail():
|
||||
unterminated = _message_start_frame("claude-haiku-4-5-20251001")[:-2]
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
|
||||
assert restamper.process(unterminated) == b""
|
||||
flushed = restamper.flush()
|
||||
|
||||
assert b'"model":"claude-auto-1"' in flushed
|
||||
assert restamper.flush() == b""
|
||||
|
||||
|
||||
def test_flush_disarms_the_restamper():
|
||||
restamper = AnthropicStreamModelRestamper("claude-auto-1")
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001")
|
||||
|
||||
assert restamper.flush() == b""
|
||||
assert restamper.process(frame) == frame
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_generator_flushes_held_tail_at_end_of_stream():
|
||||
unterminated = _message_start_frame("claude-haiku-4-5-20251001")[:-2]
|
||||
proxy_logging_obj = _proxy_logging_obj_streaming([unterminated])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_data={"model": "claude-auto-1"},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
restamp_model="claude-auto-1",
|
||||
)
|
||||
]
|
||||
|
||||
joined = b"".join(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in chunks)
|
||||
assert b'"model":"claude-auto-1"' in joined
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_generator_restamps_crlf_stream():
|
||||
frame = _message_start_frame("claude-haiku-4-5-20251001", line_end="\r\n")
|
||||
delta = b'event: content_block_delta\r\ndata: {"type":"content_block_delta","delta":{"text":"hi"}}\r\n\r\n'
|
||||
proxy_logging_obj = _proxy_logging_obj_streaming([frame, delta])
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
request_data={"model": "claude-auto-1"},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
restamp_model="claude-auto-1",
|
||||
)
|
||||
]
|
||||
|
||||
assert _model_from_frame(chunks[0]) == "claude-auto-1"
|
||||
assert chunks[1] == delta
|
||||
Loading…
Add table
Reference in a new issue