From ac40d70c3ac3a235f0329f7b6596344e01f9187b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 7 Sep 2026 00:30:05 +0200 Subject: [PATCH] clean up --- routstr/auth.py | 143 ++------- routstr/core/settings.py | 8 - routstr/payment/cost_calculation.py | 89 ++---- routstr/upstream/base.py | 300 +++++++----------- tests/integration/test_payment_invariants.py | 20 +- tests/unit/test_cost_error_after_delivery.py | 47 +++ tests/unit/test_missing_usage_policy.py | 216 ------------- tests/unit/test_pricing_rate_validation.py | 11 +- tests/unit/test_x_cashu_missing_usage.py | 141 ++++++++ .../test_x_cashu_responses_streaming_sse.py | 9 +- 10 files changed, 365 insertions(+), 619 deletions(-) create mode 100644 tests/unit/test_cost_error_after_delivery.py delete mode 100644 tests/unit/test_missing_usage_policy.py create mode 100644 tests/unit/test_x_cashu_missing_usage.py diff --git a/routstr/auth.py b/routstr/auth.py index fb573835..bf29e86e 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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") diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 8005c453..014a5795 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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 diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 7592958f..b4b0ff6b 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -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, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 0535dcbd..2ea06e0b 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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 ) diff --git a/tests/integration/test_payment_invariants.py b/tests/integration/test_payment_invariants.py index 6209aff9..3d5460dd 100644 --- a/tests/integration/test_payment_invariants.py +++ b/tests/integration/test_payment_invariants.py @@ -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 diff --git a/tests/unit/test_cost_error_after_delivery.py b/tests/unit/test_cost_error_after_delivery.py new file mode 100644 index 00000000..5540836d --- /dev/null +++ b/tests/unit/test_cost_error_after_delivery.py @@ -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 diff --git a/tests/unit/test_missing_usage_policy.py b/tests/unit/test_missing_usage_policy.py deleted file mode 100644 index 81cbf553..00000000 --- a/tests/unit/test_missing_usage_policy.py +++ /dev/null @@ -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 diff --git a/tests/unit/test_pricing_rate_validation.py b/tests/unit/test_pricing_rate_validation.py index 4f33aacf..96b478d1 100644 --- a/tests/unit/test_pricing_rate_validation.py +++ b/tests/unit/test_pricing_rate_validation.py @@ -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) diff --git a/tests/unit/test_x_cashu_missing_usage.py b/tests/unit/test_x_cashu_missing_usage.py new file mode 100644 index 00000000..19f5ae34 --- /dev/null +++ b/tests/unit/test_x_cashu_missing_usage.py @@ -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 diff --git a/tests/unit/test_x_cashu_responses_streaming_sse.py b/tests/unit/test_x_cashu_responses_streaming_sse.py index 79769ae3..08154e40 100644 --- a/tests/unit/test_x_cashu_responses_streaming_sse.py +++ b/tests/unit/test_x_cashu_responses_streaming_sse.py @@ -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")