make sure to forward to the right upstream

This commit is contained in:
9qeklajc
2026-05-03 15:15:34 +02:00
parent 985e765285
commit 37c2bea93d
4 changed files with 692 additions and 8 deletions
+174 -6
View File
@@ -71,6 +71,7 @@ from ..payment.models import (
)
from ..payment.price import sats_usd_price
from ..wallet import recieve_token, send_token
from .litellm_routing import detect_litellm_prefix
logger = get_logger(__name__)
@@ -94,7 +95,10 @@ class BaseUpstreamProvider:
platform_url: str | None = None
supports_anthropic_messages: bool = False
litellm_provider_prefix: str = "openai/"
# When None, the prefix is detected from `base_url` at dispatch time
# (see `get_litellm_provider_prefix`). Subclasses set this to lock the
# provider regardless of URL.
litellm_provider_prefix: str | None = None
base_url: str
api_key: str
@@ -116,6 +120,19 @@ class BaseUpstreamProvider:
self._models_cache = []
self._models_by_id = {}
def get_litellm_provider_prefix(self) -> str:
"""Resolve the litellm provider prefix for this provider instance.
1. If the subclass pinned `litellm_provider_prefix`, use it.
2. Otherwise infer from `base_url` (e.g. ``api.fireworks.ai`` →
``fireworks_ai/``) so custom/generic rows reach the correct
litellm backend instead of falling back to ``openai/``.
3. Default ``openai/`` for unknown OpenAI-compatible servers.
"""
if self.__class__.litellm_provider_prefix:
return self.__class__.litellm_provider_prefix
return detect_litellm_prefix(self.base_url)
@classmethod
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
@@ -1559,6 +1576,121 @@ class BaseUpstreamProvider:
events.append(obj)
return events, buffer
async def _aggregate_anthropic_events_to_message(
self, iterator: AsyncIterator[Any]
) -> dict:
"""Drain an Anthropic-Messages event iterator into a single Message
dict (the shape `litellm.anthropic.messages.acreate(stream=False)`
would have produced).
Used to transparently stream from upstream while still returning a
non-streaming response to the client. Lets us sidestep upstream
quirks (e.g. Fireworks rejects ``max_tokens > 4096`` unless
``stream=true``) without leaking any of that into client-visible
behavior.
"""
sse_buffer = b""
message: dict = {}
blocks: list[dict] = []
partial_json: dict[int, str] = {}
final_stop_reason: str | None = None
final_stop_sequence: str | None = None
final_usage: dict[str, Any] = {}
final_model: str | None = None
async for chunk in iterator:
events, sse_buffer = self._events_from_chunk(chunk, sse_buffer)
for event in events:
etype = event.get("type")
if etype == "message_start":
raw = event.get("message") or {}
if isinstance(raw, dict):
message = dict(raw)
existing = message.get("content")
blocks = list(existing) if isinstance(existing, list) else []
usage = message.get("usage")
if isinstance(usage, dict):
final_usage = dict(usage)
if isinstance(message.get("model"), str):
final_model = message["model"]
elif etype == "content_block_start":
idx = int(event.get("index") or 0)
cb = event.get("content_block") or {}
cb_dict = dict(cb) if isinstance(cb, dict) else {}
while len(blocks) <= idx:
blocks.append({})
blocks[idx] = cb_dict
elif etype == "content_block_delta":
idx = int(event.get("index") or 0)
if idx >= len(blocks):
continue
delta = event.get("delta") or {}
if not isinstance(delta, dict):
continue
dtype = delta.get("type")
block = blocks[idx]
if dtype == "text_delta":
block["text"] = (block.get("text") or "") + (
delta.get("text") or ""
)
elif dtype == "input_json_delta":
partial_json[idx] = partial_json.get(idx, "") + (
delta.get("partial_json") or ""
)
elif dtype == "thinking_delta":
block["thinking"] = (block.get("thinking") or "") + (
delta.get("thinking") or ""
)
elif dtype == "signature_delta":
block["signature"] = (block.get("signature") or "") + (
delta.get("signature") or ""
)
elif etype == "content_block_stop":
idx = int(event.get("index") or 0)
raw_json = partial_json.pop(idx, None)
if raw_json is not None and idx < len(blocks):
try:
blocks[idx]["input"] = (
json.loads(raw_json) if raw_json else {}
)
except json.JSONDecodeError:
blocks[idx]["input"] = raw_json
elif etype == "message_delta":
delta = event.get("delta") or {}
if isinstance(delta, dict):
if "stop_reason" in delta:
final_stop_reason = delta.get("stop_reason")
if "stop_sequence" in delta:
final_stop_sequence = delta.get("stop_sequence")
usage = event.get("usage")
if isinstance(usage, dict):
final_usage.update(usage)
# message_stop: nothing to merge
if not message:
# Upstream returned no message_start; expose what we can so the
# client at least sees the assembled content.
message = {
"id": "",
"type": "message",
"role": "assistant",
"content": [],
}
message["content"] = blocks
if final_model and not message.get("model"):
message["model"] = final_model
if final_stop_reason is not None:
message["stop_reason"] = final_stop_reason
if final_stop_sequence is not None:
message["stop_sequence"] = final_stop_sequence
if final_usage:
existing_usage = message.get("usage")
merged = dict(existing_usage) if isinstance(existing_usage, dict) else {}
merged.update(final_usage)
message["usage"] = merged
return message
def _events_from_chunk(
self, chunk: object, sse_buffer: bytes
) -> tuple[list[dict], bytes]:
@@ -1604,7 +1736,14 @@ class BaseUpstreamProvider:
) from exc
body.pop("model", None)
stream = bool(body.pop("stream", False))
# `stream` here is what the **client** asked for. Upstream is
# always streamed (see `upstream_stream` below); when the client
# asked for a non-streaming response we drain and aggregate the
# events into a single Anthropic Message dict before returning.
# This sidesteps provider-specific non-streaming caps (e.g.
# Fireworks rejects `max_tokens > 4096` unless `stream=true`).
client_stream = bool(body.pop("stream", False))
upstream_stream = True
# Anthropic-Messages-only fields that don't translate to OpenAI
# Chat Completions. litellm.drop_params only filters *known*
@@ -1638,13 +1777,14 @@ class BaseUpstreamProvider:
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
)
upstream_model = self.transform_model_name(model_obj.id)
litellm_model = f"{self.litellm_provider_prefix}{upstream_model}"
prefix = self.get_litellm_provider_prefix()
litellm_model = f"{prefix}{upstream_model}"
kwargs: dict = {
"model": litellm_model,
"api_base": self.base_url,
"api_key": self.api_key,
"stream": stream,
"stream": upstream_stream,
**body,
}
@@ -1652,7 +1792,9 @@ class BaseUpstreamProvider:
"Dispatching /v1/messages via litellm",
extra={
"model": litellm_model,
"stream": stream,
"resolved_provider": prefix.rstrip("/"),
"client_stream": client_stream,
"upstream_stream": upstream_stream,
**(log_extra or {}),
},
)
@@ -1687,7 +1829,33 @@ class BaseUpstreamProvider:
status_code=exc_status if isinstance(exc_status, int) else 502,
) from exc
return stream, result, requested_model
if not client_stream and hasattr(result, "__aiter__"):
# Client asked for a non-streaming response but we always
# stream from upstream — drain the events into a single
# Anthropic Message dict so the rest of the pipeline can
# treat it as if the upstream had returned non-streaming.
# Some litellm adapters return a non-streaming dict even
# when ``stream=True``; in that case, leave the result as-is.
try:
aggregated: Any = await self._aggregate_anthropic_events_to_message(
cast(AsyncIterator[Any], result)
)
except Exception as exc:
logger.error(
"Failed to aggregate streamed events into message",
extra={
"error": str(exc),
"error_type": type(exc).__name__,
"model": litellm_model,
},
)
raise UpstreamError(
f"Failed to aggregate upstream stream: {exc}",
status_code=502,
) from exc
return client_stream, aggregated, requested_model
return client_stream, result, requested_model
async def _forward_messages_via_litellm(
self,
+113
View File
@@ -0,0 +1,113 @@
"""Map an upstream `base_url` to the correct litellm provider prefix.
Used by `BaseUpstreamProvider.get_litellm_provider_prefix` so that custom /
generic provider rows (which inherit the base class default) get routed to
the right litellm backend instead of falling back to `openai/`.
The table is compiled from litellm 1.74's
`litellm/litellm_core_utils/get_llm_provider_logic.py` (the
`openai_compatible_endpoints` table) plus the providers documented at
https://docs.litellm.ai/docs/providers. Substring match is used so that
URLs with paths, ports, regional subdomains, etc. all resolve correctly.
Order matters: more specific needles must appear before more generic ones
(e.g. ``openai.azure.com`` before ``api.openai.com``).
"""
from __future__ import annotations
from urllib.parse import urlsplit
DEFAULT_PREFIX = "openai/"
LITELLM_HOST_PREFIX_MAP: tuple[tuple[str, str], ...] = (
# Azure must win over api.openai.com because the host ends with
# `openai.azure.com` and we don't want it picked up as plain OpenAI.
("openai.azure.com", "azure/"),
# Google
("generativelanguage.googleapis.com", "gemini/"),
("aiplatform.googleapis.com", "vertex_ai/"),
# First-class providers with native litellm prefixes
("api.openai.com", "openai/"),
("api.anthropic.com", "anthropic/"),
("api.groq.com", "groq/"),
("api.fireworks.ai", "fireworks_ai/"),
("api.x.ai", "xai/"),
("api.perplexity.ai", "perplexity/"),
("openrouter.ai", "openrouter/"),
("api.deepseek.com", "deepseek/"),
("api.together.xyz", "together_ai/"),
("codestral.mistral.ai", "codestral/"),
("api.mistral.ai", "mistral/"),
("api.cohere.com", "cohere_chat/"),
("api.cohere.ai", "cohere_chat/"),
("api.deepinfra.com", "deepinfra/"),
("api.endpoints.anyscale.com", "anyscale/"),
("api.cerebras.ai", "cerebras/"),
("inference.baseten.co", "baseten/"),
("api.sambanova.ai", "sambanova/"),
("api.ai21.com", "ai21_chat/"),
("api.friendli.ai", "friendliai/"),
("api.galadriel.com", "galadriel/"),
("api.llama.com", "meta_llama/"),
("api.featherless.ai", "featherless_ai/"),
("inference.api.nscale.com", "nscale/"),
("dashscope-intl.aliyuncs.com", "dashscope/"),
("api.moonshot.ai", "moonshot/"),
("api.moonshot.cn", "moonshot/"),
("api.minimax.io", "minimax/"),
("api.minimaxi.com", "minimax/"),
("platform.publicai.co", "publicai/"),
("api.synthetic.new", "synthetic/"),
("api.stima.tech", "apertis/"),
("nano-gpt.com", "nano-gpt/"),
("api.poe.com", "poe/"),
("llm.chutes.ai", "chutes/"),
("api.v0.dev", "v0/"),
("api.lambda.ai", "lambda_ai/"),
("api.hyperbolic.xyz", "hyperbolic/"),
("ai-gateway.vercel.sh", "vercel_ai_gateway/"),
("api.inference.wandb.ai", "wandb/"),
("integrate.api.nvidia.com", "nvidia_nim/"),
("api.studio.nebius.com", "nebius/"),
("api.novita.ai", "novita/"),
("ark.cn-beijing.volces.com", "volcengine/"),
("api.voyageai.com", "voyage/"),
("api.jina.ai", "jina_ai/"),
("api.aimlapi.com", "aiml/"),
("api.snowflakecomputing.com", "snowflake/"),
("databricks.com", "databricks/"),
("huggingface.co", "huggingface/"),
)
# Substrings that indicate an Ollama deployment regardless of port/scheme.
OLLAMA_HOST_HINTS: tuple[str, ...] = (
"localhost:11434",
"127.0.0.1:11434",
"ollama",
)
def detect_litellm_prefix(
base_url: str | None, default: str = DEFAULT_PREFIX
) -> str:
"""Return the litellm provider prefix (`"<provider>/"`) for `base_url`.
Falls back to `default` when the host doesn't match any known provider.
The default is `openai/` because every unmatched OpenAI-compatible
server is, by definition, an OpenAI-compatible server.
"""
if not base_url:
return default
parsed = urlsplit(base_url)
host = parsed.netloc.lower() or base_url.lower()
for needle, prefix in LITELLM_HOST_PREFIX_MAP:
if needle in host:
return prefix
if any(hint in host for hint in OLLAMA_HOST_HINTS):
return "ollama_chat/"
return default
+93
View File
@@ -0,0 +1,93 @@
"""Tests for `routstr.upstream.litellm_routing.detect_litellm_prefix`."""
from __future__ import annotations
import pytest
from routstr.upstream.litellm_routing import detect_litellm_prefix
@pytest.mark.parametrize(
"base_url,expected",
[
# The bug case: custom row pointing at Fireworks must NOT route to openai/.
("https://api.fireworks.ai/inference/v1", "fireworks_ai/"),
# Other OpenAI-compatible providers commonly plugged into the custom slot.
("https://api.groq.com/openai/v1", "groq/"),
("https://api.x.ai/v1", "xai/"),
("https://api.deepseek.com/v1", "deepseek/"),
("https://api.together.xyz/v1", "together_ai/"),
("https://api.perplexity.ai", "perplexity/"),
("https://openrouter.ai/api/v1", "openrouter/"),
("https://api.mistral.ai/v1", "mistral/"),
("https://codestral.mistral.ai/v1", "codestral/"),
("https://api.cohere.com/v1", "cohere_chat/"),
("https://api.cohere.ai/v1", "cohere_chat/"),
("https://api.deepinfra.com/v1/openai", "deepinfra/"),
("https://api.cerebras.ai/v1", "cerebras/"),
("https://api.sambanova.ai/v1", "sambanova/"),
("https://api.moonshot.cn/v1", "moonshot/"),
("https://api.moonshot.ai/v1", "moonshot/"),
("https://api.studio.nebius.com/v1", "nebius/"),
("https://api.novita.ai/v3/openai", "novita/"),
("https://api.lambda.ai/v1", "lambda_ai/"),
("https://api.aimlapi.com/v1", "aiml/"),
("https://api.featherless.ai/v1", "featherless_ai/"),
("https://integrate.api.nvidia.com/v1", "nvidia_nim/"),
("https://inference.baseten.co/v1", "baseten/"),
("https://ai-gateway.vercel.sh/v1", "vercel_ai_gateway/"),
("https://api.inference.wandb.ai/v1", "wandb/"),
("https://api.poe.com/v1", "poe/"),
("https://llm.chutes.ai/v1/", "chutes/"),
("https://api.v0.dev/v1", "v0/"),
("https://api.hyperbolic.xyz/v1", "hyperbolic/"),
("https://api.synthetic.new/openai/v1", "synthetic/"),
("https://api.stima.tech/v1", "apertis/"),
("https://nano-gpt.com/api/v1", "nano-gpt/"),
("https://api.friendli.ai/serverless/v1", "friendliai/"),
("https://api.galadriel.com/v1", "galadriel/"),
("https://api.llama.com/compat/v1", "meta_llama/"),
("https://api.minimax.io/v1", "minimax/"),
("https://api.minimaxi.com/v1", "minimax/"),
("https://platform.publicai.co/v1", "publicai/"),
("https://inference.api.nscale.com/v1", "nscale/"),
("https://dashscope-intl.aliyuncs.com/compatible-mode/v1", "dashscope/"),
("https://api.endpoints.anyscale.com/v1", "anyscale/"),
("https://api.ai21.com/studio/v1", "ai21_chat/"),
("https://ark.cn-beijing.volces.com/api/v3", "volcengine/"),
("https://api.voyageai.com/v1", "voyage/"),
("https://api.jina.ai/v1", "jina_ai/"),
("https://api.snowflakecomputing.com", "snowflake/"),
("https://my-workspace.databricks.com/serving-endpoints", "databricks/"),
("https://huggingface.co/api/inference", "huggingface/"),
# First-class providers — URL detection still produces the right prefix
# so subclasses without an explicit override stay correct.
("https://api.openai.com/v1", "openai/"),
("https://api.anthropic.com/v1", "anthropic/"),
("https://generativelanguage.googleapis.com/v1beta/openai", "gemini/"),
("https://us-central1-aiplatform.googleapis.com/v1", "vertex_ai/"),
# Azure ordering: must beat api.openai.com.
("https://my-resource.openai.azure.com/openai/deployments/foo", "azure/"),
# Ollama hints.
("http://localhost:11434/v1", "ollama_chat/"),
("http://127.0.0.1:11434/v1", "ollama_chat/"),
("http://my-ollama-host:11434/v1", "ollama_chat/"),
# Casing and trailing slash normalisation.
("HTTPS://API.FIREWORKS.AI/INFERENCE/V1/", "fireworks_ai/"),
# Unknown host falls back to openai/ (still OpenAI-compatible by convention).
("https://example.com/v1", "openai/"),
("", "openai/"),
],
)
def test_detect_litellm_prefix(base_url: str, expected: str) -> None:
assert detect_litellm_prefix(base_url) == expected
def test_detect_litellm_prefix_none_uses_default() -> None:
assert detect_litellm_prefix(None) == "openai/"
def test_detect_litellm_prefix_custom_default() -> None:
assert detect_litellm_prefix("https://example.com", default="anthropic/") == (
"anthropic/"
)
+312 -2
View File
@@ -198,7 +198,48 @@ def test_compute_refund_invalid_unit() -> None:
def test_default_provider_does_not_support_anthropic_messages() -> None:
assert BaseUpstreamProvider.supports_anthropic_messages is False
assert BaseUpstreamProvider.litellm_provider_prefix == "openai/"
# Class default is None so URL detection runs at dispatch time.
assert BaseUpstreamProvider.litellm_provider_prefix is None
def test_base_provider_resolves_prefix_from_base_url() -> None:
"""The bug case: a custom row pointing at Fireworks must resolve to
`fireworks_ai/` instead of falling back to `openai/`."""
fireworks = BaseUpstreamProvider(
base_url="https://api.fireworks.ai/inference/v1", api_key="sk-test"
)
assert fireworks.get_litellm_provider_prefix() == "fireworks_ai/"
groq = BaseUpstreamProvider(
base_url="https://api.groq.com/openai/v1", api_key="sk-test"
)
assert groq.get_litellm_provider_prefix() == "groq/"
xai = BaseUpstreamProvider(
base_url="https://api.x.ai/v1", api_key="sk-test"
)
assert xai.get_litellm_provider_prefix() == "xai/"
deepseek = BaseUpstreamProvider(
base_url="https://api.deepseek.com/v1", api_key="sk-test"
)
assert deepseek.get_litellm_provider_prefix() == "deepseek/"
unknown = BaseUpstreamProvider(
base_url="https://example.com/v1", api_key="sk-test"
)
assert unknown.get_litellm_provider_prefix() == "openai/"
def test_subclass_prefix_wins_over_url_detection() -> None:
"""A subclass override must beat URL detection. e.g. configuring an
AnthropicUpstreamProvider with a fireworks URL still produces
``anthropic/`` (defensive: subclasses pin their backend on purpose)."""
from routstr.upstream.anthropic import AnthropicUpstreamProvider
p = AnthropicUpstreamProvider(api_key="sk-test")
p.base_url = "https://api.fireworks.ai/inference/v1"
assert p.get_litellm_provider_prefix() == "anthropic/"
def test_anthropic_provider_supports_native_messages() -> None:
@@ -371,7 +412,13 @@ async def test_non_streaming_dispatches_via_litellm_and_returns_anthropic_respon
assert kwargs["model"] == "openai/openai/gpt-4o-mini"
assert kwargs["api_base"] == "http://test"
assert kwargs["api_key"] == "upstream-key"
assert kwargs["stream"] is False
# The dispatcher always streams from upstream and aggregates back
# when the client wants a non-streaming response (sidesteps
# provider-specific non-streaming caps like Fireworks's max_tokens
# > 4096). When the mock returns a plain dict, the dispatcher
# detects the lack of __aiter__ and skips aggregation, so the
# client still gets the same Response shape.
assert kwargs["stream"] is True
assert kwargs["messages"] == [{"role": "user", "content": "hi"}]
assert kwargs["max_tokens"] == 64
return upstream_response
@@ -989,3 +1036,266 @@ async def test_forward_x_cashu_request_skips_litellm_for_count_tokens() -> None:
max_cost_for_model=10_000,
model_obj=model,
)
# ---------------------------------------------------------------------------
# Upstream-always-streams + aggregate-on-non-streaming
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_aggregator_assembles_text_message_from_events() -> None:
"""When the client wants a non-streaming response, the dispatcher
drains the upstream stream and assembles a single Anthropic Message
dict. This is what makes Fireworks (which rejects max_tokens > 4096
unless stream=true) work transparently for non-streaming clients."""
provider = _make_provider()
events: list[dict] = [
{
"type": "message_start",
"message": {
"id": "msg_aggregated",
"type": "message",
"role": "assistant",
"model": "accounts/fireworks/models/glm-5",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 0},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "Hello"},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": ", world!"},
},
{"type": "content_block_stop", "index": 0},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 42},
},
{"type": "message_stop"},
]
async def event_iter() -> AsyncIterator[dict]:
for event in events:
yield event
message = await provider._aggregate_anthropic_events_to_message(event_iter())
assert message["id"] == "msg_aggregated"
assert message["role"] == "assistant"
assert message["stop_reason"] == "end_turn"
assert message["stop_sequence"] is None
assert message["content"] == [{"type": "text", "text": "Hello, world!"}]
assert message["usage"]["input_tokens"] == 10
assert message["usage"]["output_tokens"] == 42
@pytest.mark.asyncio
async def test_aggregator_parses_tool_use_input_json_delta() -> None:
"""tool_use blocks split their `input` across multiple
`input_json_delta` chunks. The aggregator must concatenate and parse
them into a single JSON object."""
provider = _make_provider()
events: list[dict] = [
{
"type": "message_start",
"message": {
"id": "msg_tool",
"type": "message",
"role": "assistant",
"model": "test-model",
"content": [],
"usage": {"input_tokens": 5, "output_tokens": 0},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "tool_use",
"id": "toolu_1",
"name": "calc",
"input": {},
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '{"x":'},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": ' 7}'},
},
{"type": "content_block_stop", "index": 0},
{
"type": "message_delta",
"delta": {"stop_reason": "tool_use"},
"usage": {"output_tokens": 3},
},
{"type": "message_stop"},
]
async def event_iter() -> AsyncIterator[dict]:
for event in events:
yield event
message = await provider._aggregate_anthropic_events_to_message(event_iter())
assert message["stop_reason"] == "tool_use"
assert message["content"][0]["type"] == "tool_use"
assert message["content"][0]["input"] == {"x": 7}
@pytest.mark.asyncio
async def test_aggregator_parses_sse_byte_chunks() -> None:
"""litellm's anthropic adapter often yields raw SSE byte chunks. The
aggregator must parse the SSE wire format, not just typed dicts."""
provider = _make_provider()
sse = (
b'event: message_start\n'
b'data: {"type":"message_start","message":{"id":"m1","type":"message",'
b'"role":"assistant","model":"x","content":[],"usage":{"input_tokens":1,"output_tokens":0}}}\n\n'
b'event: content_block_start\n'
b'data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n'
b'event: content_block_delta\n'
b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}\n\n'
b'event: content_block_stop\n'
b'data: {"type":"content_block_stop","index":0}\n\n'
b'event: message_delta\n'
b'data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":2}}\n\n'
b'event: message_stop\n'
b'data: {"type":"message_stop"}\n\n'
)
async def chunk_iter() -> AsyncIterator[bytes]:
# Split mid-event to exercise the partial-buffer path.
yield sse[:100]
yield sse[100:]
message = await provider._aggregate_anthropic_events_to_message(chunk_iter())
assert message["content"] == [{"type": "text", "text": "hi"}]
assert message["stop_reason"] == "end_turn"
assert message["usage"]["output_tokens"] == 2
@pytest.mark.asyncio
async def test_dispatch_always_streams_upstream_and_aggregates_for_non_streaming_client( # noqa: E501
) -> None:
"""End-to-end: client says stream=false; dispatcher upstream-streams
and returns an aggregated Anthropic Message dict."""
provider = _make_provider()
model = _make_model()
body = _anthropic_request_body(stream=False)
events: list[dict] = [
{
"type": "message_start",
"message": {
"id": "msg_e2e",
"type": "message",
"role": "assistant",
"model": "openai/gpt-4o-mini",
"content": [],
"usage": {"input_tokens": 4, "output_tokens": 0},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "ok"},
},
{"type": "content_block_stop", "index": 0},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 1},
},
{"type": "message_stop"},
]
async def event_iter() -> AsyncIterator[dict]:
for event in events:
yield event
captured_kwargs: dict[str, Any] = {}
async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]:
captured_kwargs.update(kwargs)
return event_iter()
with patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
):
client_stream, result, requested_model = (
await provider._dispatch_anthropic_messages(
request_body=body,
model_obj=model,
)
)
# Upstream was streamed regardless of client preference.
assert captured_kwargs["stream"] is True
# Client wanted non-streaming → aggregator ran.
assert client_stream is False
assert isinstance(result, dict)
assert result["content"] == [{"type": "text", "text": "ok"}]
assert result["stop_reason"] == "end_turn"
assert requested_model == "openai/gpt-4o-mini"
@pytest.mark.asyncio
async def test_dispatch_uses_url_detected_prefix_for_fireworks_custom_row() -> None:
"""The original bug: a custom-typed row pointing at Fireworks must
dispatch with `fireworks_ai/`, not `openai/`."""
provider = BaseUpstreamProvider(
base_url="https://api.fireworks.ai/inference/v1",
api_key="fw-key",
)
model = _make_model(model_id="accounts/fireworks/models/glm-5")
body = _anthropic_request_body(stream=True)
captured_kwargs: dict[str, Any] = {}
async def empty_iter() -> AsyncIterator[dict]:
if False:
yield {}
async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]:
captured_kwargs.update(kwargs)
return empty_iter()
with patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
):
await provider._dispatch_anthropic_messages(
request_body=body, model_obj=model
)
assert captured_kwargs["model"] == (
"fireworks_ai/accounts/fireworks/models/glm-5"
)
assert captured_kwargs["api_base"] == "https://api.fireworks.ai/inference/v1"