mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
make sure to forward to the right upstream
This commit is contained in:
+174
-6
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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/"
|
||||
)
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user