This commit is contained in:
9qeklajc
2026-05-03 15:35:49 +02:00
parent 37c2bea93d
commit 145777ffd0
4 changed files with 569 additions and 447 deletions
+5
View File
@@ -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()
+57 -447
View File
@@ -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 = "<unreadable>"
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:
+49
View File
@@ -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
+458
View File
@@ -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 = "<unreadable>"
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