normalize response

This commit is contained in:
9qeklajc
2026-04-28 00:44:22 +02:00
parent 1d22155e05
commit 66421cd17f
4 changed files with 498 additions and 676 deletions
-288
View File
@@ -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: <type>\ndata: <json>\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.
-355
View File
@@ -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/<new_revision>_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
+167 -33
View File
@@ -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:
@@ -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
# ---------------------------------------------------------------------------