update completion path

This commit is contained in:
9qeklajc
2026-09-02 22:27:17 +02:00
parent b1ffd7c018
commit cdafdcbbd3
6 changed files with 512 additions and 38 deletions
+11 -6
View File
@@ -287,16 +287,21 @@ def _sum_string_chars(node: Any) -> int:
return 0
def _count_prompt_token_ids(node: Any) -> int:
if isinstance(node, int) and not isinstance(node, bool):
return 1
if isinstance(node, list):
return sum(_count_prompt_token_ids(item) for item in node)
return 0
def estimate_prompt_tokens(body: dict) -> int:
"""Conservatively estimate prompt tokens for the whole provider-bound body.
Unlike ``estimate_tokens`` (message text only), this walks every field, so
prompt weight hidden in tool schemas, tool-call arguments, ``system``, or
any field forwarded in future cannot escape the reservation estimate. It
over-estimates rather than under-estimates: the result only shrinks a
discount against a reservation that settlement later refunds.
Every string counts, as do token IDs in legacy ``prompt`` arrays, so no
forwarded field can hide prompt weight and shrink its reservation.
"""
return _sum_string_chars(body) // 3
return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt"))
def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
+55 -28
View File
@@ -248,6 +248,13 @@ def _is_json_content_type(content_type: str | None) -> bool:
return main.startswith("application/") and main.endswith("+json")
def _openai_completion_path(path: str) -> str | None:
canonical = "/" + path.rstrip("/")
if canonical.endswith("/chat/completions"):
return "chat/completions"
return "completions" if canonical.endswith("/completions") else None
class TopupData(BaseModel):
"""Universal top-up data schema for Lightning Network invoices."""
@@ -653,17 +660,20 @@ class BaseUpstreamProvider:
return "openrouter.ai" in (self.base_url or "")
def prepare_request_body(
self, body: bytes | None, model_obj: Model
self,
body: bytes | None,
model_obj: Model,
include_stream_usage: bool = False,
) -> bytes | None:
"""Transform request body for provider-specific requirements.
Automatically transforms model names and, for streaming chat
completions, opts the upstream into emitting per-chunk ``usage``
so cost tracking can read real token counts instead of falling
back to ``MaxCostData``.
Automatically transforms model names and opts streaming OpenAI
completion endpoints into emitting per-chunk ``usage`` so cost
tracking can read real token counts.
Args:
body: Original request body bytes
include_stream_usage: Opt a streaming completion into usage chunks
Returns:
Transformed request body bytes
@@ -706,14 +716,8 @@ class BaseUpstreamProvider:
# OpenAI-compatible streaming responses omit ``usage`` unless the
# request sets ``stream_options.include_usage = true``. Without it
# we can't reconcile token counts at end of stream and must use
# the local request/response estimator. Discriminate
# chat-completions-shaped requests by the ``messages`` field so we
# don't poke unrelated endpoints.
if (
data.get("stream") is True
and "messages" in data
and isinstance(data.get("messages"), list)
):
# the local request/response estimator.
if data.get("stream") is True and include_stream_usage:
existing = data.get("stream_options")
merged = dict(existing) if isinstance(existing, dict) else {}
if merged.get("include_usage") is not True:
@@ -1022,6 +1026,7 @@ class BaseUpstreamProvider:
reservation_snapshot: ReservationSnapshot | None = None,
client: httpx.AsyncClient | None = None,
request_body: bytes | None = None,
legacy_completion: bool = False,
) -> StreamingResponse:
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
@@ -1060,6 +1065,7 @@ class BaseUpstreamProvider:
last_model_seen: str | None = None
usage_chunk_data: dict | None = None
done_seen: bool = False
stream_id: str | None = None
async def finalize_db_only() -> None:
nonlocal usage_finalized
@@ -1118,7 +1124,7 @@ class BaseUpstreamProvider:
* ``[DONE]`` is swallowed so it can be re-emitted exactly once at
end of stream.
"""
nonlocal last_model_seen, usage_chunk_data, done_seen
nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id
event = raw_event.strip(b"\r\n")
if not event:
@@ -1171,9 +1177,12 @@ class BaseUpstreamProvider:
or not isinstance(obj["id"], str)
or obj["id"] == "existing-id"
):
if not hasattr(self, "_current_stream_id"):
self._current_stream_id = f"chatcmpl-{uuid.uuid4()}"
obj["id"] = self._current_stream_id
if stream_id is None:
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
stream_id = f"{id_prefix}-{uuid.uuid4()}"
obj["id"] = stream_id
else:
stream_id = obj["id"]
if isinstance(obj.get("usage"), dict):
# Capture usage for end-of-stream cost reconciliation.
# Some models (e.g. Gemini thinking models over the
@@ -1285,11 +1294,14 @@ class BaseUpstreamProvider:
raise
if usage_chunk_data is None:
if not hasattr(self, "_current_stream_id"):
self._current_stream_id = f"chatcmpl-{uuid.uuid4()}"
if stream_id is None:
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
stream_id = f"{id_prefix}-{uuid.uuid4()}"
usage_chunk_data = {
"id": self._current_stream_id,
"object": "chat.completion.chunk",
"id": stream_id,
"object": "text_completion"
if legacy_completion
else "chat.completion.chunk",
"model": last_model_seen or "unknown",
"choices": [],
"usage": {
@@ -1361,6 +1373,7 @@ class BaseUpstreamProvider:
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
legacy_completion: bool = False,
) -> Response:
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
@@ -1400,9 +1413,11 @@ class BaseUpstreamProvider:
if requested_model:
response_json["model"] = requested_model
if "id" not in response_json or not isinstance(response_json["id"], str):
response_json["id"] = f"chatcmpl-{uuid.uuid4()}"
prefix = "cmpl" if legacy_completion else "chatcmpl"
response_json["id"] = f"{prefix}-{uuid.uuid4()}"
if not isinstance(response_json.get("usage"), dict):
usage = response_json.get("usage")
if not isinstance(usage, dict) or not usage:
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
response_json["usage"] = usage_estimator.openai_response_data(
@@ -2923,6 +2938,7 @@ class BaseUpstreamProvider:
Returns:
Response or StreamingResponse from upstream with cost tracking
"""
completion_path = _openai_completion_path(path)
path = self.normalize_request_path(path, model_obj)
if (
@@ -2951,7 +2967,11 @@ class BaseUpstreamProvider:
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
)
transformed_body = self.prepare_request_body(request_body, model_obj)
transformed_body = self.prepare_request_body(
request_body,
model_obj,
include_stream_usage=completion_path is not None,
)
logger.debug(
"Forwarding request to upstream",
@@ -3048,7 +3068,7 @@ class BaseUpstreamProvider:
return mapped_error
if (
path.endswith("chat/completions")
completion_path is not None
or path.endswith("embeddings")
or path.endswith("messages")
or path.endswith("messages/count_tokens")
@@ -3117,7 +3137,7 @@ class BaseUpstreamProvider:
await response.aclose()
await client.aclose()
if path.endswith("chat/completions"):
if completion_path is not None:
client_wants_streaming = False
if request_body:
try:
@@ -3163,6 +3183,7 @@ class BaseUpstreamProvider:
reservation_snapshot=reservation_snapshot,
client=client,
request_body=request_body,
legacy_completion=completion_path == "completions",
)
# Handle both non-streaming chat completions and embeddings
@@ -3177,6 +3198,7 @@ class BaseUpstreamProvider:
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
request_body=request_body,
legacy_completion=completion_path == "completions",
)
finally:
await response.aclose()
@@ -4213,6 +4235,7 @@ class BaseUpstreamProvider:
Returns:
Response or StreamingResponse with refund if applicable
"""
completion_path = _openai_completion_path(path)
if path.startswith("v1/"):
path = path.replace("v1/", "")
@@ -4241,7 +4264,11 @@ class BaseUpstreamProvider:
url = f"{self.base_url}/{path}"
transformed_body = self.prepare_request_body(request_body, model_obj)
transformed_body = self.prepare_request_body(
request_body,
model_obj,
include_stream_usage=completion_path is not None,
)
logger.debug(
"Forwarding request to upstream",
@@ -4338,7 +4365,7 @@ class BaseUpstreamProvider:
return error_response
if (
path.endswith("chat/completions")
completion_path is not None
or path.endswith("embeddings")
or path.endswith("messages")
or path.endswith("messages/count_tokens")
+27 -4
View File
@@ -23,7 +23,11 @@ import litellm
from fastapi.responses import Response
from ..core import get_logger
from ..payment.helpers import estimate_prompt_tokens, estimate_tokens
from ..payment.helpers import (
_count_prompt_token_ids,
estimate_prompt_tokens,
estimate_tokens,
)
from ..payment.models import Model
logger = get_logger(__name__)
@@ -46,11 +50,28 @@ def _model_name(model_obj: Model | None, body: dict[str, Any]) -> str:
return body_model if isinstance(body_model, str) else ""
def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
def _count_with_litellm(
model: str, body: dict[str, Any], include_legacy_prompt: bool = False
) -> int:
messages = body.get("messages")
if not isinstance(messages, list):
messages = []
prompt_token_ids = 0
if include_legacy_prompt:
prompt = body.get("prompt")
if isinstance(prompt, str):
prompt_texts = [prompt]
elif isinstance(prompt, list):
prompt_texts = [item for item in prompt if isinstance(item, str)]
else:
prompt_texts = []
prompt_token_ids = _count_prompt_token_ids(prompt)
messages = [
*({"role": "user", "content": text} for text in prompt_texts if text),
*messages,
]
system = body.get("system")
if isinstance(system, str) and system:
messages = [{"role": "system", "content": system}, *messages]
@@ -65,7 +86,7 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
return int(
return prompt_token_ids + int(
litellm.token_counter(
model=model,
messages=messages,
@@ -133,7 +154,9 @@ class MissingUsageEstimator:
if self._input_tokens is not None:
return self._input_tokens
try:
self._input_tokens = _count_with_litellm(self.model_name, self.body)
self._input_tokens = _count_with_litellm(
self.model_name, self.body, include_legacy_prompt=True
)
except Exception as exc:
self._input_tokens = estimate_prompt_tokens(self.body)
logger.debug(
+384
View File
@@ -0,0 +1,384 @@
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi.responses import Response
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
import routstr.auth as auth_module
from routstr.auth import ReservationSnapshot, get_reservation_snapshot, pay_for_request
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream.base import BaseUpstreamProvider
BALANCE = 100_000
RESERVED = 5_000
# 1 msat per prompt token, 2 msat per completion token.
MODEL = Model(
id="glm-test",
name="glm-test",
created=0,
description="",
context_length=64_000,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Other",
instruct_type=None,
),
pricing=Pricing(prompt=0.001, completion=0.002),
sats_pricing=Pricing(prompt=0.001, completion=0.002),
)
USAGE = {"prompt_tokens": 400, "completion_tokens": 100, "total_tokens": 500}
EXPECTED_MSATS = 400 * 1 + 100 * 2
COMPLETION_BODY = {
"model": MODEL.id,
"prompt": "Once upon a time, in a land far away, " * 20,
"max_tokens": 100,
}
CHAT_BODY = {"model": MODEL.id, "messages": [{"role": "user", "content": "hi"}]}
COMPLETION_JSON = {
"id": "cmpl-1",
"object": "text_completion",
"model": MODEL.id,
"choices": [{"text": " there was", "index": 0, "finish_reason": "stop"}],
"usage": USAGE,
}
COMPLETION_CHUNKS = [
{
"id": "cmpl-1",
"object": "text_completion",
"model": MODEL.id,
"choices": [{"text": " there", "index": 0, "finish_reason": None}],
},
{
"id": "cmpl-1",
"object": "text_completion",
"model": MODEL.id,
"choices": [{"text": " was", "index": 0, "finish_reason": "stop"}],
},
]
USAGE_CHUNK = {
"id": "cmpl-1",
"object": "text_completion",
"model": MODEL.id,
"choices": [],
"usage": USAGE,
}
def _sse(chunks: list[dict]) -> bytes:
body = b"".join(b"data: " + json.dumps(c).encode() + b"\n\n" for c in chunks)
return body + b"data: [DONE]\n\n"
@pytest.fixture(autouse=True)
def patch_sats_usd_price() -> Any:
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4):
yield
async def _engine() -> AsyncEngine:
engine = create_async_engine("sqlite+aiosqlite://")
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
return engine
def _upstream(content: bytes, content_type: str) -> httpx.Response:
return httpx.Response(
200,
content=content,
headers={"content-type": content_type},
request=httpx.Request("POST", "http://upstream"),
)
async def _drain(response: Any) -> bytes:
body = b""
if hasattr(response, "body_iterator"):
async for chunk in response.body_iterator:
body += chunk if isinstance(chunk, bytes) else chunk.encode()
else:
body = response.body
return body
async def _forward(
engine: AsyncEngine,
path: str,
body: dict,
upstream: httpx.Response,
) -> tuple[bytes, ReservationSnapshot, AsyncMock]:
"""Reserve, forward through the real ``forward_request`` and settle."""
provider = BaseUpstreamProvider(
base_url="http://upstream", api_key="k", provider_fee=1.0
)
request = MagicMock()
request.method = "POST"
request.query_params = {}
send = AsyncMock(return_value=upstream)
async with AsyncSession(engine, expire_on_commit=False) as session:
key = ApiKey(hashed_key="key", balance=BALANCE)
session.add(key)
await session.commit()
await pay_for_request(key, RESERVED, session)
snapshot = await get_reservation_snapshot(key, session)
with (
patch("httpx.AsyncClient.send", send),
patch(
"routstr.upstream.base.create_session",
side_effect=lambda: AsyncSession(engine, expire_on_commit=False),
),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
auth_module.adjust_payment_for_tokens,
),
):
response = await provider.forward_request(
request,
path,
{},
json.dumps(body).encode(),
key,
RESERVED,
session,
MODEL,
snapshot,
)
out = await _drain(response)
return out, snapshot, send
async def _ledger(
engine: AsyncEngine, snapshot: ReservationSnapshot
) -> tuple[int, int, int, str | None]:
async with AsyncSession(engine, expire_on_commit=False) as session:
key = await session.get(ApiKey, snapshot.key_hash)
record = await session.get(ReservationRelease, snapshot.release_id)
assert key is not None
return (
key.balance,
key.total_spent,
key.reserved_balance,
record.status if record else None,
)
def _sse_objects(out: bytes) -> list[dict]:
objs = []
for line in out.split(b"\n"):
if line.startswith(b"data: ") and line[6:].strip() != b"[DONE]":
objs.append(json.loads(line[6:]))
return objs
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path",
[
"completions",
"v1/completions",
"v1/completions/",
"openai/v1/completions",
],
)
async def test_non_streaming_completion_with_usage_is_charged(path: str) -> None:
engine = await _engine()
out, snapshot, _ = await _forward(
engine,
path,
COMPLETION_BODY,
_upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"),
)
balance, spent, reserved, status = await _ledger(engine, snapshot)
assert (balance, spent, reserved, status) == (
BALANCE - EXPECTED_MSATS,
EXPECTED_MSATS,
0,
"charged",
)
body = json.loads(out)
assert body["object"] == "text_completion"
assert body["usage"]["cost"]["total_msats"] == EXPECTED_MSATS
await engine.dispose()
@pytest.mark.asyncio
async def test_streaming_completion_with_final_usage_is_charged() -> None:
engine = await _engine()
out, snapshot, send = await _forward(
engine,
"v1/completions",
{**COMPLETION_BODY, "stream": True},
_upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"),
)
balance, spent, reserved, status = await _ledger(engine, snapshot)
assert (balance, spent, reserved, status) == (
BALANCE - EXPECTED_MSATS,
EXPECTED_MSATS,
0,
"charged",
)
forwarded = json.loads(send.call_args.args[0].content)
assert forwarded["stream_options"] == {"include_usage": True}
objs = _sse_objects(out)
assert [o["choices"][0]["text"] for o in objs if o["choices"]] == [
" there",
" was",
]
assert objs[-1]["object"] == "text_completion"
assert objs[-1]["usage"]["cost"]["total_msats"] == EXPECTED_MSATS
assert out.endswith(b"data: [DONE]\n\n")
await engine.dispose()
@pytest.mark.asyncio
async def test_non_streaming_completion_with_empty_usage_is_estimated() -> None:
"""Missing usage is estimated from ``prompt`` and ``text``, never free."""
engine = await _engine()
no_usage = {k: v for k, v in COMPLETION_JSON.items() if k != "id"}
no_usage["usage"] = {}
out, snapshot, _ = await _forward(
engine,
"v1/completions",
COMPLETION_BODY,
_upstream(json.dumps(no_usage).encode(), "application/json"),
)
balance, spent, reserved, status = await _ledger(engine, snapshot)
assert 0 < spent <= RESERVED
assert (BALANCE - balance, reserved, status) == (spent, 0, "charged")
body = json.loads(out)
usage = body["usage"]
assert body["id"].startswith("cmpl-")
assert usage["estimated"] is True
assert usage["prompt_tokens"] > 100
assert usage["completion_tokens"] > 0
await engine.dispose()
@pytest.mark.asyncio
async def test_streaming_completion_without_usage_is_estimated() -> None:
engine = await _engine()
out, snapshot, _ = await _forward(
engine,
"v1/completions",
{**COMPLETION_BODY, "stream": True},
_upstream(_sse(COMPLETION_CHUNKS), "text/event-stream"),
)
balance, spent, reserved, status = await _ledger(engine, snapshot)
assert 0 < spent <= RESERVED
assert (BALANCE - balance, reserved, status) == (spent, 0, "charged")
trailer = _sse_objects(out)[-1]
assert trailer["id"] == "cmpl-1"
assert trailer["object"] == "text_completion"
assert trailer["usage"]["prompt_tokens"] > 100
assert trailer["usage"]["cost"]["total_msats"] == spent
await engine.dispose()
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["v1/chat/completions/", "openai/v1/chat/completions"])
async def test_chat_completion_aliases_are_charged(path: str) -> None:
engine = await _engine()
chat_json = {
"id": "chatcmpl-1",
"object": "chat.completion",
"model": MODEL.id,
"choices": [
{
"message": {"role": "assistant", "content": "hi"},
"index": 0,
"finish_reason": "stop",
}
],
"usage": USAGE,
}
out, snapshot, _ = await _forward(
engine,
path,
CHAT_BODY,
_upstream(json.dumps(chat_json).encode(), "application/json"),
)
balance, spent, reserved, status = await _ledger(engine, snapshot)
assert (balance, spent, reserved, status) == (
BALANCE - EXPECTED_MSATS,
EXPECTED_MSATS,
0,
"charged",
)
assert json.loads(out)["usage"]["cost"]["total_msats"] == EXPECTED_MSATS
await engine.dispose()
def test_stream_usage_option_is_scoped_to_completion_endpoints() -> None:
provider = BaseUpstreamProvider(base_url="http://upstream", api_key="k")
body = json.dumps({"prompt": "draw this", "stream": True}).encode()
assert provider.prepare_request_body(body, MODEL) == body
async def _forward_x_cashu(
path: str, body: dict, upstream: httpx.Response
) -> tuple[AsyncMock, AsyncMock]:
"""Run ``forward_x_cashu_request`` with the settlement handler stubbed out."""
provider = BaseUpstreamProvider(
base_url="http://upstream", api_key="k", provider_fee=1.0
)
request = MagicMock()
request.method = "POST"
request.query_params = {}
request.state.request_id = "req-1"
request.body = AsyncMock(return_value=json.dumps(body).encode())
send = AsyncMock(return_value=upstream)
settle = AsyncMock(return_value=Response(content=b"{}", status_code=200))
with (
patch("httpx.AsyncClient.send", send),
patch.object(provider, "handle_x_cashu_chat_completion", settle),
):
await provider.forward_x_cashu_request(
request, path, {}, 10, "sat", RESERVED, MODEL
)
return settle, send
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["completions", "v1/completions", "v1/completions/"])
async def test_x_cashu_legacy_completion_is_settled(path: str) -> None:
"""Legacy completions must reach refund settlement, not raw passthrough."""
settle, _ = await _forward_x_cashu(
path,
COMPLETION_BODY,
_upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"),
)
assert settle.await_count == 1
@pytest.mark.asyncio
async def test_x_cashu_streaming_completion_requests_usage() -> None:
_, send = await _forward_x_cashu(
"v1/completions",
{**COMPLETION_BODY, "stream": True},
_upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"),
)
forwarded = json.loads(send.call_args.args[0].content)
assert forwarded["stream_options"] == {"include_usage": True}
+9
View File
@@ -229,6 +229,15 @@ def test_missing_usage_estimator_openai_dialect() -> None:
}
def test_missing_usage_estimator_counts_legacy_token_prompt() -> None:
model = _make_model()
request_body = _body({"model": model.id, "prompt": [[1, 2], [3, 4, 5]]})
usage = MissingUsageEstimator(request_body, model).openai_response_data()["usage"]
assert usage["prompt_tokens"] >= 5
def test_missing_usage_estimator_does_not_count_response_metadata() -> None:
estimator = MissingUsageEstimator(b"{}", None)
estimator.observe(
+26
View File
@@ -183,6 +183,32 @@ def test_estimate_prompt_tokens_counts_every_string_in_the_body() -> None:
# value prefix can buy a discount, so both still count in full.
assert estimate_prompt_tokens({"tools": [{"data": hidden}]}) >= 1_000
assert estimate_prompt_tokens({"system": "data:" + hidden}) >= 1_000
assert estimate_prompt_tokens({"prompt": [[1, 2], [3, 4, 5]]}) >= 5
async def test_discount_counts_legacy_token_id_prompt() -> None:
from routstr.payment.helpers import calculate_discounted_max_cost
pricing = Mock()
pricing.prompt = 0.001
pricing.completion = 0.0
pricing.max_prompt_cost = 50.0
pricing.max_completion_cost = 0.0
model_obj = Mock()
model_obj.sats_pricing = pricing
model_obj.top_provider = None
model_obj.context_length = None
body = {"model": "test-model", "prompt": list(range(50_000)), "max_tokens": 0}
with (
patch.object(settings, "fixed_pricing", False),
patch.object(settings, "tolerance_percentage", 0),
patch.object(settings, "min_request_msat", 1000),
):
cost = await calculate_discounted_max_cost(50_000, body, model_obj)
assert cost == 50_000
async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None: