From 145777ffd0caa33cd8c9bde046475fdb44d98721 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 3 May 2026 15:35:49 +0200 Subject: [PATCH] clean up --- routstr/core/main.py | 5 + routstr/upstream/base.py | 504 +++----------------------- routstr/upstream/litellm_routing.py | 49 +++ routstr/upstream/messages_dispatch.py | 458 +++++++++++++++++++++++ 4 files changed, 569 insertions(+), 447 deletions(-) create mode 100644 routstr/upstream/messages_dispatch.py diff --git a/routstr/core/main.py b/routstr/core/main.py index e22173df..60913447 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -23,6 +23,7 @@ from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..upstream.auto_topup import periodic_auto_topup +from ..upstream.litellm_routing import configure_litellm from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout from .admin import admin_router from .db import create_session, init_db, run_migrations @@ -56,6 +57,10 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: routstr_fee_task = None try: + # Apply litellm-wide settings (drop_params, chat-completions URL, + # debug logging) before any upstream provider dispatches a request. + configure_litellm() + # Run database migrations on startup run_migrations() diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 242781f9..3aca949e 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3,7 +3,6 @@ from __future__ import annotations import asyncio import hashlib import json -import os import re import traceback import uuid @@ -11,35 +10,6 @@ from collections.abc import AsyncGenerator, AsyncIterator from typing import Any, Mapping, cast import httpx -import litellm - -if os.getenv("LITELLM_DEBUG") == "1": - try: - litellm._turn_on_debug() # type: ignore[no-untyped-call] - except Exception: - pass - -# Force litellm's Anthropic-messages adapter to use OpenAI Chat Completions -# (POST /chat/completions) instead of OpenAI Responses API (POST /responses) -# for openai-prefixed providers. OpenAI-compatible upstreams like Google's -# generativelanguage compat endpoint expose /chat/completions but not -# /responses, which produces a 404. Override with -# `LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES=1` if a future upstream -# requires the Responses API. -if os.getenv("LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES") != "1": - try: - litellm.use_chat_completions_url_for_anthropic_messages = True - except Exception: - pass - -# Silently drop Anthropic-Messages-only parameters (e.g. `context_management`, -# `cache_control`, `thinking`) when translating to providers that don't -# accept them. Without this, litellm raises UnsupportedParamsError for any -# unrecognized field and rejects the whole request. Override with -# `LITELLM_STRICT_PARAMS=1` if an integration depends on the strict -# behavior. -if os.getenv("LITELLM_STRICT_PARAMS") != "1": - litellm.drop_params = True from fastapi import BackgroundTasks, HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel @@ -71,6 +41,7 @@ from ..payment.models import ( ) from ..payment.price import sats_usd_price from ..wallet import recieve_token, send_token +from . import messages_dispatch from .litellm_routing import detect_litellm_prefix logger = get_logger(__name__) @@ -1520,195 +1491,24 @@ class BaseUpstreamProvider: except Exception: raise - @staticmethod - def _coerce_litellm_payload(payload: object) -> dict: - """Convert a litellm event into a plain dict. + # ------------------------------------------------------------------ + # Litellm /v1/messages dispatch (thin wrappers) + # + # The actual translation logic lives in ``messages_dispatch``. These + # method shims exist so subclasses and tests can keep the original + # provider-bound API. + # ------------------------------------------------------------------ - Non-streaming responses come back as Anthropic-shaped pydantic - models or dicts. Streaming may yield raw bytes/str (SSE-encoded); - those go through ``_events_from_chunk`` instead, not here. - """ - if isinstance(payload, dict): - return dict(payload) - if hasattr(payload, "model_dump"): - return cast(dict, payload.model_dump()) - raise TypeError(f"Cannot coerce {type(payload).__name__} to dict") - - @staticmethod - def _parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: - """Parse complete SSE event blocks out of a byte buffer. - - Returns (events, remaining_buffer). Events are JSON objects parsed - from one or more `data:` lines per block. Comments, blank lines, - and `[DONE]` sentinels are ignored. A trailing partial block is - preserved in remaining_buffer. - """ - events: list[dict] = [] - while True: - sep = buffer.find(b"\n\n") - if sep < 0: - sep_rn = buffer.find(b"\r\n\r\n") - if sep_rn < 0: - break - block = buffer[:sep_rn] - buffer = buffer[sep_rn + 4 :] - else: - block = buffer[:sep] - buffer = buffer[sep + 2 :] - - data_lines: list[str] = [] - for raw_line in block.replace(b"\r\n", b"\n").split(b"\n"): - line = raw_line.decode("utf-8", errors="replace") - if line.startswith(":"): - continue - if line.startswith("data:"): - data_lines.append(line[5:].lstrip()) - if not data_lines: - continue - payload = "\n".join(data_lines).strip() - if not payload or payload == "[DONE]": - continue - try: - obj = json.loads(payload) - except json.JSONDecodeError: - continue - if isinstance(obj, dict): - events.append(obj) - return events, buffer + _coerce_litellm_payload = staticmethod(messages_dispatch.coerce_litellm_payload) + _parse_sse_blocks = staticmethod(messages_dispatch.parse_sse_blocks) + _events_from_chunk = staticmethod(messages_dispatch.events_from_chunk) 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]: - """Normalize a stream chunk into one or more event dicts. - - litellm.anthropic.messages.acreate(stream=True) yields raw SSE - bytes in practice; some adapters yield strings or typed events. - Handle all three. - """ - if isinstance(chunk, (bytes, bytearray)): - sse_buffer += bytes(chunk) - events, sse_buffer = self._parse_sse_blocks(sse_buffer) - return events, sse_buffer - if isinstance(chunk, str): - sse_buffer += chunk.encode("utf-8") - events, sse_buffer = self._parse_sse_blocks(sse_buffer) - return events, sse_buffer - return [self._coerce_litellm_payload(chunk)], sse_buffer + return await messages_dispatch.aggregate_anthropic_events_to_message( + iterator + ) async def _dispatch_anthropic_messages( self, @@ -1717,145 +1517,15 @@ class BaseUpstreamProvider: *, log_extra: dict[str, Any] | None = None, ) -> tuple[bool, Any, str | None]: - """Call litellm.anthropic.messages.acreate and return - (stream, result, requested_model). - - Shared by bearer-key and x-cashu paths. Raises UpstreamError on - bad input or upstream failure. - """ - 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` 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* - # unsupported params; these newer/extension fields get passed - # through verbatim and the upstream rejects them with a 400. - # Pop them here so the request reaches the upstream cleanly. - anthropic_only_fields = ( - "thinking", - "cache_control", - "context_management", - "output_config", - "mcp_servers", - "service_tier", - "anthropic_version", - "anthropic_beta", + return await messages_dispatch.dispatch_anthropic_messages( + request_body=request_body, + model_obj=model_obj, + base_url=self.base_url, + api_key=self.api_key, + provider_prefix=self.get_litellm_provider_prefix(), + transform_model_name=self.transform_model_name, + log_extra=log_extra, ) - dropped: dict[str, Any] = {} - for field in anthropic_only_fields: - if field in body: - dropped[field] = body.pop(field) - if dropped: - logger.debug( - "Dropped anthropic-only fields before litellm dispatch", - extra={"dropped_keys": sorted(dropped.keys())}, - ) - - # Convention: `model.id` is the canonical upstream model name; - # `forwarded_model_id` is the public alias the internal API - # exposes and echoes back to the client. - 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) - 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": upstream_stream, - **body, - } - - logger.info( - "Dispatching /v1/messages via litellm", - extra={ - "model": litellm_model, - "resolved_provider": prefix.rstrip("/"), - "client_stream": client_stream, - "upstream_stream": upstream_stream, - **(log_extra or {}), - }, - ) - - try: - result = await litellm.anthropic.messages.acreate(**kwargs) - except Exception as exc: - exc_message = getattr(exc, "message", None) or str(exc) or repr(exc) - exc_status = getattr(exc, "status_code", None) - exc_response = getattr(exc, "response", None) - response_text = None - if exc_response is not None: - try: - response_text = getattr(exc_response, "text", str(exc_response)) - except Exception: - response_text = "" - logger.error( - "litellm dispatch failed", - extra={ - "error": exc_message, - "error_type": type(exc).__name__, - "status_code": exc_status, - "llm_provider": getattr(exc, "llm_provider", None), - "body": getattr(exc, "body", None), - "response_text": response_text, - "model": litellm_model, - "api_base": self.base_url, - }, - ) - raise UpstreamError( - f"Upstream error via litellm: {exc_message}", - status_code=exc_status if isinstance(exc_status, int) else 502, - ) from exc - - 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, @@ -1885,7 +1555,7 @@ class BaseUpstreamProvider: requested_model, ) - response_json = self._coerce_litellm_payload(result) + response_json = messages_dispatch.coerce_litellm_payload(result) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -1934,7 +1604,7 @@ class BaseUpstreamProvider: request_id, ) - response_json = self._coerce_litellm_payload(result) + response_json = messages_dispatch.coerce_litellm_payload(result) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -1947,7 +1617,7 @@ class BaseUpstreamProvider: response_headers: dict[str, str] = {} if cost_data: - refund_amount = self._compute_refund( + refund_amount = messages_dispatch.compute_refund( amount, unit, cost_data.total_msats ) if refund_amount > 0: @@ -1975,13 +1645,7 @@ class BaseUpstreamProvider: media_type="application/json", ) - @staticmethod - def _compute_refund(amount: int, unit: str, cost_msats: int) -> int: - if unit == "msat": - return amount - cost_msats - if unit == "sat": - return amount - (cost_msats + 999) // 1000 - raise ValueError(f"Invalid unit: {unit}") + _compute_refund = staticmethod(messages_dispatch.compute_refund) def _stream_litellm_messages( self, @@ -1990,8 +1654,8 @@ class BaseUpstreamProvider: 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.""" + """Re-emit a litellm Anthropic-event iterator as live SSE bytes + with cost reconciliation appended at end of stream.""" async def stream_with_cost() -> AsyncGenerator[bytes, None]: usage_finalized = False @@ -2030,51 +1694,15 @@ class BaseUpstreamProvider: usage_finalized = True return None - sse_buffer = b"" try: - async for chunk in iterator: - events, sse_buffer = self._events_from_chunk( - chunk, sse_buffer - ) - for event in events: - 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() + async for annotated in messages_dispatch.stream_annotated_events( + iterator, requested_model + ): + if annotated.model: + last_model_seen = annotated.model + input_tokens += annotated.input_tokens + output_tokens += annotated.output_tokens + yield annotated.sse_bytes if input_tokens > 0 or output_tokens > 0: async with create_session() as new_session: @@ -2134,50 +1762,32 @@ class BaseUpstreamProvider: payment_token_hash: str | None, request_id: str | None, ) -> StreamingResponse: - """Buffer a litellm Anthropic-event iterator, compute cost, refund - on overspend, and re-emit the events as SSE with X-Cashu set on - the response header. + """Buffer a litellm stream end-to-end, compute cost, then replay. + + Note this is **not** true streaming — the full event sequence is + accumulated into memory before a single byte is sent to the + client. The constraint is the ``X-Cashu`` refund token, which must + be set as a response *header* and therefore has to be known before + the response begins. The bearer-key path + (:meth:`_stream_litellm_messages`) avoids this by emitting cost as + a trailing ``event: cost`` SSE message; switching x-cashu to the + same trailing-event contract would let this path stream live, at + the cost of a wire-format change for clients that read ``X-Cashu`` + from headers today. """ buffered: list[bytes] = [] last_model_seen: str | None = None input_tokens = 0 output_tokens = 0 - sse_buffer = b"" - async for chunk in iterator: - events, sse_buffer = self._events_from_chunk(chunk, sse_buffer) - for event in events: - 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: - buffered.append( - f"event: {event_type}\ndata: {payload}\n\n".encode() - ) - else: - buffered.append(f"data: {payload}\n\n".encode()) + async for annotated in messages_dispatch.stream_annotated_events( + iterator, requested_model + ): + if annotated.model: + last_model_seen = annotated.model + input_tokens += annotated.input_tokens + output_tokens += annotated.output_tokens + buffered.append(annotated.sse_bytes) response_headers: dict[str, str] = { "Cache-Control": "no-cache", @@ -2197,7 +1807,7 @@ class BaseUpstreamProvider: response_data, max_cost_for_model ) if cost_data: - refund_amount = self._compute_refund( + refund_amount = messages_dispatch.compute_refund( amount, unit, cost_data.total_msats ) if refund_amount > 0: diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py index 6ce2687c..0b2a92a3 100644 --- a/routstr/upstream/litellm_routing.py +++ b/routstr/upstream/litellm_routing.py @@ -16,8 +16,11 @@ Order matters: more specific needles must appear before more generic ones from __future__ import annotations +import os from urllib.parse import urlsplit +import litellm + DEFAULT_PREFIX = "openai/" LITELLM_HOST_PREFIX_MAP: tuple[tuple[str, str], ...] = ( @@ -111,3 +114,49 @@ def detect_litellm_prefix( return "ollama_chat/" return default + + +_configured = False + + +def configure_litellm() -> None: + """Apply litellm global settings used by the messages-dispatch path. + + Idempotent: safe to call from both app startup and module-level + initializers without side effects on the second invocation. + + Settings applied: + + * ``LITELLM_DEBUG=1`` enables litellm's verbose debug logger. + * Forces the Anthropic-messages adapter to call OpenAI Chat Completions + (POST ``/chat/completions``) instead of the Responses API (POST + ``/responses``) for ``openai/``-prefixed providers. OpenAI-compatible + upstreams like Google's generativelanguage compat endpoint expose + ``/chat/completions`` but not ``/responses``, which would 404. Set + ``LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES=1`` to opt out. + * Silently drops Anthropic-Messages-only parameters (``thinking``, + ``cache_control``, ``context_management``, ...) when translating to + providers that don't accept them, instead of raising + ``UnsupportedParamsError``. Set ``LITELLM_STRICT_PARAMS=1`` to opt + out. + """ + global _configured + if _configured: + return + + if os.getenv("LITELLM_DEBUG") == "1": + try: + litellm._turn_on_debug() # type: ignore[no-untyped-call] + except Exception: + pass + + if os.getenv("LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES") != "1": + try: + litellm.use_chat_completions_url_for_anthropic_messages = True + except Exception: + pass + + if os.getenv("LITELLM_STRICT_PARAMS") != "1": + litellm.drop_params = True + + _configured = True diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py new file mode 100644 index 00000000..a492e481 --- /dev/null +++ b/routstr/upstream/messages_dispatch.py @@ -0,0 +1,458 @@ +"""Pure helpers for translating ``/v1/messages`` to upstream chat completions +via litellm. + +This module owns the litellm/Anthropic-Messages translation layer: + +* SSE parsing (``parse_sse_blocks``, ``events_from_chunk``) +* Payload coercion (``coerce_litellm_payload``) +* Stream aggregation (``aggregate_anthropic_events_to_message``) — drains + a streamed Anthropic event sequence into a single Message dict +* Per-event annotation for streaming (``annotate_event``, + ``stream_annotated_events``) — handles the model-rewrite + token-tally + bookkeeping shared by the bearer-key and x-cashu streaming paths +* The dispatch entry point (``dispatch_anthropic_messages``) +* Refund math (``compute_refund``) + +Nothing in here touches ``BaseUpstreamProvider``; the thin instance methods +on the provider class forward to these functions and only retain logic that +genuinely needs ``self`` (cost adjustment, metadata injection, refund +sending). +""" + +from __future__ import annotations + +import json +from collections.abc import AsyncGenerator, AsyncIterator +from typing import Any, Callable, NamedTuple, cast + +import litellm + +from ..core import get_logger +from ..core.exceptions import UpstreamError +from ..payment.models import Model + +logger = get_logger(__name__) + +# Anthropic-Messages-only fields that don't translate to OpenAI +# Chat Completions. ``litellm.drop_params`` only filters *known* +# unsupported params; these newer/extension fields get passed through +# verbatim and the upstream rejects them with a 400. Pop them here so the +# request reaches the upstream cleanly. +ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = ( + "thinking", + "cache_control", + "context_management", + "output_config", + "mcp_servers", + "service_tier", + "anthropic_version", + "anthropic_beta", +) + + +def coerce_litellm_payload(payload: object) -> dict: + """Convert a litellm event into a plain dict. + + Non-streaming responses come back as Anthropic-shaped pydantic models + or dicts. Streaming may yield raw bytes/str (SSE-encoded); those go + through ``events_from_chunk`` instead, not here. + """ + if isinstance(payload, dict): + return dict(payload) + if hasattr(payload, "model_dump"): + return cast(dict, payload.model_dump()) + raise TypeError(f"Cannot coerce {type(payload).__name__} to dict") + + +def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: + """Parse complete SSE event blocks out of a byte buffer. + + Returns (events, remaining_buffer). Events are JSON objects parsed from + one or more ``data:`` lines per block. Comments, blank lines, and + ``[DONE]`` sentinels are ignored. A trailing partial block is preserved + in remaining_buffer. + """ + events: list[dict] = [] + while True: + sep = buffer.find(b"\n\n") + if sep < 0: + sep_rn = buffer.find(b"\r\n\r\n") + if sep_rn < 0: + break + block = buffer[:sep_rn] + buffer = buffer[sep_rn + 4 :] + else: + block = buffer[:sep] + buffer = buffer[sep + 2 :] + + data_lines: list[str] = [] + for raw_line in block.replace(b"\r\n", b"\n").split(b"\n"): + line = raw_line.decode("utf-8", errors="replace") + if line.startswith(":"): + continue + if line.startswith("data:"): + data_lines.append(line[5:].lstrip()) + if not data_lines: + continue + payload = "\n".join(data_lines).strip() + if not payload or payload == "[DONE]": + continue + try: + obj = json.loads(payload) + except json.JSONDecodeError: + continue + if isinstance(obj, dict): + events.append(obj) + return events, buffer + + +def events_from_chunk( + chunk: object, sse_buffer: bytes +) -> tuple[list[dict], bytes]: + """Normalize a stream chunk into one or more event dicts. + + ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE + bytes in practice; some adapters yield strings or typed events. Handle + all three. + """ + if isinstance(chunk, (bytes, bytearray)): + sse_buffer += bytes(chunk) + events, sse_buffer = parse_sse_blocks(sse_buffer) + return events, sse_buffer + if isinstance(chunk, str): + sse_buffer += chunk.encode("utf-8") + events, sse_buffer = parse_sse_blocks(sse_buffer) + return events, sse_buffer + return [coerce_litellm_payload(chunk)], sse_buffer + + +async def aggregate_anthropic_events_to_message( + iterator: AsyncIterator[Any], +) -> dict: + """Drain an Anthropic-Messages event iterator into a single Message dict. + + Produces the shape ``litellm.anthropic.messages.acreate(stream=False)`` + would have returned. 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 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 = 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 + + +class AnnotatedEvent(NamedTuple): + """One Anthropic SSE event after model-rewrite + token-tally bookkeeping. + + ``sse_bytes`` is the wire-ready ``event:`` / ``data:`` block; the two + streaming paths in ``BaseUpstreamProvider`` consume ``sse_bytes`` plus + the tallies and only differ in whether they stream live or buffer + first. + """ + + event: dict + sse_bytes: bytes + input_tokens: int + output_tokens: int + model: str | None + + +def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent: + """Rewrite ``model`` fields and extract per-event token / model info. + + Mutates ``event`` in place when ``requested_model`` is set so the + upstream's true model name doesn't leak to the client. + """ + 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 + + in_tokens = 0 + out_tokens = 0 + model: str | None = None + + msg_for_meta = event.get("message") + if isinstance(msg_for_meta, dict): + if msg_for_meta.get("model"): + model = str(msg_for_meta["model"]) + usage = msg_for_meta.get("usage") + if isinstance(usage, dict): + in_tokens += int(usage.get("input_tokens") or 0) + out_tokens += int(usage.get("output_tokens") or 0) + + if isinstance(event.get("usage"), dict): + usage = event["usage"] + in_tokens += int(usage.get("input_tokens") or 0) + out_tokens += int(usage.get("output_tokens") or 0) + + event_type = str(event.get("type") or "") + payload = json.dumps(event) + if event_type: + sse_bytes = f"event: {event_type}\ndata: {payload}\n\n".encode() + else: + sse_bytes = f"data: {payload}\n\n".encode() + + return AnnotatedEvent(event, sse_bytes, in_tokens, out_tokens, model) + + +async def stream_annotated_events( + iterator: AsyncIterator[Any], + requested_model: str | None, +) -> AsyncGenerator[AnnotatedEvent, None]: + """Yield annotated, SSE-serialized events from a litellm stream. + + Both streaming paths in ``BaseUpstreamProvider`` consume this; the only + divergence between them — yield-as-you-go vs buffer-then-replay — stays + in the caller. + """ + sse_buffer = b"" + async for chunk in iterator: + events, sse_buffer = events_from_chunk(chunk, sse_buffer) + for event in events: + yield annotate_event(event, requested_model) + + +def compute_refund(amount: int, unit: str, cost_msats: int) -> int: + if unit == "msat": + return amount - cost_msats + if unit == "sat": + return amount - (cost_msats + 999) // 1000 + raise ValueError(f"Invalid unit: {unit}") + + +async def dispatch_anthropic_messages( + *, + request_body: bytes | None, + model_obj: Model, + base_url: str, + api_key: str, + provider_prefix: str, + transform_model_name: Callable[[str], str], + log_extra: dict[str, Any] | None = None, +) -> tuple[bool, Any, str | None]: + """Call ``litellm.anthropic.messages.acreate`` and return + ``(client_stream, result, requested_model)``. + + Shared by the bearer-key and x-cashu paths. Raises :class:`UpstreamError` + on bad input or upstream failure. + """ + 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` 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 + + dropped: dict[str, Any] = {} + for field in ANTHROPIC_ONLY_FIELDS: + if field in body: + dropped[field] = body.pop(field) + if dropped: + logger.debug( + "Dropped anthropic-only fields before litellm dispatch", + extra={"dropped_keys": sorted(dropped.keys())}, + ) + + # Convention: `model.id` is the canonical upstream model name; + # `forwarded_model_id` is the public alias the internal API exposes + # and echoes back to the client. + requested_model = ( + (model_obj.forwarded_model_id or model_obj.id) if model_obj else None + ) + upstream_model = transform_model_name(model_obj.id) + litellm_model = f"{provider_prefix}{upstream_model}" + + kwargs: dict = { + "model": litellm_model, + "api_base": base_url, + "api_key": api_key, + "stream": upstream_stream, + **body, + } + + logger.info( + "Dispatching /v1/messages via litellm", + extra={ + "model": litellm_model, + "resolved_provider": provider_prefix.rstrip("/"), + "client_stream": client_stream, + "upstream_stream": upstream_stream, + **(log_extra or {}), + }, + ) + + try: + result = await litellm.anthropic.messages.acreate(**kwargs) + except Exception as exc: + exc_message = getattr(exc, "message", None) or str(exc) or repr(exc) + exc_status = getattr(exc, "status_code", None) + exc_response = getattr(exc, "response", None) + response_text = None + if exc_response is not None: + try: + response_text = getattr(exc_response, "text", str(exc_response)) + except Exception: + response_text = "" + logger.error( + "litellm dispatch failed", + extra={ + "error": exc_message, + "error_type": type(exc).__name__, + "status_code": exc_status, + "llm_provider": getattr(exc, "llm_provider", None), + "body": getattr(exc, "body", None), + "response_text": response_text, + "model": litellm_model, + "api_base": base_url, + }, + ) + raise UpstreamError( + f"Upstream error via litellm: {exc_message}", + status_code=exc_status if isinstance(exc_status, int) else 502, + ) from exc + + 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 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 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