This commit is contained in:
9qeklajc
2026-09-07 00:30:05 +02:00
parent 49f2256e3b
commit ac40d70c3a
10 changed files with 365 additions and 619 deletions
+21 -122
View File
@@ -1191,13 +1191,27 @@ async def adjust_payment_for_tokens(
calculated_cost = await calculate_cost(
response_data, deducted_max_cost, model_obj, provider_fee
)
if not isinstance(calculated_cost, CostDataError):
if not await _claim_reservation_for_charge(reservation, session):
# A prior charge or release already owns this reservation. Returning
# the calculated metadata is safe; the aggregate balances must not
# be modified a second time.
calculated_cost.charged_msats = 0
return calculated_cost.dict()
if isinstance(calculated_cost, CostDataError):
# Content was already served, so release instead of raising a 400.
logger.error(
"Cost calculation error during payment adjustment, releasing reservation",
extra={
"key_hash": key_log_hash,
"model": model,
"error_message": calculated_cost.message,
"error_code": calculated_cost.code,
},
)
calculated_cost = MaxCostData(
base_msats=0, input_msats=0, output_msats=0, total_msats=0
)
if not await _claim_reservation_for_charge(reservation, session):
# A prior charge or release already owns this reservation. Returning
# the calculated metadata is safe; the aggregate balances must not
# be modified a second time.
calculated_cost.charged_msats = 0
return calculated_cost.dict()
match calculated_cost:
case MaxCostData() as cost:
@@ -1522,121 +1536,6 @@ async def adjust_payment_for_tokens(
return cost.dict()
case CostDataError() as error:
# Pricing derivation failed AFTER the upstream served the request.
# Raising here would hand the client a 400 for content it already
# received (streaming) while the provider eats the upstream cost.
# Apply missing_usage_policy instead: keep the pre-authorized
# ceiling (charge_max), or release without charging
# (estimate/refund) and let the response complete.
policy = (settings.missing_usage_policy or "charge_max").strip().lower()
logger.error(
"Cost calculation error during payment adjustment — applying "
"missing_usage_policy=%s instead of raising",
policy,
extra={
"key_hash": key.hashed_key[:8] + "...",
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
if policy == "charge_max":
charged = await _charge_reservation_rows(
session,
billing_key_hash=billing_key.hashed_key,
reserved_msats=deducted_max_cost,
charge_msats=deducted_max_cost,
)
if charged:
await session.commit()
await _stop_reservation_heartbeat(reservation.release_id)
await session.refresh(billing_key)
await _accumulate_fee(deducted_max_cost)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_log_hash,
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": deducted_max_cost,
"input_tokens": 0,
"output_tokens": 0,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "missing_usage_policy",
},
)
return {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": deducted_max_cost,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"cache_read_msats": 0,
"cache_creation_msats": 0,
"charged_msats": deducted_max_cost,
"reason": "missing_usage",
"estimated": True,
}
logger.error(
"Failed to charge reservation under missing_usage_policy=charge_max "
"— releasing instead",
extra={
"key_hash": key_log_hash,
"model": model,
"error_message": error.message,
},
)
else:
if policy in ("estimate", "refund"):
logger.warning(
"Releasing reservation without charging under "
"missing_usage_policy=%s",
policy,
extra={
"key_hash": key_log_hash,
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
else:
logger.warning(
"Unknown missing_usage_policy %r — treating as 'estimate'",
policy,
extra={
"key_hash": key_log_hash,
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
await release_reservation_only()
return {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"cache_read_msats": 0,
"cache_creation_msats": 0,
"charged_msats": 0,
"reason": "missing_usage",
"estimated": True,
"error": {"message": error.message, "code": error.code},
}
# All calculate_cost variants are handled above.
raise AssertionError("Unreachable: unhandled calculate_cost result")
-8
View File
@@ -79,14 +79,6 @@ class Settings(BaseSettings):
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
# Minimum per-request charge in millisatoshis when model pricing is free/zero
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
# Policy when an upstream response carries no usable usage AND no usable
# pricing (content was served, cost cannot be measured or derived):
# estimate — charge whatever a local token estimate yields (may be 0)
# charge_max — keep the prepayment/reservation (user pre-authorized it)
# refund — release the reservation / refund the full prepayment
missing_usage_policy: str = Field(
default="charge_max", env="MISSING_USAGE_POLICY"
)
reset_reserved_balance_on_startup: bool = Field(
default=True, env="RESET_RESERVED_BALANCE_ON_STARTUP"
) # deactivate in horizontal scaling setups
+19 -70
View File
@@ -41,26 +41,7 @@ class CostData(BaseModel):
class MaxCostData(CostData):
"""Reservation-ceiling billing.
Two distinct meanings ride on this class:
- ``reason="max_cost"`` — pricing is usable but the response is empty or
the upstream reports a USD cost with zero tokens; the ceiling is the
agreed charge for a served-but-unmeasurable request.
- ``reason="missing_usage"`` — usage AND pricing were both unusable and
the ``missing_usage_policy`` setting chose the ceiling (or a refund).
``total_msats == 0`` under this reason means "charge nothing and
release", per the ``refund`` policy.
Callers that need to distinguish these (dashboards, estimated markers)
read ``reason``; billing behavior only reads ``total_msats``.
"""
reason: str = "max_cost"
# Integer because `_cost_field` only handles numeric fields; 1 marks a
# charge derived under missing_usage_policy rather than measured usage.
estimated_flag: int = 0
pass
class CostDataError(BaseModel):
@@ -90,45 +71,6 @@ def _empty_cost(cls: type[CostData] = CostData) -> CostData:
)
def _missing_usage_max_cost(max_cost: int) -> MaxCostData:
"""Apply ``missing_usage_policy`` when usage AND pricing are unusable.
The content was served, so the request cannot be free by accident of the
upstream omitting its usage trailer. ``estimate`` keeps legacy behavior
(charge 0, reservation released); ``charge_max`` bills the pre-authorized
ceiling; ``refund`` bills 0 explicitly. The zero-charge variants carry
``reason="missing_usage"`` so callers can mark the charge as estimated
rather than measured.
"""
policy = (settings.missing_usage_policy or "charge_max").strip().lower()
if policy == "charge_max" and max_cost > 0:
logger.warning(
"No usage data and no usable pricing — applying "
"missing_usage_policy=charge_max: billing the pre-authorized "
"reservation ceiling.",
extra={"max_cost_msats": max_cost},
)
return MaxCostData(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=max_cost,
total_usd=0.0,
reason="missing_usage",
estimated_flag=1,
)
if policy not in ("estimate", "charge_max", "refund"):
logger.warning(
"Unknown missing_usage_policy %r — treating as 'estimate'",
policy,
)
zero = _empty_cost(MaxCostData)
assert isinstance(zero, MaxCostData)
zero.reason = "missing_usage"
zero.estimated_flag = 1
return zero
async def calculate_cost(
response_data: dict,
max_cost: int,
@@ -182,7 +124,7 @@ async def calculate_cost(
else None,
},
)
return _missing_usage_max_cost(max_cost)
return _empty_cost(MaxCostData)
usage_data = response_data.get("usage") or {}
if not isinstance(usage_data, dict):
@@ -311,8 +253,11 @@ async def calculate_cost(
rates = (input_rate, output_rate, cache_read_rate, cache_creation_rate)
if not all(is_usable_rate(rate) for rate in rates):
logger.warning(
"No usable token pricing — applying missing_usage_policy instead of "
"treating the reservation ceiling as the charge or releasing for free.",
"No usable token pricing — releasing the reservation instead of "
"treating its ceiling as the charge. Token counts %s in the "
"upstream response but cannot be converted to money; the request "
"will appear in dashboards with raw counts and a zero charge.",
"are present" if (input_tokens > 0 or output_tokens > 0) else "are zero",
extra={
"base_cost_msats": max_cost,
"model": response_data.get("model", "unknown"),
@@ -322,14 +267,18 @@ async def calculate_cost(
"output_rate": output_rate,
},
)
missing = _missing_usage_max_cost(max_cost)
# Preserve the raw token counts for dashboards even when no usable
# rate exists to convert them to money.
missing.input_tokens = input_tokens
missing.output_tokens = output_tokens
missing.cache_read_input_tokens = cache_read_tokens
missing.cache_creation_input_tokens = cache_creation_tokens
return missing
return MaxCostData(
base_msats=0,
input_msats=0,
output_msats=0,
total_msats=0,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=0,
cache_creation_msats=0,
)
return _calculate_from_tokens(
input_tokens,
+123 -177
View File
@@ -22,7 +22,6 @@ from ..auth import (
release_reservation,
)
from ..core import get_logger
from ..core.settings import settings
from ..core.db import (
ApiKey,
AsyncSession,
@@ -153,8 +152,15 @@ def _inject_cost_response_headers(
total_usd = float(_cost_field(cost_data, "total_usd", 0.0))
if total_usd:
headers["X-Routstr-Cost-Usd"] = str(total_usd)
if _cost_field(cost_data, "estimated_flag", 0) == 1:
headers["X-Routstr-Cost-Estimated"] = "true"
def _estimated_usage(
estimator: MissingUsageEstimator, model: str | None
) -> dict[str, Any] | None:
"""Local usage estimate, or None when the upstream generated no text."""
if not estimator.output_text:
return None
return estimator.response_data(model)["usage"]
def _parse_sse_events(content: str) -> list[tuple[list[str], str]]:
@@ -3731,23 +3737,17 @@ class BaseUpstreamProvider:
)
return cost
case CostDataError() as error:
# Content was already served, so refund instead of raising.
logger.error(
"Cost calculation error",
"Cost calculation error, refunding the prepayment",
extra={
"model": model,
"error_message": error.message,
"error_code": error.code,
},
)
raise HTTPException(
status_code=400,
detail={
"error": {
"message": error.message,
"type": "invalid_request_error",
"code": error.code,
}
},
return MaxCostData(
base_msats=0, input_msats=0, output_msats=0, total_msats=0
)
return None
@@ -3862,10 +3862,6 @@ class BaseUpstreamProvider:
usage_data = None
model = None
cost_data: CostData | MaxCostData | None = None
# Local estimator fed with every streamed text event — used to build
# an auditable usage estimate when the upstream omits its usage
# trailer, instead of silently keeping the full prepayment.
usage_estimator = MissingUsageEstimator(request_body, model_obj)
lines = content_str.strip().split("\n")
@@ -3897,152 +3893,97 @@ class BaseUpstreamProvider:
usage_estimator.observe(data_json)
if usage_data is None:
# No usage trailer: bill from the local token estimate instead of
# keeping the whole prepayment (legacy behavior). Only when the
# estimator produced nothing at all do we fall through with no
# usage, letting `missing_usage_policy` decide.
estimated = usage_estimator.response_data(model)
if estimated["usage"]["output_tokens"] > 0 or (
estimated["usage"]["input_tokens"] > 0
):
usage_data = _estimated_usage(usage_estimator, model)
if usage_data:
logger.warning(
"No usage in streaming x-cashu response — billing from "
"local token estimate",
"No usage in streaming response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"estimated_usage": estimated["usage"],
"estimated_usage": usage_data,
},
)
usage_data = estimated["usage"]
model = model or estimated["model"]
if usage_data and model:
logger.debug(
"Found usage data in streaming response",
logger.debug(
"Calculating cost for streaming response",
extra={
"model": model,
"usage_data": usage_data,
"amount": amount,
"unit": unit,
},
)
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 response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
},
)
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 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 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 response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"usage_data": usage_data,
"amount": amount,
"unit": unit,
},
)
response_data = {"usage": usage_data, "model": model}
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 response",
extra={
"original_amount": amount,
"cost_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"unit": unit,
"model": model,
},
)
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 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 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 response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"model": model,
"amount": amount,
"unit": unit,
},
)
else:
# Still nothing billable (no usage, no estimate): let
# missing_usage_policy decide instead of silently keeping the
# prepayment (legacy behavior).
try:
cost_data = await self.get_x_cashu_cost(
{"usage": None, "model": model or "unknown"},
max_cost_for_model,
model_obj,
)
if cost_data:
_inject_cost_response_headers(response_headers, cost_data)
refund_amount = (
amount - cost_data.total_msats
if unit == "msat"
else amount - (cost_data.total_msats + 999) // 1000
)
if refund_amount > 0:
refund_token = await self.send_refund(
refund_amount,
unit,
mint,
request_id=request_id,
)
response_headers["X-Cashu"] = refund_token
logger.warning(
"No usage and no estimate in streaming x-cashu "
"response — applied missing_usage_policy",
extra={
"policy": settings.missing_usage_policy,
"charge_msats": cost_data.total_msats,
"refund_amount": refund_amount,
"model": model,
},
)
except Exception as e:
logger.error(
"Error applying missing_usage_policy for streaming x-cashu response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"amount": amount,
"unit": unit,
},
)
for i, line in enumerate(lines):
if line.startswith("data: "):
try:
@@ -4103,29 +4044,23 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
if not isinstance(response_json.get("usage"), dict) or not response_json[
"usage"
]:
# No upstream usage: bill from the local token estimate rather
# than keeping the full prepayment (legacy behavior).
if not response_json.get("usage"):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
estimated = usage_estimator.response_data(response_json.get("model"))
if (
estimated["usage"]["output_tokens"] > 0
or estimated["usage"]["input_tokens"] > 0
):
estimated = _estimated_usage(
usage_estimator, response_json.get("model")
)
if estimated:
logger.warning(
"No usage in non-streaming x-cashu response — billing "
"from local token estimate",
"No usage in non-streaming response, billing from local token estimate",
extra={
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated["usage"],
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated["usage"]
response_json["usage"] = estimated
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
@@ -4260,6 +4195,7 @@ class BaseUpstreamProvider:
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse | Response:
"""Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming.
@@ -4285,10 +4221,6 @@ class BaseUpstreamProvider:
is_streaming = _is_sse_body(
response.headers.get("content-type"), content_str
)
# The original request body is not reachable at this settlement
# seam; pass None so the missing-usage estimator falls back to
# output-text counting only (no prompt-token estimate).
request_body: bytes | None = None
logger.debug(
"Chat completion response analysis",
@@ -4517,6 +4449,7 @@ class BaseUpstreamProvider:
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
request_body=request_body,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4807,6 +4740,7 @@ class BaseUpstreamProvider:
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
request_body=request_body,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4858,6 +4792,7 @@ class BaseUpstreamProvider:
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
request_body: bytes | None = None,
) -> StreamingResponse | Response:
"""Handle Responses API completion response for X-Cashu payment.
@@ -4884,10 +4819,6 @@ class BaseUpstreamProvider:
is_streaming = _is_sse_body(
response.headers.get("content-type"), content_str
)
# The original request body is not reachable at this settlement
# seam; pass None so the missing-usage estimator falls back to
# output-text counting only (no prompt-token estimate).
request_body: bytes | None = None
logger.debug(
"Responses API completion response analysis",
@@ -4977,6 +4908,7 @@ class BaseUpstreamProvider:
model: str | None = None
reasoning_tokens = 0
cost_data: CostData | MaxCostData | None = None
usage_estimator = MissingUsageEstimator(request_body, model_obj)
for _fields, data in events:
if data.strip() == "[DONE]":
@@ -4987,6 +4919,7 @@ class BaseUpstreamProvider:
continue
if not isinstance(data_json, dict):
continue
usage_estimator.observe(data_json)
# Canonical Responses API events carry model and usage nested under
# "response" (response.completed/incomplete); older shapes put them
# at the top level.
@@ -5003,18 +4936,14 @@ class BaseUpstreamProvider:
model = payload["model"]
if usage_data is None:
# No usage in the stream: with no measured tokens there is no
# auditable estimate, so `missing_usage_policy` decides —
# charge_max keeps the pre-authorized ceiling, refund/estimate
# release it. The charge itself comes from get_x_cashu_cost ->
# calculate_cost below.
usage_data = _estimated_usage(usage_estimator, model)
logger.warning(
"No usage in streaming Responses API response — applying missing_usage_policy",
"No usage in streaming Responses API response, billing from local token estimate",
extra={
"model": model,
"amount": amount,
"unit": unit,
"max_cost_msats": max_cost_for_model,
"estimated_usage": usage_data,
},
)
else:
@@ -5150,6 +5079,23 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
if not response_json.get("usage"):
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_estimator.observe(response_json)
estimated = _estimated_usage(
usage_estimator, response_json.get("model")
)
if estimated:
logger.warning(
"No usage in non-streaming Responses API response, billing from local token estimate",
extra={
"model": response_json.get("model", "unknown"),
"amount": amount,
"unit": unit,
"estimated_usage": estimated,
},
)
response_json["usage"] = estimated
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
+9 -11
View File
@@ -12,7 +12,6 @@ import uuid
from unittest.mock import patch
import pytest
from fastapi import HTTPException
from sqlmodel import col, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
@@ -305,21 +304,20 @@ async def test_cost_error_releases_the_reservation_without_charging(
"routstr.auth.calculate_cost",
return_value=CostDataError(message="no pricing", code="pricing_error"),
):
with pytest.raises(HTTPException) as exc:
await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert exc.value.status_code == 400
cost = await adjust_payment_for_tokens(
key,
_response(),
integration_session,
reserved,
reservation_snapshot=reservation,
)
assert cost["charged_msats"] == 0
key = await integration_session.get(ApiKey, key_hash)
assert key is not None
assert key.balance == 10_000, "a pricing failure must not charge the user"
assert key.total_spent == 0
assert key.reserved_balance == 0, "funds must not stay locked after a 400"
assert key.reserved_balance == 0, "funds must not stay locked"
assert await _active_reservations(integration_session) == 0
@@ -0,0 +1,47 @@
"""A pricing failure after the upstream served content must not raise a 400."""
from unittest.mock import AsyncMock, patch
import pytest
from routstr.auth import ReservationSnapshot, adjust_payment_for_tokens
from routstr.core.db import ApiKey
from routstr.payment import cost_calculation
@pytest.mark.asyncio
async def test_cost_data_error_releases_without_raising() -> None:
key_hash = "a" * 64
reservation = ReservationSnapshot(
release_id="rel-1",
key_hash=key_hash,
billing_key_hash=key_hash,
reserved_msats=7_000,
)
with (
patch.object(
cost_calculation,
"_get_pricing_rates",
side_effect=ValueError("no pricing for model"),
),
patch("routstr.auth._validate_reservation_snapshot", new=AsyncMock()),
patch("routstr.auth._stop_reservation_heartbeat", new=AsyncMock()),
patch(
"routstr.auth._claim_reservation_for_charge",
new=AsyncMock(return_value=True),
),
patch(
"routstr.auth._charge_reservation_rows", new=AsyncMock(return_value=True)
),
patch("routstr.auth.accumulate_routstr_fee", new=AsyncMock()),
):
cost = await adjust_payment_for_tokens(
ApiKey(hashed_key=key_hash),
{"model": "gpt-4o", "usage": {"prompt_tokens": 10, "completion_tokens": 5}},
session=AsyncMock(),
deducted_max_cost=7_000,
reservation_snapshot=reservation,
)
assert cost["total_msats"] == 0
assert cost["charged_msats"] == 0
-216
View File
@@ -1,216 +0,0 @@
"""Missing-usage billing policy tests.
Covers the three money paths touched by `missing_usage_policy`:
1. `calculate_cost` with NO usage at all (the `_empty_cost(MaxCostData)` dead-end).
2. `calculate_cost` with token counts but unusable pricing (the second dead-end).
3. `adjust_payment_for_tokens` `CostDataError` handling — must NOT raise a
post-delivery 400.
4. X-Cashu non-streaming handler bills from the local estimate when the
upstream omits usage.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from routstr.payment import cost_calculation
from routstr.payment.cost_calculation import (
CostData,
MaxCostData,
calculate_cost,
)
from routstr.payment.usage import normalize_usage
def _response(usage=None):
data = {"model": "gpt-4o", "id": "x", "object": "chat.completion"}
if usage is not None:
data["usage"] = usage
return data
# ---------------------------------------------------------------------------
# calculate_cost: no usage at all
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_no_usage_charge_max_bills_ceiling(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "charge_max")
cost = await calculate_cost(_response(), max_cost=50_000, model_obj=None)
assert isinstance(cost, MaxCostData)
assert cost.total_msats == 50_000
assert cost.reason == "missing_usage"
assert cost.estimated_flag == 1
@pytest.mark.asyncio
async def test_no_usage_estimate_bills_zero(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "estimate")
cost = await calculate_cost(_response(), max_cost=50_000, model_obj=None)
assert isinstance(cost, MaxCostData)
assert cost.total_msats == 0
assert cost.reason == "missing_usage"
@pytest.mark.asyncio
async def test_no_usage_refund_bills_zero(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "refund")
cost = await calculate_cost(_response(), max_cost=50_000, model_obj=None)
assert cost.total_msats == 0
assert cost.reason == "missing_usage"
@pytest.mark.asyncio
async def test_no_usage_unknown_policy_treated_as_estimate(monkeypatch):
monkeypatch.setattr(
cost_calculation.settings, "missing_usage_policy", "garbage"
)
cost = await calculate_cost(_response(), max_cost=50_000, model_obj=None)
assert cost.total_msats == 0
@pytest.mark.asyncio
async def test_no_usage_charge_max_zero_ceiling_stays_zero(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "charge_max")
cost = await calculate_cost(_response(), max_cost=0, model_obj=None)
assert cost.total_msats == 0
# ---------------------------------------------------------------------------
# calculate_cost: tokens present but pricing unusable
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_tokens_without_pricing_charge_max(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "charge_max")
# NaN pricing rates fail the usable-rate gate -> policy applies.
with patch.object(
cost_calculation,
"_get_pricing_rates",
return_value=(float("nan"), 1.0, 1.0, 1.0),
):
usage = {"prompt_tokens": 100, "completion_tokens": 50}
cost = await calculate_cost(_response(usage), max_cost=9_999, model_obj=None)
assert isinstance(cost, MaxCostData)
assert cost.total_msats == 9_999
assert cost.reason == "missing_usage"
assert cost.input_tokens == 100
assert cost.output_tokens == 50
@pytest.mark.asyncio
async def test_tokens_without_pricing_estimate(monkeypatch):
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "estimate")
with patch.object(
cost_calculation,
"_get_pricing_rates",
return_value=(float("nan"), 1.0, 1.0, 1.0),
):
usage = {"prompt_tokens": 100, "completion_tokens": 50}
cost = await calculate_cost(_response(usage), max_cost=9_999, model_obj=None)
assert cost.total_msats == 0
# Token counts still surface for dashboards.
assert cost.input_tokens == 100
assert cost.output_tokens == 50
# ---------------------------------------------------------------------------
# adjust_payment_for_tokens: CostDataError must not raise post-delivery
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_cost_data_error_charge_max_charges_and_returns(monkeypatch):
from routstr.auth import ReservationSnapshot, adjust_payment_for_tokens
from routstr.core.db import ApiKey
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "charge_max")
key = ApiKey(hashed_key="a" * 64, balance=1_000_000, reserved_balance=50_000)
reservation = ReservationSnapshot(
release_id="rel-1",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=50_000,
)
with (
patch("routstr.auth._validate_reservation_snapshot", new=AsyncMock()),
patch("routstr.auth._stop_reservation_heartbeat", new=AsyncMock()),
patch("routstr.auth._claim_reservation_for_charge", new=AsyncMock(return_value=True)),
patch("routstr.auth._charge_reservation_rows", new=AsyncMock(return_value=True)),
patch("routstr.auth.get_reservation_snapshot", new=AsyncMock(return_value=reservation)),
patch(
"routstr.auth.accumulate_routstr_fee",
new=AsyncMock(),
) as accumulate_fee,
):
cost = await adjust_payment_for_tokens(
key,
_response(), # no usage -> policy path via MaxCostData
session=AsyncMock(),
deducted_max_cost=50_000,
reservation_snapshot=reservation,
)
# The MaxCostData path bills the ceiling through normal finalization.
assert cost["total_msats"] == 50_000
assert cost["charged_msats"] == 50_000
@pytest.mark.asyncio
async def test_cost_data_error_path_returns_dict_not_raise(monkeypatch):
"""Force a genuine CostDataError (pricing ValueError) under 'refund'."""
from routstr.auth import ReservationSnapshot, adjust_payment_for_tokens
from routstr.core.db import ApiKey
monkeypatch.setattr(cost_calculation.settings, "missing_usage_policy", "refund")
usage = {"prompt_tokens": 10, "completion_tokens": 5}
# No model_obj and no fixed pricing -> usable-rate gate... but tokens are
# present, so to force a CostDataError we patch _get_pricing_rates.
with (
patch.object(
cost_calculation,
"_get_pricing_rates",
side_effect=ValueError("no pricing for model"),
),
patch("routstr.auth._validate_reservation_snapshot", new=AsyncMock()),
patch("routstr.auth._stop_reservation_heartbeat", new=AsyncMock()),
patch("routstr.auth._claim_reservation_for_charge", new=AsyncMock(return_value=True)),
patch("routstr.auth._charge_reservation_rows", new=AsyncMock(return_value=True)),
patch("routstr.auth.release_reservation", new=AsyncMock(return_value=True)),
patch(
"routstr.auth.accumulate_routstr_fee",
new=AsyncMock(),
),
):
cost = await adjust_payment_for_tokens(
ApiKey(hashed_key="a" * 64),
_response(usage),
session=AsyncMock(),
deducted_max_cost=7_000,
reservation_snapshot=ReservationSnapshot(
release_id="rel-2",
key_hash="a" * 64,
billing_key_hash="a" * 64,
reserved_msats=7_000,
),
)
assert isinstance(cost, dict)
assert cost["total_msats"] == 0
assert cost["reason"] == "missing_usage"
assert cost["estimated"] is True
assert cost["error"]["code"] == "pricing_error"
# ---------------------------------------------------------------------------
# X-Cashu: estimator wiring
# ---------------------------------------------------------------------------
def test_normalize_usage_rejects_none():
assert normalize_usage(None) is None
+2 -9
View File
@@ -77,20 +77,13 @@ def _usage_response() -> dict[str, Any]:
async def test_unusable_token_rate_never_charges_the_reservation(
bad_rate: float,
) -> None:
"""An unusable configured rate must not turn authorization into usage.
Under the default ``charge_max`` policy the request is still billed the
pre-authorized ceiling (never MORE than it), with the raw token counts
preserved for dashboards. A zero rate remains a price (see
``test_a_rate_of_zero_is_billed_as_free_not_as_missing``).
"""
"""An unusable configured rate must not turn authorization into usage."""
model = _model(Pricing(prompt=bad_rate, completion=1.0))
cost = await calculate_cost(_usage_response(), max_cost=1234, model_obj=model)
assert isinstance(cost, MaxCostData)
assert cost.total_msats == 1234
assert cost.reason == "missing_usage"
assert cost.total_msats == 0
assert (cost.input_tokens, cost.output_tokens) == (1000, 500)
+141
View File
@@ -0,0 +1,141 @@
"""X-Cashu billing when the upstream omits usage.
The local token estimator bills from the request body and the generated text.
When nothing can be estimated the prepayment is refunded in full.
"""
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.upstream.base import BaseUpstreamProvider # noqa: E402
REQUEST_BODY = json.dumps(
{"model": "gpt-4o", "messages": [{"role": "user", "content": "Tell me a joke"}]}
).encode()
def _sse(events: list[dict[str, Any]]) -> httpx.Response:
body = "".join(f"data: {json.dumps(e)}\n\n" for e in events) + "data: [DONE]\n\n"
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, content=body.encode()
)
def _json(payload: dict[str, Any]) -> httpx.Response:
return httpx.Response(
200, headers={"content-type": "application/json"}, content=json.dumps(payload)
)
async def _settle(
response: httpx.Response,
*,
responses_api: bool = False,
request_body: bytes | None = REQUEST_BODY,
) -> tuple[Any, AsyncMock, AsyncMock]:
provider = BaseUpstreamProvider(base_url="http://test", api_key="test-key")
get_cost = AsyncMock(side_effect=provider.get_x_cashu_cost)
send_refund = AsyncMock(return_value="cashuBrefund")
handler = (
provider.handle_x_cashu_responses_completion
if responses_api
else provider.handle_x_cashu_chat_completion
)
with (
patch.object(provider, "get_x_cashu_cost", new=get_cost),
patch.object(provider, "send_refund", new=send_refund),
):
result = await handler(
response=response,
amount=10_000,
unit="msat",
max_cost_for_model=9_000,
mint=None,
request_body=request_body,
)
return result, get_cost, send_refund
def _billed_usage(get_cost: AsyncMock) -> dict[str, Any] | None:
assert get_cost.await_args is not None
return get_cost.await_args.args[0].get("usage")
@pytest.mark.asyncio
async def test_streaming_chat_without_usage_bills_from_estimate() -> None:
events = [
{"model": "gpt-4o", "choices": [{"delta": {"content": "Why did the "}}]},
{"model": "gpt-4o", "choices": [{"delta": {"content": "chicken cross"}}]},
]
_, get_cost, _ = await _settle(_sse(events))
usage = _billed_usage(get_cost)
assert usage is not None
assert usage["input_tokens"] > 0
assert usage["output_tokens"] > 0
assert usage["estimated"] is True
@pytest.mark.asyncio
async def test_streaming_chat_without_text_refunds_everything() -> None:
_, get_cost, send_refund = await _settle(
_sse([{"model": "gpt-4o"}]), request_body=None
)
assert _billed_usage(get_cost) is None
assert send_refund.await_args is not None
assert send_refund.await_args.args[0] == 10_000
@pytest.mark.asyncio
async def test_non_streaming_chat_without_usage_bills_from_estimate() -> None:
payload = {
"model": "gpt-4o",
"choices": [
{"message": {"role": "assistant", "content": "To get to the other side."}}
],
}
_, get_cost, _ = await _settle(_json(payload))
usage = _billed_usage(get_cost)
assert usage is not None
assert usage["output_tokens"] > 0
assert usage["estimated"] is True
@pytest.mark.asyncio
async def test_streaming_responses_without_usage_bills_from_estimate() -> None:
events = [
{"type": "response.created", "response": {"model": "gpt-5-mini"}},
{"type": "response.output_text.delta", "delta": "Why did the chicken"},
{"type": "response.output_text.done", "text": "Why did the chicken"},
]
_, get_cost, _ = await _settle(_sse(events), responses_api=True)
usage = _billed_usage(get_cost)
assert usage is not None
assert usage["output_tokens"] > 0
assert usage["estimated"] is True
@pytest.mark.asyncio
async def test_non_streaming_responses_without_usage_bills_from_estimate() -> None:
payload = {
"model": "gpt-5-mini",
"output": [
{"type": "message", "content": [{"type": "output_text", "text": "Hi"}]}
],
}
_, get_cost, _ = await _settle(_json(payload), responses_api=True)
usage = _billed_usage(get_cost)
assert usage is not None
assert usage["output_tokens"] > 0
@@ -187,8 +187,6 @@ async def test_multiline_data_payload_is_parsed_and_reframed() -> None:
@pytest.mark.asyncio
async def test_missing_usage_refunds_instead_of_charging_authorized_max() -> None:
"""No usage in the stream: ``missing_usage_policy`` (default charge_max)
bills the pre-authorized ceiling and refunds only the difference."""
chunks = [
b'data: {"type":"response.created","response":{"model":"gpt-5-mini"}}\r\n\r\n',
b"data: [DONE]\r\n\r\n",
@@ -200,10 +198,9 @@ async def test_missing_usage_refunds_instead_of_charging_authorized_max() -> Non
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 send_refund.await_args.args[0] == 10_000
assert response.headers["x-cashu"] == "cashuBrefundtoken0123456789"
assert response.headers["x-routstr-cost-msats"] == "9000"
assert response.headers["x-routstr-cost-estimated"] == "true"
assert response.headers["x-routstr-cost-msats"] == "0"
@pytest.mark.asyncio
@@ -218,7 +215,7 @@ async def test_malformed_events_do_not_retain_whole_token() -> None:
)
assert send_refund.await_args is not None
assert send_refund.await_args.args[0] == 10_000 - 9_000
assert send_refund.await_args.args[0] == 10_000
body = await _collect(response)
assert b"\\n" not in body
assert body.endswith(b"\n\n")