From 66421cd17f3e36b09e6a6580b38af27c5d439285 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Apr 2026 22:07:29 +0200 Subject: [PATCH] normalize response --- docs/plan-messages-to-chat-completions.md | 288 --------------- plans/other-mints-balance-tracking-plan.md | 355 ------------------- routstr/upstream/base.py | 200 +++++++++-- tests/unit/test_messages_litellm_dispatch.py | 331 +++++++++++++++++ 4 files changed, 498 insertions(+), 676 deletions(-) delete mode 100644 docs/plan-messages-to-chat-completions.md delete mode 100644 plans/other-mints-balance-tracking-plan.md diff --git a/docs/plan-messages-to-chat-completions.md b/docs/plan-messages-to-chat-completions.md deleted file mode 100644 index 91e07fa6..00000000 --- a/docs/plan-messages-to-chat-completions.md +++ /dev/null @@ -1,288 +0,0 @@ -# Plan: Route Anthropic `/v1/messages` requests to OpenAI-compatible upstreams via litellm - -## Problem - -Claude Code (and other Anthropic-SDK clients) issues requests against -`/v1/messages`. In `routstr-core`, every request is dispatched through -`BaseUpstreamProvider.forward_request` (or `forward_x_cashu_request`), -which sends the request body verbatim to the upstream's `/messages` -endpoint. - -Only two upstream providers natively accept `/v1/messages`: -- `AnthropicUpstreamProvider` -- `OpenRouterUpstreamProvider` (OpenRouter exposes an Anthropic-compat - endpoint at `/api/v1/messages`) - -All other providers expose only `/chat/completions` (OpenAI-compatible). -Today these requests fail with 404 / "endpoint not found" when a Claude -Code user routes them through `routstr-core` to an OpenAI-compatible -upstream. - -Affected providers: -`openai`, `groq`, `xai`, `fireworks`, `perplexity`, `gemini`, `ollama`, -`azure`, `ppqai`, `routstr`, `generic`. - -## Goal - -When the client sends `/v1/messages` and the resolved upstream **does -not** natively support that endpoint, `routstr-core` should: - -1. Translate the Anthropic-format request body into OpenAI Chat - Completions format. -2. Forward it to the upstream's `/chat/completions` endpoint. -3. Translate the response (streaming or non-streaming) back to Anthropic - Messages format before returning it to the client. - -Cost tracking, payment deduction, model rewriting, and provider -fallback must keep working unchanged. - -## Approach: use `litellm` as the translation engine - -After auditing the problem we determined the translator (request + -response + SSE event sequencing + tool_call buffering across deltas + -usage propagation + provider quirks like OpenAI's 64-char tool-name -limit) is non-trivial to maintain. - -[`litellm`](https://docs.litellm.ai/docs/anthropic_unified) ships -exactly this capability as an in-process Python SDK call: - -```python -import litellm -response = await litellm.anthropic.messages.acreate( - model="openai/gpt-4o-mini", # or groq/..., xai/..., gemini/..., etc. - api_base="https://api.groq.com/openai/v1", - api_key="...", - messages=[...], - max_tokens=1024, - stream=True, -) -``` - -It accepts the Anthropic `/v1/messages` body, calls any LiteLLM-known -provider, and returns an Anthropic-shaped response (or an -`AsyncIterator` of Anthropic-shaped chunks for `stream=True`). - -**No separate proxy process required.** litellm is just a Python library; -we `import litellm` and call the function in-process. The deployment -story is unchanged: still one `routstr-core` process. - -## Constraints - -- **Least changes / simplest design.** Reuse litellm; do not maintain our - own translator. -- Must not break existing `/v1/messages` flow for Anthropic and - OpenRouter (already working). -- Must respect `CLAUDE.md`: clean `uv run ruff check . --fix` and - `uv run mypy .` after every edit; affected unit tests pass. - -## Design - -### 1. Dependency - -Add `litellm` to `pyproject.toml` `[project.dependencies]`. - -### 2. Provider opt-in: `supports_anthropic_messages` - -On `BaseUpstreamProvider` add a class attribute, default `False`: - -```python -supports_anthropic_messages: bool = False -``` - -Override on the two providers that natively serve `/v1/messages`: -- `routstr/upstream/anthropic.py` → `supports_anthropic_messages = True` -- `routstr/upstream/openrouter.py` → `supports_anthropic_messages = True` - -### 3. Provider → litellm prefix mapping: `litellm_provider_prefix` - -On `BaseUpstreamProvider` add a class attribute, default `"openai/"` -(safe default — most non-listed providers are OpenAI-compatible and can -be reached by passing `api_base`): - -```python -litellm_provider_prefix: str = "openai/" -``` - -Per-provider overrides where litellm has a native provider: -- `groq.py` → `"groq/"` -- `xai.py` → `"xai/"` -- `fireworks.py` → `"fireworks_ai/"` -- `perplexity.py` → `"perplexity/"` -- `gemini.py` → `"gemini/"` -- `ollama.py` → `"ollama/"` -- `azure.py` → `"azure/"` -- `openai.py`, `ppqai.py`, `routstr.py`, `generic.py` → keep default `"openai/"` - -litellm derives auth from `api_base` + `api_key` we pass, so -"openai/" + custom `api_base` works for any OpenAI-compatible upstream. - -### 4. New helper on `BaseUpstreamProvider`: `_forward_messages_via_litellm` - -Single private method that owns the full litellm path: build kwargs, -call `litellm.anthropic.messages.acreate`, run cost tracking, return -`Response` or `StreamingResponse`. Pure addition — does not touch -existing chat/completions or messages paths. - -Signature: - -```python -async def _forward_messages_via_litellm( - self, - request_body: bytes | None, - key_or_payment: ApiKey | XCashuPaymentContext, - session: AsyncSession | None, - max_cost_for_model: int, - model_obj: Model, - *, - request_id: str | None = None, - is_x_cashu: bool = False, - mint: str | None = None, - payment_token_hash: str | None = None, -) -> Response | StreamingResponse: - ... -``` - -Behaviour: -1. Parse `request_body` as JSON. Pop `model` (we override with our - transformed id). Read `stream` flag. -2. Build litellm kwargs: - - `model = f"{self.litellm_provider_prefix}{self.transform_model_name(model_obj.id)}"` - - `api_key = self.api_key` - - `api_base = self.base_url` - - splat the rest of the body (messages, max_tokens, system, tools, - tool_choice, temperature, top_p, top_k, stop_sequences, metadata, - thinking, stream). -3. `result = await litellm.anthropic.messages.acreate(**kwargs)`. -4. Non-streaming branch (`stream=False`): `result` is an Anthropic - `AnthropicMessagesResponse` dict-shaped object. Wrap into the same - shape `handle_non_streaming_messages_completion` produces: - - Run `adjust_payment_for_tokens(key, response_dict, session, max_cost_for_model)` - for the bearer-key path, or compute x-cashu refund for the x-cashu - path. - - Call `inject_cost_metadata(response_dict, cost_data, key)` (bearer - path) — for x-cashu, mirror `handle_x_cashu_chat_completion`'s - `X-Cashu` refund header logic. - - Return `Response(content=json.dumps(response_dict).encode(), media_type="application/json", ...)`. -5. Streaming branch (`stream=True`): `result` is an - `AsyncIterator` of Anthropic event dicts. We: - - Wrap it in an async generator that: - - For each event, serialize as SSE - (`event: \ndata: \n\n`). - - Tap `message_start` for `usage.input_tokens`, accumulate - `output_tokens` from `message_delta` `usage.output_tokens`. - - On stream end, run cost reconciliation - (`adjust_payment_for_tokens` for bearer; refund/lock for - x-cashu) using the captured usage. - - Return `StreamingResponse(generator(), media_type="text/event-stream")`. - -Reusing `adjust_payment_for_tokens` and `inject_cost_metadata` keeps -cost calculation identical between native and translated paths — they -already operate on Anthropic Messages response shape. - -### 5. Wire into `forward_request` (base.py:1403) - -At the start of the `if path.endswith("messages"):` branch, before the -existing httpx call, add: - -```python -if path.endswith("messages") and not self.supports_anthropic_messages: - return await self._forward_messages_via_litellm( - request_body=request_body, - key_or_payment=key, - session=session, - max_cost_for_model=max_cost_for_model, - model_obj=model_obj, - request_id=getattr(request.state, "request_id", None), - ) -``` - -Critical predicate ordering: this branch must be reached **only** for -`/messages` (not `/messages/count_tokens`). The current code already -checks `messages/count_tokens` separately, but our shortcut goes -**before** the httpx call so we must explicitly exclude count_tokens -ourselves, e.g. `path.endswith("messages") and not path.endswith("count_tokens")`. - -### 6. Wire into `forward_x_cashu_request` (base.py:2583) - -Same shortcut at the same logical location. We pass -`is_x_cashu=True` plus `mint` and `payment_token_hash` so the helper -takes the x-cashu refund branch. - -### 7. count_tokens is out of scope - -`/messages/count_tokens` requests against non-supporting providers stay -unchanged — they get the existing 404 response from the upstream. Same -behaviour as today, no regression. - -## Files touched - -- **Edit**: `pyproject.toml` (add `litellm`) -- **Edit**: `routstr/upstream/base.py` - - Add `supports_anthropic_messages: bool = False` - - Add `litellm_provider_prefix: str = "openai/"` - - Add `_forward_messages_via_litellm` method - - Branch into it from `forward_request` and `forward_x_cashu_request` -- **Edit**: `routstr/upstream/anthropic.py` (1 line — flag) -- **Edit**: `routstr/upstream/openrouter.py` (1 line — flag) -- **Edit**: `routstr/upstream/{groq,xai,fireworks,perplexity,gemini,ollama,azure}.py` - (1 line each — `litellm_provider_prefix`) -- **New**: `tests/unit/test_messages_litellm_dispatch.py` - -No new translator module, no per-provider override of forward paths, -no SSE rewriting code — it all lives inside litellm. - -## Tests - -`tests/unit/test_messages_litellm_dispatch.py`: -- Bearer-key path: dispatches to litellm when `supports_anthropic_messages=False`, - bypasses litellm when `True`. -- Non-streaming response: usage extracted, `adjust_payment_for_tokens` - called, cost metadata injected. -- Streaming response: usage extracted from `message_delta` events, - cost reconciled at end. -- x-cashu path: refund header populated. - -We mock `litellm.anthropic.messages.acreate` so tests do not require -network or upstream API keys. - -## Out of scope - -- `/v1/messages/count_tokens` translation — left to upstream to 404 as - today. -- Per-provider tuning of litellm options - (e.g. `litellm.use_chat_completions_url_for_anthropic_messages`). - Sane defaults; revisit only if a specific upstream misbehaves. - -## Verification - -After implementation, per `CLAUDE.md`: - -```bash -uv run ruff check . --fix -uv run mypy . -uv run pytest tests/unit/test_messages_litellm_dispatch.py -uv run pytest tests/unit # full unit suite still green -``` - -## Spike before merge - -Before opening the PR, run a manual one-shot: - -```python -import asyncio, litellm -async def main(): - r = await litellm.anthropic.messages.acreate( - model="groq/llama-3.3-70b-versatile", - api_base="https://api.groq.com/openai/v1", - api_key=os.environ["GROQ_API_KEY"], - messages=[{"role": "user", "content": "say hi"}], - max_tokens=64, - stream=True, - ) - async for chunk in r: - print(chunk) -asyncio.run(main()) -``` - -Confirms streaming, usage propagation and tool-format on a real -OpenAI-compatible upstream before we trust it in production. diff --git a/plans/other-mints-balance-tracking-plan.md b/plans/other-mints-balance-tracking-plan.md deleted file mode 100644 index b688b0ba..00000000 --- a/plans/other-mints-balance-tracking-plan.md +++ /dev/null @@ -1,355 +0,0 @@ -# Plan: Track unsupported incoming mints in `other_mints` and include them in balances - -## Goal - -When a Cashu token arrives from a mint that is **not** in `settings.cashu_mints`, we currently create/use a wallet for that mint and swap value into the primary mint. Any change/surplus left behind after `melt()` remains in the foreign `token_wallet`, but that balance is not surfaced by `fetch_all_balances()` because it only iterates over configured mints. - -This plan adds persistent tracking for those foreign mints in a new database table called `other_mints`, and updates balance reporting to include them. - ---- - -## Current behavior - -### Incoming unsupported mint flow - -In `routstr/wallet.py`: - -- `recieve_token()` deserializes the token -- if `token_obj.mint not in settings.cashu_mints`, it calls `swap_to_primary_mint(token_obj, wallet)` -- `swap_to_primary_mint()` calls `token_wallet.melt(...)` to pay the primary mint invoice - -### Important detail: change is retained, not discarded - -The underlying Cashu wallet library keeps any melt change: - -- `Wallet.melt()` constructs blank outputs for change -- when the melt succeeds, returned change is reconstructed into proofs -- those proofs are appended to `self.proofs` and stored in the wallet DB - -So surplus from unsupported mints is **not discarded**, but it may become invisible operationally. - -### Visibility problem - -`fetch_all_balances()` currently only loops over: - -- `settings.cashu_mints` -- units `sat` and `msat` - -This means balances left on unsupported mints are not shown in admin balance reporting. - ---- - -## Proposed design - -## 1. Add a new DB table: `other_mints` - -Add a small table in `routstr/core/db.py` to persist unsupported mints we have seen in incoming tokens. - -Suggested schema: - -- `mint_url: str` primary key -- `created_at: int` -- `last_seen_at: int` - -Minimal model: - -```python -class OtherMint(SQLModel, table=True): - __tablename__ = "other_mints" - - mint_url: str = Field(primary_key=True) - created_at: int = Field(default_factory=lambda: int(time.time())) - last_seen_at: int = Field(default_factory=lambda: int(time.time())) -``` - -Why minimal: - -- the only required function is mint discovery/tracking -- unit handling can remain dynamic via existing balance queries over `sat` and `msat` - ---- - -## 2. Add DB helpers for `other_mints` - -In `routstr/core/db.py`, add helper functions: - -### `register_other_mint(mint_url: str) -> None` - -Behavior: - -- if the mint is not present, insert it -- if it already exists, update `last_seen_at` - -### `list_other_mints(session) -> list[str]` - -Behavior: - -- return all tracked unsupported mint URLs - -Optional later: - -- `delete_other_mint(...)` -- admin cleanup helpers - ---- - -## 3. Register unsupported mints during token receipt - -Update `recieve_token()` in `routstr/wallet.py`. - -Current logic: - -```python -if token_obj.mint not in settings.cashu_mints: - return await swap_to_primary_mint(token_obj, wallet) -``` - -Planned logic: - -```python -if token_obj.mint not in settings.cashu_mints: - await db.register_other_mint(token_obj.mint) - return await swap_to_primary_mint(token_obj, wallet) -``` - -Why here: - -- this is the earliest reliable point where we know the mint came in via an actual token -- this is exactly the path that can leave foreign-mint change behind -- it avoids needing to infer unsupported mints later from wallet internals - ---- - -## 4. Update `fetch_all_balances()` to include `other_mints` - -Current behavior only includes configured mints. - -Planned behavior: - -- load tracked unsupported mints from DB -- combine them with `settings.cashu_mints` -- dedupe while preserving order -- fetch balances for all tracked mints across requested units - -Conceptual flow: - -```python -tracked_mints = dedupe(settings.cashu_mints + other_mints_from_db) -``` - -Then existing per-mint/per-unit balance logic can remain mostly unchanged. - -This ensures that retained change on unsupported mints becomes visible in admin balance reporting. - ---- - -## 5. Add a balance source marker - -Extend `BalanceDetail` in `routstr/wallet.py` to identify whether a balance row comes from a configured mint or an `other_mints` entry. - -Suggested field: - -- `source: str` with values: - - `"configured"` - - `"other"` - -Updated shape: - -```python -class BalanceDetail(TypedDict, total=False): - mint_url: str - unit: str - source: str - wallet_balance: int - user_balance: int - owner_balance: int - error: str -``` - -Why this helps: - -- admin can distinguish normal configured wallet balances from foreign/unsupported balances -- avoids confusion if unexpected mint URLs show up in the balances API/UI - ---- - -## 6. Admin/API impact - -Backend impact is minimal because `/admin/api/balances` already returns `fetch_all_balances()` output. - -Effects: - -- supported mints continue to show as before -- tracked unsupported mints will also appear -- UI can optionally display the new `source` field - -No API contract break is expected if the frontend ignores unknown fields. - ---- - -## 7. Payout behavior: do not change in phase 1 - -`periodic_payout()` currently only iterates over `settings.cashu_mints`. - -Recommendation for this change: - -- **do not** expand `periodic_payout()` to include `other_mints` yet -- only improve visibility through balance reporting - -Reason: - -- automatic payout from unsupported/foreign mints may be operationally undesirable -- visibility should come first, automation second - -Possible future phase: - -- add optional sweeping/payout support for `other_mints` -- or provide an admin-triggered withdrawal/sweep flow - ---- - -## 8. Logging improvements (optional) - -Optional follow-up improvement in `swap_to_primary_mint()`: - -- capture the return value from `token_wallet.melt(...)` -- if feasible, log any reported change amount -- otherwise, rely on wallet balance reporting to surface residual amounts - -This is useful but not required for the first implementation. - ---- - -## Files to change - -### `routstr/core/db.py` - -Add: - -- `OtherMint` SQLModel -- `register_other_mint()` -- `list_other_mints()` - -### `migrations/versions/_add_other_mints_table.py` - -Create migration to add the `other_mints` table. - -### `routstr/wallet.py` - -Update: - -- `recieve_token()` to register unsupported mints -- `BalanceDetail` to include `source` -- `fetch_all_balances()` to include both configured and tracked unsupported mints - -### `routstr/core/admin.py` - -Likely no backend changes required unless a dedicated `other_mints` API is desired. - ---- - -## Behavior rules - -### Register a mint when - -- an incoming token is processed -- the token mint is not in `settings.cashu_mints` - -### Do not remove automatically when - -- balance reaches zero - -Reason: - -- historical visibility is useful -- avoids flapping entries in the admin balance list -- mint may receive additional unsupported tokens later - -Potential future enhancement: - -- admin endpoint to prune zero-balance `other_mints` - ---- - -## Edge cases - -### A mint later becomes configured - -If a mint in `other_mints` is later added to `settings.cashu_mints`: - -- deduplication prevents duplicate balance rows -- `source` should resolve to `configured` - -### Unsupported mint with zero balance - -A tracked unsupported mint may show zero balances. - -Initial recommendation: - -- allow it to appear -- consider later filtering zero-balance `other` rows if the UI becomes noisy - -### Units - -Balance fetching can continue to query both `sat` and `msat` for each tracked mint. - -If a mint has no proofs in one unit, current error/zero handling can continue to apply. - ---- - -## Test plan - -### DB tests - -- registering a new unsupported mint inserts a row -- registering the same mint again updates `last_seen_at` without duplication -- listing other mints returns expected mint URLs - -### Wallet tests - -#### `recieve_token()` - -- when mint is unsupported, `db.register_other_mint()` is called before swap -- when mint is configured, `db.register_other_mint()` is not called - -#### `fetch_all_balances()` - -- includes configured mints -- includes `other_mints` from DB -- dedupes if a mint exists in both configured and other lists -- sets `source` correctly - -### Regression tests - -- existing trusted mint balance reporting remains unchanged -- `/admin/api/balances` continues to work - ---- - -## Recommended implementation order - -1. Add `OtherMint` model to `routstr/core/db.py` -2. Add Alembic migration for `other_mints` -3. Add `register_other_mint()` and `list_other_mints()` helpers -4. Update `recieve_token()` to register unsupported mints -5. Update `fetch_all_balances()` to union configured + tracked other mints -6. Add `source` to `BalanceDetail` -7. Add/adjust tests - ---- - -## Summary - -This change solves an operational visibility problem: - -- unsupported incoming mints can leave retained change in foreign wallets -- those funds are currently preserved but not surfaced in balance reporting -- introducing `other_mints` makes those mints discoverable and auditable -- expanding `fetch_all_balances()` ensures their balances are visible in admin tooling - -Recommended scope for the first pass: - -- track unsupported mints in DB -- include them in balance reporting -- mark them as `source="other"` -- do not yet change payout/sweeping behavior diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 6872efd0..6d3933f8 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -6,6 +6,7 @@ import json import re import traceback import uuid +import warnings from collections.abc import AsyncGenerator, AsyncIterator from typing import Any, Mapping, cast @@ -44,6 +45,18 @@ from ..wallet import recieve_token, send_token logger = get_logger(__name__) +# litellm response models declare nested fields as pydantic types +# (e.g. `usage: ResponseAPIUsage`) but populate them with plain dicts at +# runtime. Whenever such a model is dumped — by us, by litellm, or by a +# downstream client — pydantic-core emits a benign UserWarning that floods +# the logs once per request. We re-normalize `usage` ourselves, so the +# warning is irrelevant; suppress it process-wide. +warnings.filterwarnings( + "ignore", + message="Pydantic serializer warnings:", + category=UserWarning, +) + class TopupData(BaseModel): """Universal top-up data schema for Lightning Network invoices.""" @@ -1408,13 +1421,125 @@ class BaseUpstreamProvider: if isinstance(payload, dict): return dict(payload) if hasattr(payload, "model_dump"): - return payload.model_dump() # type: ignore[no-any-return] + # Suppress at source: a downstream `simplefilter("always")` + # would otherwise reinstate the benign serializer warning that + # litellm's typed-but-dict-valued fields trigger. + with warnings.catch_warnings(): + warnings.filterwarnings( + "ignore", + message="Pydantic serializer warnings:", + category=UserWarning, + ) + return payload.model_dump(serialize_as_any=True) # 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") + @staticmethod + def _alias_usage(usage: Mapping[str, Any]) -> dict: + """Backfill Anthropic-shaped keys (input_tokens/output_tokens) + from OpenAI-shaped equivalents (prompt_tokens/completion_tokens). + """ + aliased = dict(usage) + for target, source in ( + ("input_tokens", "prompt_tokens"), + ("output_tokens", "completion_tokens"), + ): + if target not in aliased and source in aliased: + aliased[target] = aliased[source] + return aliased + + @staticmethod + def _normalize_litellm_payload(payload: object) -> dict: + """Coerce a litellm response/event into a dict with Anthropic-shaped + usage keys at the top level and inside any nested ``message.usage``. + """ + payload_dict = BaseUpstreamProvider._coerce_litellm_payload(payload) + + usage = payload_dict.get("usage") + if isinstance(usage, Mapping): + payload_dict = { + **payload_dict, + "usage": BaseUpstreamProvider._alias_usage(usage), + } + + message = payload_dict.get("message") + if isinstance(message, Mapping): + msg_usage = message.get("usage") + if isinstance(msg_usage, Mapping): + payload_dict = { + **payload_dict, + "message": { + **message, + "usage": BaseUpstreamProvider._alias_usage(msg_usage), + }, + } + return payload_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. Trailing partial block is preserved. + """ + events: list[dict] = [] + while True: + sep = buffer.find(b"\n\n") + if sep < 0: + # Tolerate \r\n\r\n separators too. + 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( + self, chunk: object, sse_buffer: bytes + ) -> tuple[list[dict], bytes]: + """Normalize an upstream stream chunk into one or more event dicts. + + Handles three cases: + - dict / pydantic model / object → single coerced event + - bytes / bytearray → one or more SSE blocks (buffered) + - str → encoded then handled as bytes + """ + 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 + async def _forward_messages_via_litellm( self, request_body: bytes | None, @@ -1492,7 +1617,7 @@ class BaseUpstreamProvider: requested_model, ) - response_json = self._coerce_litellm_payload(result) + response_json = self._normalize_litellm_payload(result) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -1554,43 +1679,52 @@ class BaseUpstreamProvider: usage_finalized = True return None + sse_buffer = b"" try: async for chunk in iterator: - event = self._coerce_litellm_payload(chunk) - event_type = str(event.get("type") or "") + events, sse_buffer = self._events_from_chunk( + chunk, sse_buffer + ) + for raw_event in events: + event = self._normalize_litellm_payload(raw_event) + 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 + 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"]) + 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) + 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() + 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: diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index fa964286..8b92389f 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -111,6 +111,156 @@ def test_coerce_litellm_payload_handles_pydantic_v2() -> None: assert out == {"x": 42} +def test_parse_sse_blocks_extracts_full_events() -> None: + buffer = ( + b"event: message_start\n" + b'data: {"type":"message_start","message":{"id":"m1"}}\n\n' + b"event: message_stop\n" + b'data: {"type":"message_stop"}\n\n' + ) + events, remaining = BaseUpstreamProvider._parse_sse_blocks(buffer) + assert remaining == b"" + assert events == [ + {"type": "message_start", "message": {"id": "m1"}}, + {"type": "message_stop"}, + ] + + +def test_parse_sse_blocks_preserves_partial_trailing_block() -> None: + buffer = b'data: {"type":"a"}\n\nevent: b\ndata: {"type":"b"}' + events, remaining = BaseUpstreamProvider._parse_sse_blocks(buffer) + assert events == [{"type": "a"}] + assert remaining == b'event: b\ndata: {"type":"b"}' + + +def test_parse_sse_blocks_skips_done_and_comments() -> None: + buffer = b": ping\n\ndata: [DONE]\n\ndata: {\"type\":\"x\"}\n\n" + events, remaining = BaseUpstreamProvider._parse_sse_blocks(buffer) + assert remaining == b"" + assert events == [{"type": "x"}] + + +def test_events_from_chunk_handles_bytes_chunks() -> None: + provider = _make_provider() + chunk = ( + b"event: a\ndata: {\"type\":\"a\"}\n\n" + b"event: b\ndata: {\"type\":\"b\"}" + ) + events, buf = provider._events_from_chunk(chunk, b"") + assert events == [{"type": "a"}] + # second event still partial because no trailing \n\n + assert buf == b'event: b\ndata: {"type":"b"}' + + # Feed remainder of the second event + events, buf = provider._events_from_chunk(b"\n\n", buf) + assert events == [{"type": "b"}] + assert buf == b"" + + +def test_events_from_chunk_handles_str_chunks() -> None: + provider = _make_provider() + events, buf = provider._events_from_chunk( + 'event: a\ndata: {"type":"a"}\n\n', b"" + ) + assert events == [{"type": "a"}] + assert buf == b"" + + +# --------------------------------------------------------------------------- +# Pydantic serializer warning + usage normalization +# --------------------------------------------------------------------------- + + +def test_coerce_litellm_payload_silences_pydantic_serializer_warning() -> None: + """litellm response models emit a benign UserWarning when their nested + pydantic field (e.g. usage = ResponseAPIUsage) holds a plain dict. + _coerce_litellm_payload must silence it locally.""" + import warnings as _warnings + + obj = MagicMock() + + def _emit_warning(*args: Any, **kwargs: Any) -> dict: + _warnings.warn( + "Pydantic serializer warnings:\n Expected `ResponseAPIUsage`", + UserWarning, + stacklevel=2, + ) + return {"id": "x", "usage": {"completion_tokens": 5}} + + obj.model_dump.side_effect = _emit_warning + + with _warnings.catch_warnings(record=True) as caught: + _warnings.simplefilter("always") + out = BaseUpstreamProvider._coerce_litellm_payload(obj) + + assert out == {"id": "x", "usage": {"completion_tokens": 5}} + serializer_warnings = [ + w for w in caught if "Pydantic serializer warnings" in str(w.message) + ] + assert serializer_warnings == [], ( + "Pydantic serializer warning should be suppressed at source" + ) + + +def test_normalize_litellm_payload_maps_openai_keys_to_anthropic() -> None: + event = { + "type": "message_delta", + "usage": { + "prompt_tokens": 11, + "completion_tokens": 22, + "total_tokens": 33, + }, + "message": { + "model": "x", + "usage": {"prompt_tokens": 7, "completion_tokens": 13}, + }, + } + out = BaseUpstreamProvider._normalize_litellm_payload(event) + + # Top-level usage gets canonical Anthropic keys mirrored in + assert out["usage"]["input_tokens"] == 11 + assert out["usage"]["output_tokens"] == 22 + # Original keys preserved (non-destructive) + assert out["usage"]["prompt_tokens"] == 11 + assert out["usage"]["completion_tokens"] == 22 + + # Nested message.usage normalized too + assert out["message"]["usage"]["input_tokens"] == 7 + assert out["message"]["usage"]["output_tokens"] == 13 + + # Original event must not be mutated + assert "input_tokens" not in event["usage"] + + +def test_normalize_litellm_payload_keeps_anthropic_keys_intact() -> None: + event = { + "type": "message_delta", + "usage": {"input_tokens": 4, "output_tokens": 9}, + } + out = BaseUpstreamProvider._normalize_litellm_payload(event) + assert out["usage"] == {"input_tokens": 4, "output_tokens": 9} + + +def test_normalize_litellm_payload_aliases_dumped_usage() -> None: + """A litellm result whose model_dump emits OpenAI-style usage keys + should be normalized to Anthropic-style keys for downstream cost + reconciliation, with the original keys preserved alongside.""" + result = MagicMock() + result.model_dump.return_value = { + "id": "abc", + "model": "openai/gpt-4o-mini", + "usage": {"prompt_tokens": 12, "completion_tokens": 34}, + } + + out = BaseUpstreamProvider._normalize_litellm_payload(result) + + assert out["model"] == "openai/gpt-4o-mini" + assert out["usage"]["input_tokens"] == 12 + assert out["usage"]["output_tokens"] == 34 + # OpenAI keys preserved alongside the canonical mapping + assert out["usage"]["prompt_tokens"] == 12 + + # --------------------------------------------------------------------------- # Provider gating # --------------------------------------------------------------------------- @@ -319,6 +469,187 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: assert combined["model"] == "openai/gpt-4o-mini" +@pytest.mark.asyncio +async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: + """Regression: litellm sometimes yields already-SSE-encoded bytes. + + Previously this raised TypeError("Cannot coerce bytes to dict"). The + stream loop must parse SSE blocks (even split across chunks) and still + perform cost reconciliation. + """ + provider = _make_provider() + key = _make_key() + model = _make_model() + session = _make_session() + body = _anthropic_request_body(stream=True) + + async def fake_byte_chunks() -> AsyncIterator[bytes]: + # Whole event in one chunk + yield ( + b"event: message_start\n" + b'data: {"type":"message_start","message":{"id":"m1",' + b'"model":"openai/gpt-4o-mini",' + b'"usage":{"input_tokens":3,"output_tokens":0}}}\n\n' + ) + # Event split across two chunks (boundary inside the data line) + yield b'event: message_delta\ndata: {"type":"message_delta",' + yield b'"delta":{},"usage":{"output_tokens":4}}\n\n' + # SSE comment + DONE sentinel must be ignored + yield b": keepalive\n\ndata: [DONE]\n\n" + yield b"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n" + + fake_cost = {"total_msats": 999, "total_usd": 0.0001} + captured: dict[str, Any] = {} + + async def fake_adjust( + fresh_key: Any, combined_data: Any, sess: Any, max_cost: int + ) -> dict: + captured["combined_data"] = combined_data + return fake_cost + + fake_session = MagicMock() + fake_session.get = AsyncMock(return_value=key) + + class FakeSessionCtx: + async def __aenter__(self) -> Any: + return fake_session + + async def __aexit__(self, *args: Any) -> None: + return None + + with ( + patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(return_value=fake_byte_chunks()), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(side_effect=fake_adjust), + ), + patch( + "routstr.upstream.base.create_session", + new=lambda: FakeSessionCtx(), + ), + ): + result = await provider._forward_messages_via_litellm( + request_body=body, + key=key, + session=session, + max_cost_for_model=10_000, + model_obj=model, + ) + + assert isinstance(result, StreamingResponse) + emitted: list[bytes] = [] + async for chunk in result.body_iterator: + if isinstance(chunk, bytes): + emitted.append(chunk) + elif isinstance(chunk, memoryview): + emitted.append(bytes(chunk)) + else: + emitted.append(chunk.encode()) + + joined = b"".join(emitted).decode() + assert "event: message_start" in joined + assert "event: message_delta" in joined + assert "event: message_stop" in joined + assert "event: cost" in joined + # [DONE] sentinel and SSE comments must NOT be re-emitted as data + assert "[DONE]" not in joined + + combined = captured["combined_data"] + assert combined["usage"]["input_tokens"] == 3 + assert combined["usage"]["output_tokens"] == 4 + assert combined["model"] == "openai/gpt-4o-mini" + + +@pytest.mark.asyncio +async def test_streaming_normalizes_openai_usage_keys_for_cost() -> None: + """Regression: when litellm yields chunks with OpenAI-style usage keys + (prompt_tokens/completion_tokens), the stream loop must map them to + Anthropic input_tokens/output_tokens for cost reconciliation.""" + provider = _make_provider() + key = _make_key() + model = _make_model() + session = _make_session() + body = _anthropic_request_body(stream=True) + + async def fake_chunks() -> AsyncIterator[dict]: + yield { + "type": "message_start", + "message": { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "openai/gpt-4o-mini", + "content": [], + # OpenAI-style usage on the message + "usage": {"prompt_tokens": 9, "completion_tokens": 0}, + }, + } + yield { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "hi"}, + } + yield { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + # OpenAI-style usage on the delta + "usage": {"prompt_tokens": 0, "completion_tokens": 17}, + } + yield {"type": "message_stop"} + + fake_cost = {"total_msats": 100, "total_usd": 0.0001} + captured: dict[str, Any] = {} + + async def fake_adjust( + fresh_key: Any, combined_data: Any, sess: Any, max_cost: int + ) -> dict: + captured["combined_data"] = combined_data + return fake_cost + + fake_session = MagicMock() + fake_session.get = AsyncMock(return_value=key) + + class FakeSessionCtx: + async def __aenter__(self) -> Any: + return fake_session + + async def __aexit__(self, *args: Any) -> None: + return None + + with ( + patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(return_value=fake_chunks()), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(side_effect=fake_adjust), + ), + patch( + "routstr.upstream.base.create_session", + new=lambda: FakeSessionCtx(), + ), + ): + result = await provider._forward_messages_via_litellm( + request_body=body, + key=key, + session=session, + max_cost_for_model=10_000, + model_obj=model, + ) + assert isinstance(result, StreamingResponse) + async for _ in result.body_iterator: + pass + + combined = captured["combined_data"] + # Token counts come from the OpenAI-style fields after normalization + assert combined["usage"]["input_tokens"] == 9 + assert combined["usage"]["output_tokens"] == 17 + + # --------------------------------------------------------------------------- # forward_request gating # ---------------------------------------------------------------------------