From 85a3d3adc0cbbc437d4d64a135ca9c215df1ade3 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Apr 2026 22:45:22 +0200 Subject: [PATCH] feat: route /v1/messages via litellm when upstream lacks native support --- routstr/upstream/anthropic.py | 2 + routstr/upstream/azure.py | 1 + routstr/upstream/base.py | 262 ++++++++++++++++++++++++++++++++- routstr/upstream/fireworks.py | 1 + routstr/upstream/gemini.py | 1 + routstr/upstream/groq.py | 1 + routstr/upstream/ollama.py | 1 + routstr/upstream/openrouter.py | 2 + routstr/upstream/perplexity.py | 1 + routstr/upstream/xai.py | 1 + 10 files changed, 271 insertions(+), 2 deletions(-) diff --git a/routstr/upstream/anthropic.py b/routstr/upstream/anthropic.py index 3f228e9c..5e48f058 100644 --- a/routstr/upstream/anthropic.py +++ b/routstr/upstream/anthropic.py @@ -13,6 +13,8 @@ class AnthropicUpstreamProvider(BaseUpstreamProvider): provider_type = "anthropic" default_base_url = "https://api.anthropic.com/v1" platform_url = "https://console.anthropic.com/settings/keys" + supports_anthropic_messages = True + litellm_provider_prefix = "anthropic/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index e412c466..11cee17c 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -13,6 +13,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider): provider_type = "azure" default_base_url = None platform_url = "https://portal.azure.com/" + litellm_provider_prefix = "azure/" def __init__( self, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index c2e9f8dd..6872efd0 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -6,8 +6,8 @@ import json import re import traceback import uuid -from collections.abc import AsyncGenerator -from typing import Mapping +from collections.abc import AsyncGenerator, AsyncIterator +from typing import Any, Mapping, cast import httpx from fastapi import BackgroundTasks, HTTPException, Request @@ -63,6 +63,9 @@ class BaseUpstreamProvider: default_base_url: str | None = None platform_url: str | None = None + supports_anthropic_messages: bool = False + litellm_provider_prefix: str = "openai/" + base_url: str api_key: str provider_fee: float = 1.05 @@ -1400,6 +1403,242 @@ class BaseUpstreamProvider: except Exception: raise + @staticmethod + def _coerce_litellm_payload(payload: object) -> dict: + if isinstance(payload, dict): + return dict(payload) + if hasattr(payload, "model_dump"): + return payload.model_dump() # type: ignore[no-any-return] + if hasattr(payload, "dict") and callable(payload.dict): # type: ignore[union-attr] + return payload.dict() # type: ignore[no-any-return,union-attr] + if hasattr(payload, "__dict__"): + return dict(payload.__dict__) + raise TypeError(f"Cannot coerce {type(payload).__name__} to dict") + + async def _forward_messages_via_litellm( + self, + request_body: bytes | None, + key: ApiKey, + session: AsyncSession, + max_cost_for_model: int, + model_obj: Model, + ) -> Response | StreamingResponse: + """Translate /v1/messages to upstream chat/completions via litellm. + + Used when the upstream provider does not natively serve Anthropic + Messages (i.e. supports_anthropic_messages is False). Cost + tracking and metadata injection mirror the native messages path. + """ + import litellm + + if not request_body: + raise UpstreamError( + "Missing request body for /v1/messages", status_code=400 + ) + + try: + body: dict = json.loads(request_body) + except json.JSONDecodeError as exc: + raise UpstreamError( + f"Invalid JSON in /v1/messages body: {exc}", status_code=400 + ) from exc + + body.pop("model", None) + stream = bool(body.pop("stream", False)) + + requested_model = ( + (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}" + + kwargs: dict = { + "model": litellm_model, + "api_base": self.base_url, + "api_key": self.api_key, + "stream": stream, + **body, + } + + logger.info( + "Dispatching /v1/messages via litellm", + extra={ + "model": litellm_model, + "stream": stream, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + try: + result = await litellm.anthropic.messages.acreate(**kwargs) + except Exception as exc: + logger.error( + "litellm dispatch failed", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "model": litellm_model, + }, + ) + raise UpstreamError( + f"Upstream error via litellm: {exc}", status_code=502 + ) from exc + + if stream: + return self._stream_litellm_messages( + cast(AsyncIterator[Any], result), + key, + max_cost_for_model, + requested_model, + ) + + response_json = self._coerce_litellm_payload(result) + if requested_model and "model" in response_json: + response_json["model"] = requested_model + + cost_data = await adjust_payment_for_tokens( + key, response_json, session, max_cost_for_model + ) + self.inject_cost_metadata(response_json, cost_data, key) + + return Response( + content=json.dumps(response_json).encode(), + status_code=200, + media_type="application/json", + ) + + def _stream_litellm_messages( + self, + iterator: AsyncIterator[Any], + key: ApiKey, + max_cost_for_model: int, + requested_model: str | None, + ) -> StreamingResponse: + """Re-emit a litellm Anthropic-event iterator as SSE bytes with + cost reconciliation at end of stream.""" + + async def stream_with_cost() -> AsyncGenerator[bytes, None]: + usage_finalized = False + last_model_seen: str | None = None + input_tokens = 0 + output_tokens = 0 + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if not fresh_key: + usage_finalized = True + return None + try: + fallback: dict = { + "model": last_model_seen or "unknown", + "usage": None, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, + fallback, + new_session, + max_cost_for_model, + ) + usage_finalized = True + return ( + f"event: cost\ndata: " + f"{json.dumps({'cost': cost_data})}\n\n" + ).encode() + except Exception: + usage_finalized = True + return None + + try: + async for chunk in iterator: + event = self._coerce_litellm_payload(chunk) + event_type = str(event.get("type") or "") + + if requested_model: + msg = event.get("message") + if isinstance(msg, dict) and "model" in msg: + msg["model"] = requested_model + if "model" in event: + event["model"] = requested_model + + msg_for_meta = event.get("message") + if ( + isinstance(msg_for_meta, dict) + and msg_for_meta.get("model") + ): + last_model_seen = str(msg_for_meta["model"]) + + if isinstance(msg_for_meta, dict) and isinstance( + msg_for_meta.get("usage"), dict + ): + usage = msg_for_meta["usage"] + input_tokens += int(usage.get("input_tokens") or 0) + output_tokens += int(usage.get("output_tokens") or 0) + if isinstance(event.get("usage"), dict): + usage = event["usage"] + input_tokens += int(usage.get("input_tokens") or 0) + output_tokens += int(usage.get("output_tokens") or 0) + + payload = json.dumps(event) + if event_type: + yield ( + f"event: {event_type}\ndata: {payload}\n\n" + ).encode() + else: + yield f"data: {payload}\n\n".encode() + + if input_tokens > 0 or output_tokens > 0: + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if fresh_key: + try: + combined_data: dict = { + "model": last_model_seen or "unknown", + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + }, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, + combined_data, + new_session, + max_cost_for_model, + ) + self.inject_cost_metadata( + combined_data, cost_data, fresh_key + ) + usage_finalized = True + yield ( + f"event: cost\ndata: " + f"{json.dumps({'cost': cost_data})}\n\n" + ).encode() + except Exception: + pass + + if not usage_finalized: + cost_event = await finalize_without_usage() + if cost_event is not None: + yield cost_event + + except Exception: + if not usage_finalized: + await finalize_without_usage() + raise + + return StreamingResponse( + stream_with_cost(), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, + ) + async def forward_request( self, request: Request, @@ -1426,6 +1665,20 @@ class BaseUpstreamProvider: Response or StreamingResponse from upstream with cost tracking """ path = self.normalize_request_path(path, model_obj) + + if ( + path.endswith("messages") + and not path.endswith("count_tokens") + and not self.supports_anthropic_messages + ): + return await self._forward_messages_via_litellm( + request_body=request_body, + key=key, + session=session, + max_cost_for_model=max_cost_for_model, + model_obj=model_obj, + ) + url = self.build_request_url(path, model_obj) original_model_id = ( @@ -2609,6 +2862,11 @@ class BaseUpstreamProvider: if path.startswith("v1/"): path = path.replace("v1/", "") + # TODO: route /messages via litellm here too when + # supports_anthropic_messages is False. Bearer-key path already + # does this via _forward_messages_via_litellm; x-cashu needs the + # additional refund-on-overspend reconciliation logic. + url = f"{self.base_url}/{path}" request_body = await request.body() diff --git a/routstr/upstream/fireworks.py b/routstr/upstream/fireworks.py index ed1ea053..e0bb1fec 100644 --- a/routstr/upstream/fireworks.py +++ b/routstr/upstream/fireworks.py @@ -12,6 +12,7 @@ class FireworksUpstreamProvider(BaseUpstreamProvider): provider_type = "fireworks" default_base_url = "https://api.fireworks.ai/inference/v1" platform_url = "https://app.fireworks.ai/settings/users/api-keys" + litellm_provider_prefix = "fireworks_ai/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index 8a416bfa..eaeb6d9d 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -23,6 +23,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): provider_type = "gemini" default_base_url = "https://generativelanguage.googleapis.com/v1beta" platform_url = "https://aistudio.google.com/app/apikey" + litellm_provider_prefix = "gemini/" def __init__( self, diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index 11ab8ce4..4020b2df 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -12,6 +12,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider): provider_type = "groq" default_base_url = "https://api.groq.com/openai/v1" platform_url = "https://console.groq.com/keys" + litellm_provider_prefix = "groq/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index eff5d5bb..b8df6f3f 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -21,6 +21,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): provider_type = "ollama" default_base_url = "http://localhost:11434" platform_url = None + litellm_provider_prefix = "ollama_chat/" def __init__( self, diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1736f1d7..2cbc34de 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -15,6 +15,8 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): provider_type = "openrouter" default_base_url = "https://openrouter.ai/api/v1" platform_url = "https://openrouter.ai/settings/keys" + supports_anthropic_messages = True + litellm_provider_prefix = "openrouter/" def __init__(self, api_key: str, provider_fee: float = 1.06): """Initialize OpenRouter provider with API key. diff --git a/routstr/upstream/perplexity.py b/routstr/upstream/perplexity.py index b73881d8..55a9116e 100644 --- a/routstr/upstream/perplexity.py +++ b/routstr/upstream/perplexity.py @@ -13,6 +13,7 @@ class PerplexityUpstreamProvider(BaseUpstreamProvider): provider_type = "perplexity" default_base_url = "https://api.perplexity.ai/" platform_url = "https://www.perplexity.ai/account/api/keys" + litellm_provider_prefix = "perplexity/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__( diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 99e2d35a..b46676cb 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -13,6 +13,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider): provider_type = "x-ai" default_base_url = "https://api.x.ai/v1" platform_url = "https://console.x.ai/" + litellm_provider_prefix = "xai/" def __init__(self, api_key: str, provider_fee: float = 1.01): super().__init__(