Merge pull request #691 from Routstr/fix/responses-xcashu-sse-refund

fix: settle X-Cashu streaming Responses instead of keeping the token
This commit is contained in:
9qeklajc
2026-08-24 23:59:13 +02:00
committed by GitHub
3 changed files with 399 additions and 112 deletions
+176 -110
View File
@@ -127,6 +127,49 @@ def _inject_cost_response_headers(
headers["X-Routstr-Cost-Usd"] = str(total_usd)
def _parse_sse_events(content: str) -> list[tuple[list[str], str]]:
"""Split a buffered SSE body into ``(field_lines, data)`` pairs.
``data`` is the newline-joined payload the SSE spec reassembles from every
``data:`` line of one event, so multi-line JSON survives. Comment/keepalive
lines are dropped and events carrying no data at all are skipped; the
remaining ``event:``/``id:``/``retry:`` fields stay attached to their event
so Responses API framing is preserved on re-emission. A trailing event
without its blank-line terminator is still returned.
"""
events: list[tuple[list[str], str]] = []
normalized = content.replace("\r\n", "\n").replace("\r", "\n")
for raw_event in normalized.split("\n\n"):
field_lines: list[str] = []
data_lines: list[str] = []
for line in raw_event.split("\n"):
if line.startswith("data:"):
data_lines.append(line[len("data:") :].lstrip(" "))
elif line and not line.startswith(":"):
field_lines.append(line)
if not data_lines:
continue
events.append((field_lines, "\n".join(data_lines)))
return events
def _responses_usage_payload(data_json: dict) -> dict:
"""Return the object carrying a Responses API event's model and usage.
Canonical events nest them under ``response`` (``response.completed`` /
``response.incomplete``); legacy and compat shapes keep them at top level.
"""
nested = data_json.get("response")
return nested if isinstance(nested, dict) else data_json
def _render_sse_event(field_lines: list[str], data: str) -> str:
"""Re-frame one parsed event, re-prefixing every line of a multi-line data."""
body = "".join(f"{line}\n" for line in field_lines)
body += "".join(f"data: {line}\n" for line in data.split("\n"))
return body + "\n"
def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> None:
"""Inject cost breakdown into the response body's ``usage.cost`` object.
@@ -4670,12 +4713,14 @@ class BaseUpstreamProvider:
Similar to regular streaming but handles Responses API specific tokens like reasoning_tokens.
"""
events = _parse_sse_events(content_str)
logger.debug(
"Processing streaming Responses API response",
extra={
"amount": amount,
"unit": unit,
"content_lines": len(content_str.strip().split("\\n")),
"event_count": len(events),
},
)
@@ -4685,30 +4730,49 @@ class BaseUpstreamProvider:
if "content-encoding" in response_headers:
del response_headers["content-encoding"]
usage_data = None
model = None
usage_data: dict | None = None
model: str | None = None
reasoning_tokens = 0
cost_data: CostData | MaxCostData | None = None
lines = content_str.strip().split("\\n")
for line in lines:
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if "usage" in data_json:
usage_data = data_json["usage"]
model = data_json.get("model")
# Track reasoning tokens for Responses API
if (
isinstance(usage_data, dict)
and "reasoning_tokens" in usage_data
):
reasoning_tokens = usage_data.get("reasoning_tokens", 0)
elif "model" in data_json and not model:
model = data_json["model"]
except json.JSONDecodeError:
continue
for _fields, data in events:
if data.strip() == "[DONE]":
continue
try:
data_json = json.loads(data)
except json.JSONDecodeError:
continue
if not isinstance(data_json, dict):
continue
# Canonical Responses API events carry model and usage nested under
# "response" (response.completed/incomplete); older shapes put them
# at the top level.
payload = _responses_usage_payload(data_json)
if isinstance(payload.get("usage"), dict):
usage_data = payload["usage"]
model = payload.get("model") or model
details = usage_data.get("output_tokens_details")
if isinstance(details, dict):
reasoning_tokens = details.get("reasoning_tokens", 0)
elif "reasoning_tokens" in usage_data:
reasoning_tokens = usage_data["reasoning_tokens"]
elif not model and payload.get("model"):
model = payload["model"]
if usage_data and model:
if usage_data is None:
# Settlement invariant: a terminal request is never silently
# zero-billed and never silently keeps the whole token. Unmeasured
# usage settles at the authorization ceiling and refunds the rest.
logger.warning(
"No usage in streaming Responses API response — settling at authorized max",
extra={
"model": model,
"amount": amount,
"unit": unit,
"max_cost_msats": max_cost_for_model,
},
)
else:
logger.debug(
"Found usage data in streaming Responses API response",
extra={
@@ -4720,97 +4784,99 @@ class BaseUpstreamProvider:
},
)
response_data = {"usage": usage_data, "model": model}
response_data = {"usage": usage_data, "model": model or "unknown"}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
if refund_amount > 0:
logger.debug(
"Processing refund for streaming Responses API response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
"reasoning_tokens": reasoning_tokens,
},
)
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming Responses API response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
else:
logger.debug(
"No refund needed for streaming Responses API response",
extra={
"amount": amount,
"cost_msats": cost_data.total_msats,
"model": model,
},
)
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e:
logger.error(
"Error calculating cost for streaming Responses API response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"amount": amount,
"unit": unit,
},
)
for i, (fields, data) in enumerate(events):
if data.strip() == "[DONE]":
continue
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
refund_amount = amount - cost_data.total_msats
elif unit == "sat":
refund_amount = amount - (cost_data.total_msats + 999) // 1000
else:
raise ValueError(f"Invalid unit: {unit}")
if refund_amount > 0:
logger.debug(
"Processing refund for streaming Responses API response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
"reasoning_tokens": reasoning_tokens,
},
)
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming Responses API response",
extra={
"refund_amount": refund_amount,
"unit": unit,
"refund_token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
},
)
else:
logger.debug(
"No refund needed for streaming Responses API response",
extra={
"amount": amount,
"cost_msats": cost_data.total_msats,
"model": model,
},
)
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e:
logger.error(
"Error calculating cost for streaming Responses API response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"amount": amount,
"unit": unit,
},
)
for i, line in enumerate(lines):
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if not isinstance(data_json, dict):
continue
changed = False
if "provider" not in data_json:
self._apply_provider_field(data_json)
changed = True
if cost_data and "usage" in data_json and data_json["usage"]:
_inject_cost_into_usage(data_json, cost_data)
changed = True
if changed:
lines[i] = "data: " + json.dumps(data_json)
except json.JSONDecodeError:
pass
data_json = json.loads(data)
except json.JSONDecodeError:
continue
if not isinstance(data_json, dict):
continue
changed = False
if "provider" not in data_json:
self._apply_provider_field(data_json)
changed = True
payload = _responses_usage_payload(data_json)
if cost_data and isinstance(payload.get("usage"), dict):
_inject_cost_into_usage(payload, cost_data)
changed = True
if changed:
events[i] = (fields, json.dumps(data_json))
async def generate() -> AsyncGenerator[bytes, None]:
for line in lines:
yield (line + "\\n").encode("utf-8")
for fields, data in events:
yield _render_sse_event(fields, data).encode("utf-8")
return StreamingResponse(
generate(),
+2 -2
View File
@@ -339,7 +339,7 @@ async def test_slow_mints_do_not_exhaust_a_single_connection_pool(
f"sqlite+aiosqlite:///{tmp_path / 'pool-pressure.db'}",
pool_size=1,
max_overflow=0,
pool_timeout=0.2,
pool_timeout=0.5,
)
async with engine.begin() as connection:
await connection.run_sync(SQLModel.metadata.create_all)
@@ -350,7 +350,7 @@ async def test_slow_mints_do_not_exhaust_a_single_connection_pool(
yield session
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
await asyncio.sleep(0.3)
await asyncio.sleep(1.0)
return proofs
try:
@@ -0,0 +1,221 @@
"""X-Cashu settlement for streaming ``/v1/responses``.
The stream is real SSE: CRLF delimiters, comment keepalives, ``event:`` fields,
multi-line ``data:`` payloads and a ``[DONE]`` sentinel. Canonical Responses API
usage arrives nested under ``response`` on ``response.completed``.
"""
import json
import os
from typing import Any
from unittest.mock import AsyncMock, patch
import httpx
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
def _make_provider() -> BaseUpstreamProvider:
return BaseUpstreamProvider(base_url="http://test", api_key="test-key")
def _make_cost_data(total_msats: int = 4000) -> CostData:
return CostData(
base_msats=0,
input_msats=2500,
output_msats=1500,
total_msats=total_msats,
total_usd=0.0002,
input_tokens=12,
output_tokens=8,
)
def _sse_response(chunks: list[bytes]) -> httpx.Response:
"""Build the upstream response from wire chunks that split events."""
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content=b"".join(chunks),
)
COMPLETED_EVENT = {
"type": "response.completed",
"response": {
"model": "gpt-5-mini",
"usage": {
"input_tokens": 12,
"output_tokens": 8,
"total_tokens": 20,
"output_tokens_details": {"reasoning_tokens": 3},
},
},
}
def _canonical_chunks() -> list[bytes]:
"""CRLF stream whose completed event straddles two wire chunks."""
completed = json.dumps(COMPLETED_EVENT).encode()
return [
b": keepalive\r\n\r\n",
b"event: response.created\r\n"
b'data: {"type":"response.created","response":{"model":"gpt-5-mini"}}\r\n\r\n',
b"event: response.completed\r\ndata: " + completed[:40],
completed[40:] + b"\r\n\r\n",
b"data: [DONE]\r\n\r\n",
]
async def _collect(response: Any) -> bytes:
body = b""
async for chunk in response.body_iterator:
body += chunk
return body
async def _settle(
chunks: list[bytes],
*,
amount: int = 10_000,
max_cost_for_model: int = 9_000,
cost_data: CostData | None = None,
) -> tuple[Any, AsyncMock, AsyncMock]:
provider = _make_provider()
get_cost = (
AsyncMock(return_value=cost_data)
if cost_data is not None
else AsyncMock(side_effect=provider.get_x_cashu_cost)
)
send_refund = AsyncMock(return_value="cashuBrefundtoken0123456789")
with (
patch.object(provider, "get_x_cashu_cost", new=get_cost),
patch.object(provider, "send_refund", new=send_refund),
):
response = await provider.handle_x_cashu_responses_completion(
response=_sse_response(chunks),
amount=amount,
unit="msat",
max_cost_for_model=max_cost_for_model,
mint=None,
)
return response, get_cost, send_refund
@pytest.mark.asyncio
async def test_fragmented_crlf_stream_refunds_and_sets_cost_headers() -> None:
response, _, send_refund = await _settle(
_canonical_chunks(), cost_data=_make_cost_data(4000)
)
send_refund.assert_awaited_once()
assert send_refund.await_args is not None
assert send_refund.await_args.args[0] == 10_000 - 4000
assert response.headers["x-cashu"] == "cashuBrefundtoken0123456789"
assert response.headers["x-routstr-cost-msats"] == "4000"
assert response.headers["x-routstr-input-cost-msats"] == "2500"
assert response.headers["x-routstr-output-cost-msats"] == "1500"
@pytest.mark.asyncio
async def test_nested_completion_usage_drives_cost_calculation() -> None:
_, get_cost, _ = await _settle(_canonical_chunks(), cost_data=_make_cost_data(4000))
assert get_cost.await_args is not None
response_data = get_cost.await_args.args[0]
assert response_data["model"] == "gpt-5-mini"
assert response_data["usage"]["input_tokens"] == 12
assert response_data["usage"]["output_tokens"] == 8
@pytest.mark.asyncio
async def test_reemitted_stream_is_valid_sse() -> None:
response, _, _ = await _settle(_canonical_chunks(), cost_data=_make_cost_data(4000))
body = await _collect(response)
assert b"\\n" not in body
assert body.endswith(b"\n\n")
assert b": keepalive" not in body
events = [e for e in body.split(b"\n\n") if e.strip()]
payloads = []
for event in events:
data_lines = [
line[len(b"data:") :].lstrip()
for line in event.split(b"\n")
if line.startswith(b"data:")
]
assert data_lines, f"event carries no data line: {event!r}"
payloads.append(b"\n".join(data_lines))
assert payloads[-1] == b"[DONE]"
assert any(b"event: response.completed" in event for event in events)
completed = json.loads(payloads[-2])
assert completed["type"] == "response.completed"
assert completed["response"]["usage"]["cost"]["total_msats"] == 4000
@pytest.mark.asyncio
async def test_multiline_data_payload_is_parsed_and_reframed() -> None:
completed = json.dumps(COMPLETED_EVENT)
head, tail = completed[:30], completed[30:]
chunks = [
("data: " + head + "\r\ndata: " + tail + "\r\n\r\n").encode(),
b"data: [DONE]\r\n\r\n",
]
response, get_cost, send_refund = await _settle(
chunks, cost_data=_make_cost_data(4000)
)
assert get_cost.await_args is not None
assert send_refund.await_args is not None
assert get_cost.await_args.args[0]["usage"]["input_tokens"] == 12
assert send_refund.await_args.args[0] == 6000
body = await _collect(response)
for event in body.split(b"\n\n"):
for line in event.split(b"\n"):
if line.strip():
assert line.startswith(b"data:") or line.startswith(b"event:")
@pytest.mark.asyncio
async def test_missing_usage_settles_at_authorized_max() -> None:
chunks = [
b'data: {"type":"response.created","response":{"model":"gpt-5-mini"}}\r\n\r\n',
b"data: [DONE]\r\n\r\n",
]
response, _, send_refund = await _settle(
chunks, amount=10_000, max_cost_for_model=9_000
)
send_refund.assert_awaited_once()
assert send_refund.await_args is not None
assert send_refund.await_args.args[0] == 10_000 - 9_000
assert response.headers["x-cashu"] == "cashuBrefundtoken0123456789"
assert response.headers["x-routstr-cost-msats"] == "9000"
@pytest.mark.asyncio
async def test_malformed_events_do_not_retain_whole_token() -> None:
chunks = [
b"data: {not json\r\n\r\n",
b"data: [DONE]\r\n\r\n",
]
response, _, send_refund = await _settle(
chunks, amount=10_000, max_cost_for_model=9_000
)
assert send_refund.await_args is not None
assert send_refund.await_args.args[0] == 1000
body = await _collect(response)
assert b"\\n" not in body
assert body.endswith(b"\n\n")