mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
normalize response
This commit is contained in:
@@ -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.
|
||||
@@ -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
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user