mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: settle X-Cashu streaming Responses instead of keeping the token
The streaming /v1/responses X-Cashu handler split the SSE body on the two-character sequence backslash-n and re-emitted with the same literal, so a real event stream was never framed into events. Even when a line happened to parse, usage and model were read at the top level, but the canonical Responses API nests them under "response" on response.completed. Both together meant usage was never found, the refund branch never ran, and the customer's whole token was retained. cost_data was also only bound inside that branch, so any path that did reach the re-emission loop would have raised NameError. Parse the body with a real SSE reader (CRLF, comment/keepalive lines, multi-line data fields, trailing event without terminator, [DONE] sentinel), resolve model and usage from the nested response object, and re-frame each event with genuine newlines so downstream clients get valid SSE. Settlement now always runs: unmeasured usage goes through the existing bounded fallback and settles at the authorized maximum with the remainder refunded, rather than silently keeping the full token.
This commit is contained in:
+176
-110
@@ -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(),
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
"""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.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))
|
||||
|
||||
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.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.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.args[0] == 1000
|
||||
body = await _collect(response)
|
||||
assert b"\\n" not in body
|
||||
assert body.endswith(b"\n\n")
|
||||
Reference in New Issue
Block a user