diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 18731695..242781f9 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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, diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py new file mode 100644 index 00000000..6ce2687c --- /dev/null +++ b/routstr/upstream/litellm_routing.py @@ -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 (`"/"`) 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 diff --git a/tests/unit/test_litellm_routing.py b/tests/unit/test_litellm_routing.py new file mode 100644 index 00000000..264a142d --- /dev/null +++ b/tests/unit/test_litellm_routing.py @@ -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/" + ) diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 4403afb8..20482bb1 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -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"