From 5c6871ad6ed0701f4a0885c80c0588606f98c4bd Mon Sep 17 00:00:00 2001 From: Ashen <310210685+ashen0x@users.noreply.github.com> Date: Sun, 27 Sep 2026 19:17:48 +0530 Subject: [PATCH] feat: record completed request usage after settlement --- ...8e4a1f2b3d5_add_terminal_outcome_ledger.py | 133 +++ routstr/auth.py | 77 +- routstr/core/db.py | 112 ++- routstr/core/terminal_outcome_writer.py | 756 +++++++++++++++++ routstr/core/terminal_outcomes.py | 150 ++++ routstr/payment/cost_calculation.py | 88 +- routstr/payment/usage.py | 108 +++ routstr/upstream/base.py | 600 ++++++++++++- routstr/upstream/ehbp.py | 253 +++++- tests/integration/test_failover_billing.py | 9 +- tests/unit/test_cost_calculation_caching.py | 67 +- tests/unit/test_ehbp_finalize_payment.py | 162 ++++ tests/unit/test_messages_litellm_dispatch.py | 31 +- .../test_streaming_billing_finalization.py | 380 ++++++++- tests/unit/test_streaming_sse_providers.py | 66 +- tests/unit/test_terminal_outcome_migration.py | 190 +++++ tests/unit/test_terminal_outcomes.py | 790 ++++++++++++++++++ tests/unit/test_tinfoil_integration.py | 40 +- tests/unit/test_usage_normalization.py | 73 +- tests/unit/test_x_cashu_cost_sats.py | 170 +++- tests/unit/test_x_cashu_missing_usage.py | 23 + .../test_x_cashu_responses_streaming_sse.py | 35 +- 22 files changed, 4185 insertions(+), 128 deletions(-) create mode 100644 migrations/versions/c8e4a1f2b3d5_add_terminal_outcome_ledger.py create mode 100644 routstr/core/terminal_outcome_writer.py create mode 100644 routstr/core/terminal_outcomes.py create mode 100644 tests/unit/test_terminal_outcome_migration.py create mode 100644 tests/unit/test_terminal_outcomes.py diff --git a/migrations/versions/c8e4a1f2b3d5_add_terminal_outcome_ledger.py b/migrations/versions/c8e4a1f2b3d5_add_terminal_outcome_ledger.py new file mode 100644 index 00000000..8de0522e --- /dev/null +++ b/migrations/versions/c8e4a1f2b3d5_add_terminal_outcome_ledger.py @@ -0,0 +1,133 @@ +"""add terminal outcome ledger + +Revision ID: c8e4a1f2b3d5 +Revises: a73d19b6c204 +Create Date: 2026-08-31 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "c8e4a1f2b3d5" +down_revision = "a73d19b6c204" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "terminal_outcome_epochs", + sa.Column("epoch", sa.BigInteger(), nullable=False), + sa.Column("coverage_start_day", sa.Date(), nullable=False), + sa.Column("coverage_end_day", sa.Date(), nullable=True), + sa.Column("current_slot", sa.Integer(), nullable=True), + sa.CheckConstraint( + "epoch >= 0", + name="ck_terminal_outcome_epochs_nonnegative", + ), + sa.CheckConstraint( + "current_slot IS NULL OR current_slot = 1", + name="ck_terminal_outcome_epochs_current_slot", + ), + sa.CheckConstraint( + "(current_slot = 1 AND coverage_end_day IS NULL) OR " + "(current_slot IS NULL AND coverage_end_day IS NOT NULL)", + name="ck_terminal_outcome_epochs_state", + ), + sa.PrimaryKeyConstraint("epoch"), + sa.UniqueConstraint("current_slot", name="uq_terminal_outcome_epochs_current"), + ) + op.create_table( + "terminal_outcome_writer_runs", + sa.Column("run_id", sa.String(), nullable=False), + sa.Column("status", sa.String(), nullable=False), + sa.Column("started_at_ms", sa.BigInteger(), nullable=False), + sa.Column("heartbeat_at_ms", sa.BigInteger(), nullable=False), + sa.Column("flushed_through_ms", sa.BigInteger(), nullable=True), + sa.Column("closed_at_ms", sa.BigInteger(), nullable=True), + sa.Column("loss_day", sa.Date(), nullable=True), + sa.CheckConstraint( + "status IN ('active', 'degraded', 'clean', 'lost', 'recovered')", + name="ck_terminal_outcome_writer_runs_status", + ), + sa.CheckConstraint( + "started_at_ms >= 0 AND heartbeat_at_ms >= 0 " + "AND (closed_at_ms IS NULL OR closed_at_ms >= 0)", + name="ck_terminal_outcome_writer_runs_nonnegative", + ), + sa.CheckConstraint( + "(status IN ('active', 'degraded') AND closed_at_ms IS NULL) OR " + "(status IN ('clean', 'lost', 'recovered') " + "AND closed_at_ms IS NOT NULL)", + name="ck_terminal_outcome_writer_runs_state", + ), + sa.CheckConstraint( + "status NOT IN ('lost', 'recovered') OR loss_day IS NOT NULL", + name="ck_terminal_outcome_writer_runs_lost_day", + ), + sa.PrimaryKeyConstraint("run_id"), + ) + op.create_index( + "ix_terminal_outcome_writer_runs_status_heartbeat", + "terminal_outcome_writer_runs", + ["status", "heartbeat_at_ms"], + unique=False, + ) + op.create_table( + "terminal_outcomes", + sa.Column("outcome_id", sa.String(), nullable=False), + sa.Column("terminal_at_ms", sa.BigInteger(), nullable=False), + sa.Column("terminal_day", sa.Date(), nullable=False), + sa.Column("model_identifier", sa.String(), nullable=True), + sa.Column("served_model_identifier", sa.String(), nullable=True), + sa.Column("pricing_source", sa.String(), nullable=True), + sa.Column( + "input_source", sa.String(), nullable=False, server_default="missing" + ), + sa.Column( + "output_source", sa.String(), nullable=False, server_default="missing" + ), + sa.Column( + "cache_read_source", sa.String(), nullable=False, server_default="missing" + ), + sa.Column( + "cache_creation_source", + sa.String(), + nullable=False, + server_default="missing", + ), + sa.Column("input_tokens", sa.BigInteger(), nullable=False), + sa.Column("output_tokens", sa.BigInteger(), nullable=False), + sa.Column("cache_read_input_tokens", sa.BigInteger(), nullable=False), + sa.Column("cache_creation_input_tokens", sa.BigInteger(), nullable=False), + sa.Column("revenue_msats", sa.BigInteger(), nullable=False), + sa.CheckConstraint( + "terminal_at_ms >= 0 AND input_tokens >= 0 " + "AND output_tokens >= 0 AND cache_read_input_tokens >= 0 " + "AND cache_creation_input_tokens >= 0 AND revenue_msats >= 0", + name="ck_terminal_outcomes_nonnegative", + ), + sa.PrimaryKeyConstraint("outcome_id"), + ) + op.create_index( + "ix_terminal_outcomes_terminal_day_terminal_at_ms", + "terminal_outcomes", + ["terminal_day", "terminal_at_ms"], + unique=False, + ) + + +def downgrade() -> None: + op.drop_index( + "ix_terminal_outcomes_terminal_day_terminal_at_ms", + table_name="terminal_outcomes", + ) + op.drop_table("terminal_outcomes") + op.drop_index( + "ix_terminal_outcome_writer_runs_status_heartbeat", + table_name="terminal_outcome_writer_runs", + ) + op.drop_table("terminal_outcome_writer_runs") + op.drop_table("terminal_outcome_epochs") diff --git a/routstr/auth.py b/routstr/auth.py index 99c1d18e..4ccab255 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -5,7 +5,7 @@ import time import uuid from contextlib import suppress from contextvars import ContextVar -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, Optional from fastapi import HTTPException @@ -22,12 +22,19 @@ from .core.db import ( create_session, ) from .core.settings import settings +from .core.terminal_outcomes import ( + TerminalOutcomeContext, + mark_terminal_outcome_loss, + record_terminal_outcome, +) from .payment.cost_calculation import ( CostData, CostDataError, MaxCostData, calculate_cost, + unpriced_cost, ) +from .payment.usage import UsageFieldPresence, usage_field_presence from .redemption_cache import ( TERMINAL_REDEMPTION_CODES, CachedRedemptionFailure, @@ -1097,6 +1104,8 @@ async def release_reservation( snapshot: ReservationSnapshot, session: AsyncSession, reserved_msats: int, + *, + idempotent_success: bool = True, ) -> bool: """Release one durable reservation exactly once without charging.""" if reserved_msats <= 0 or reserved_msats != snapshot.reserved_msats: @@ -1105,7 +1114,7 @@ async def release_reservation( snapshot, session, decrement_requests=False, - idempotent_success=True, + idempotent_success=idempotent_success, ) @@ -1189,6 +1198,9 @@ async def _adjust_payment_for_tokens( model_obj: "Model | None" = None, provider_fee: float | None = None, reservation_snapshot: ReservationSnapshot | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, + usage_presence: UsageFieldPresence | None = None, + terminal_usage: dict | None = None, ) -> dict: """ Adjusts the payment based on token usage in the response. @@ -1201,6 +1213,7 @@ async def _adjust_payment_for_tokens( The response's usage object is normalized with the default union parser in ``calculate_cost``. + ``terminal_usage`` supplies reported stats without changing that billing input. """ billing_key = key reservation = reservation_snapshot or await get_reservation_snapshot(key, session) @@ -1265,8 +1278,48 @@ async def _adjust_payment_for_tokens( extra={"error": str(e), "fee_msats": fee_msats}, ) + async def _commit_settlement(cost: CostData, revenue_msats: int) -> None: + try: + await session.commit() + except BaseException: + if terminal_outcome is not None: + mark_terminal_outcome_loss("prepaid_commit_ambiguous") + raise + if terminal_outcome is not None: + recorded_usage = response_data.get("usage") + context = replace( + terminal_outcome, + pricing_source=cost.pricing_source, + input_source=cost.input_source, + output_source=cost.output_source, + cache_read_source=cost.cache_read_source, + cache_creation_source=cost.cache_creation_source, + ) + if terminal_usage is not None: + recorded_usage = {**(recorded_usage or {}), **terminal_usage} + presence = UsageFieldPresence( + input_source=cost.input_source, + output_source=cost.output_source, + cache_read_source=cost.cache_read_source, + cache_creation_source=cost.cache_creation_source, + ).merged(usage_field_presence(terminal_usage)) + context = replace(context, **presence.sources_dict()) + record_terminal_outcome( + context, + input_tokens=cost.input_tokens, + output_tokens=cost.output_tokens, + cache_read_input_tokens=cost.cache_read_input_tokens, + cache_creation_input_tokens=cost.cache_creation_input_tokens, + revenue_msats=revenue_msats, + usage=recorded_usage, + ) + calculated_cost = await calculate_cost( - response_data, deducted_max_cost, model_obj, provider_fee + response_data, + deducted_max_cost, + model_obj, + provider_fee, + usage_presence, ) if isinstance(calculated_cost, CostDataError): # Content was already served, so release instead of raising a 400. @@ -1279,9 +1332,7 @@ async def _adjust_payment_for_tokens( "error_code": calculated_cost.code, }, ) - calculated_cost = MaxCostData( - base_msats=0, input_msats=0, output_msats=0, total_msats=0 - ) + calculated_cost = unpriced_cost(response_data, usage_presence) if not await _claim_reservation_for_charge(reservation, session): # A prior charge or release already owns this reservation. Returning @@ -1309,7 +1360,7 @@ async def _adjust_payment_for_tokens( charge_msats=cost.total_msats, ) if charged: - await session.commit() + await _commit_settlement(cost, cost.total_msats) await _stop_reservation_heartbeat(reservation.release_id) if not charged: logger.error( @@ -1410,7 +1461,7 @@ async def _adjust_payment_for_tokens( await release_reservation_only() return cost.dict() - await session.commit() + await _commit_settlement(cost, total_cost_msats) await _stop_reservation_heartbeat(reservation.release_id) cost.charged_msats = total_cost_msats await session.refresh(billing_key) @@ -1493,7 +1544,7 @@ async def _adjust_payment_for_tokens( await session.rollback() raise RuntimeError("Could not atomically finalize cost overrun") - await session.commit() + await _commit_settlement(cost, actual_charge_msats) await _stop_reservation_heartbeat(reservation.release_id) await session.refresh(billing_key) @@ -1559,7 +1610,7 @@ async def _adjust_payment_for_tokens( charge_msats=total_cost_msats, ) if charged: - await session.commit() + await _commit_settlement(cost, total_cost_msats) await _stop_reservation_heartbeat(reservation.release_id) if not charged: @@ -1625,6 +1676,9 @@ async def adjust_payment_for_tokens( model_obj: "Model | None" = None, provider_fee: float | None = None, reservation_snapshot: ReservationSnapshot | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, + usage_presence: UsageFieldPresence | None = None, + terminal_usage: dict | None = None, ) -> dict: """Settle payment while exposing latency for every import path.""" started = time.perf_counter() @@ -1639,6 +1693,9 @@ async def adjust_payment_for_tokens( model_obj, provider_fee, reservation_snapshot, + terminal_outcome, + usage_presence, + terminal_usage, ) succeeded = True return result diff --git a/routstr/core/db.py b/routstr/core/db.py index 84307a7e..591d6a7c 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -6,13 +6,25 @@ import sqlite3 import time import uuid from contextlib import asynccontextmanager +from datetime import date from enum import Enum from typing import AsyncGenerator from alembic import command from alembic.config import Config from alembic.util.exc import CommandError -from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text +from sqlalchemy import ( + BigInteger, + CheckConstraint, + Date, + Index, + UniqueConstraint, + case, + delete, + event, + or_, + text, +) from sqlalchemy.engine import make_url from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio import AsyncEngine @@ -796,6 +808,104 @@ class ReservationRelease(SQLModel, table=True): # type: ignore created_at: int = Field(default_factory=lambda: int(time.time())) +class TerminalOutcome(SQLModel, table=True): # type: ignore + __tablename__ = "terminal_outcomes" + __table_args__ = ( + Index( + "ix_terminal_outcomes_terminal_day_terminal_at_ms", + "terminal_day", + "terminal_at_ms", + ), + CheckConstraint( + "terminal_at_ms >= 0 AND input_tokens >= 0 " + "AND output_tokens >= 0 AND cache_read_input_tokens >= 0 " + "AND cache_creation_input_tokens >= 0 AND revenue_msats >= 0", + name="ck_terminal_outcomes_nonnegative", + ), + ) + + outcome_id: str = Field(primary_key=True) + terminal_at_ms: int = Field(sa_type=BigInteger) + terminal_day: date = Field(sa_type=Date) + model_identifier: str | None = Field(default=None, nullable=True) + served_model_identifier: str | None = Field(default=None, nullable=True) + pricing_source: str | None = Field(default=None, nullable=True) + input_source: str = Field(default="missing") + output_source: str = Field(default="missing") + cache_read_source: str = Field(default="missing") + cache_creation_source: str = Field(default="missing") + input_tokens: int = Field(sa_type=BigInteger) + output_tokens: int = Field(sa_type=BigInteger) + cache_read_input_tokens: int = Field(sa_type=BigInteger) + cache_creation_input_tokens: int = Field(sa_type=BigInteger) + revenue_msats: int = Field(sa_type=BigInteger) + + +class TerminalOutcomeEpoch(SQLModel, table=True): # type: ignore + __tablename__ = "terminal_outcome_epochs" + __table_args__ = ( + UniqueConstraint("current_slot", name="uq_terminal_outcome_epochs_current"), + CheckConstraint( + "epoch >= 0", + name="ck_terminal_outcome_epochs_nonnegative", + ), + CheckConstraint( + "current_slot IS NULL OR current_slot = 1", + name="ck_terminal_outcome_epochs_current_slot", + ), + CheckConstraint( + "(current_slot = 1 AND coverage_end_day IS NULL) OR " + "(current_slot IS NULL AND coverage_end_day IS NOT NULL)", + name="ck_terminal_outcome_epochs_state", + ), + ) + + epoch: int = Field(primary_key=True, sa_type=BigInteger) + coverage_start_day: date = Field(sa_type=Date) + coverage_end_day: date | None = Field(default=None, nullable=True, sa_type=Date) + current_slot: int | None = Field(default=None, nullable=True) + + +class TerminalOutcomeWriterRun(SQLModel, table=True): # type: ignore + __tablename__ = "terminal_outcome_writer_runs" + __table_args__ = ( + Index( + "ix_terminal_outcome_writer_runs_status_heartbeat", + "status", + "heartbeat_at_ms", + ), + CheckConstraint( + "status IN ('active', 'degraded', 'clean', 'lost', 'recovered')", + name="ck_terminal_outcome_writer_runs_status", + ), + CheckConstraint( + "started_at_ms >= 0 AND heartbeat_at_ms >= 0 " + "AND (closed_at_ms IS NULL OR closed_at_ms >= 0)", + name="ck_terminal_outcome_writer_runs_nonnegative", + ), + CheckConstraint( + "(status IN ('active', 'degraded') AND closed_at_ms IS NULL) OR " + "(status IN ('clean', 'lost', 'recovered') " + "AND closed_at_ms IS NOT NULL)", + name="ck_terminal_outcome_writer_runs_state", + ), + CheckConstraint( + "status NOT IN ('lost', 'recovered') OR loss_day IS NOT NULL", + name="ck_terminal_outcome_writer_runs_lost_day", + ), + ) + + run_id: str = Field(primary_key=True) + status: str + started_at_ms: int = Field(sa_type=BigInteger) + heartbeat_at_ms: int = Field(sa_type=BigInteger) + flushed_through_ms: int | None = Field( + default=None, nullable=True, sa_type=BigInteger + ) + closed_at_ms: int | None = Field(default=None, nullable=True, sa_type=BigInteger) + loss_day: date | None = Field(default=None, nullable=True, sa_type=Date) + + class RoutstrFee(SQLModel, table=True): # type: ignore __tablename__ = "routstr_fees" id: int = Field(default=1, primary_key=True) diff --git a/routstr/core/terminal_outcome_writer.py b/routstr/core/terminal_outcome_writer.py new file mode 100644 index 00000000..1586d82d --- /dev/null +++ b/routstr/core/terminal_outcome_writer.py @@ -0,0 +1,756 @@ +from __future__ import annotations + +import asyncio +import time +import uuid +from collections.abc import Callable +from dataclasses import asdict, dataclass, fields +from datetime import UTC, date, datetime, timedelta +from enum import Enum +from typing import AsyncContextManager + +from sqlalchemy.exc import IntegrityError +from sqlmodel import col, select, update +from sqlmodel.ext.asyncio.session import AsyncSession + +from .db import ( + TerminalOutcome, + TerminalOutcomeEpoch, + TerminalOutcomeWriterRun, + create_session, +) +from .logging import get_logger + +logger = get_logger(__name__) + +_QUEUE_SIZE = 4096 +_RETRY_SECONDS = 1.0 +_HEARTBEAT_SECONDS = 15.0 +_LEASE_TIMEOUT_SECONDS = 90.0 + +SessionFactory = Callable[[], AsyncContextManager[AsyncSession]] +Clock = Callable[[], int] + + +@dataclass(frozen=True) +class _QueuedOutcome: + outcome_id: str + terminal_at_ms: int + terminal_day: date + model_identifier: str | None + served_model_identifier: str | None + pricing_source: str | None + input_source: str + output_source: str + cache_read_source: str + cache_creation_source: str + input_tokens: int + output_tokens: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + revenue_msats: int + + +class _PersistResult(Enum): + STORED = "stored" + CONFLICT = "conflict" + RETRY = "retry" + + +def _same_outcome(existing: TerminalOutcome, queued: _QueuedOutcome) -> bool: + return all( + getattr(existing, field.name) == getattr(queued, field.name) + for field in fields(queued) + ) + + +def _valid_nonnegative_int(value: object, maximum: int | None = None) -> bool: + return type(value) is int and value >= 0 and (maximum is None or value <= maximum) + + +def _utc_day_from_ms(timestamp_ms: int) -> date: + return datetime.fromtimestamp(timestamp_ms / 1000, tz=UTC).date() + + +async def _current_epoch(session: AsyncSession) -> TerminalOutcomeEpoch | None: + return ( + await session.exec( + select(TerminalOutcomeEpoch).where( + col(TerminalOutcomeEpoch.current_slot) == 1 + ) + ) + ).first() + + +class TerminalOutcomeWriter: + def __init__( + self, + *, + session_factory: SessionFactory = create_session, + queue_size: int = _QUEUE_SIZE, + retry_seconds: float = _RETRY_SECONDS, + heartbeat_seconds: float = _HEARTBEAT_SECONDS, + lease_timeout_seconds: float = _LEASE_TIMEOUT_SECONDS, + clock: Clock | None = None, + ) -> None: + if queue_size <= 0: + raise ValueError("Terminal outcome queue size must be positive") + if retry_seconds <= 0 or heartbeat_seconds <= 0: + raise ValueError("Terminal outcome retry intervals must be positive") + if lease_timeout_seconds <= heartbeat_seconds: + raise ValueError("Terminal outcome lease must exceed its heartbeat") + self._session_factory = session_factory + self._queue_size = queue_size + self._retry_seconds = retry_seconds + self._heartbeat_seconds = heartbeat_seconds + self._lease_timeout_ms = int(lease_timeout_seconds * 1000) + self._clock = clock or (lambda: int(time.time() * 1000)) + self._queue: asyncio.Queue[_QueuedOutcome] | None = None + self._task: asyncio.Task[None] | None = None + self._idle: asyncio.Event | None = None + self._wake: asyncio.Event | None = None + self._run_id: str | None = None + self._epoch = 0 + self._enabled = False + self._accepting = False + self._stopping = False + self._loss_pending = False + self._loss_day: date | None = None + + @property + def epoch(self) -> int: + return self._epoch + + @property + def loss_pending(self) -> bool: + return self._loss_pending + + @property + def loss_day(self) -> date | None: + return self._loss_day + + @property + def running(self) -> bool: + return self._task is not None and not self._task.done() + + async def start(self, *, serving: bool = False) -> bool: + """``serving`` means this process settled requests while not collecting.""" + if self.running: + return True + previous_loss_day = self._loss_day + self._enabled = True + self._accepting = False + self._stopping = False + self._loss_pending = False + self._loss_day = None + self._queue = asyncio.Queue(maxsize=self._queue_size) + self._idle = asyncio.Event() + self._idle.set() + self._wake = asyncio.Event() + try: + now = self._now_ms() + await self._ensure_epoch(now) + await self._recover_stale_runs(now) + if previous_loss_day is not None: + await self._rotate_epoch(previous_loss_day) + await self._recover_unattended_coverage(now) + if serving: + await self._void_served_coverage(now) + await self._create_run(now, "active") + except Exception: + self._loss_pending = True + self._loss_day = previous_loss_day or _utc_day_from_ms(self._now_ms()) + self._clear_runtime() + logger.critical("Terminal outcome writer could not start", exc_info=True) + return False + self._accepting = True + self._task = self._new_task() + return True + + async def stop(self, *, timeout: float = 5.0, close_coverage: bool = False) -> bool: + closed_day = _utc_day_from_ms(self._now_ms()) - timedelta(days=1) + if not self._enabled: + if close_coverage: + try: + await self._recover_unattended_coverage(self._now_ms()) + await self._close_coverage(closed_day) + except Exception: + logger.critical( + "Terminal outcome coverage closure failed", exc_info=True + ) + return False + return True + self._stopping = True + self._accepting = False + if self._wake is not None: + self._wake.set() + clean = False + drained = False + task = self._task + try: + if self._queue is not None: + await asyncio.wait_for(self._queue.join(), timeout=timeout) + if self._idle is not None: + await asyncio.wait_for(self._idle.wait(), timeout=timeout) + drained = not self._loss_pending and task is not None and not task.done() + except (TimeoutError, asyncio.CancelledError): + pass + except Exception: + logger.critical("Terminal outcome queue shutdown failed", exc_info=True) + if task is not None and not task.done(): + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + try: + if close_coverage: + await self._close_coverage(closed_day) + clean = drained and not self._loss_pending and await self._close_run() + except (TimeoutError, asyncio.CancelledError): + pass + except Exception: + logger.critical("Terminal outcome clean shutdown failed", exc_info=True) + self._clear_runtime() + return clean + + async def flush(self, *, timeout: float = 5.0) -> bool: + if not self.running or self._queue is None or self._idle is None: + return False + try: + await asyncio.wait_for(self._queue.join(), timeout=timeout) + await asyncio.wait_for(self._idle.wait(), timeout=timeout) + return not self._loss_pending and await self._touch_run(drained=True) + except (TimeoutError, asyncio.CancelledError): + return False + except Exception: + logger.error("Terminal outcome flush failed", exc_info=True) + return False + + def submit(self, outcome: _QueuedOutcome) -> bool: + if not self._enabled: + return False + if not self._accepting or not self.running or self._queue is None: + self.declare_loss( + "terminal outcome writer unavailable", outcome.terminal_day + ) + return False + try: + self._queue.put_nowait(outcome) + except asyncio.QueueFull: + self.declare_loss("terminal outcome queue full", outcome.terminal_day) + return False + if self._idle is not None: + self._idle.clear() + if self._wake is not None: + self._wake.set() + return True + + def declare_loss(self, reason: str, lost_day: date | None = None) -> None: + if not self._enabled: + return + day = lost_day or _utc_day_from_ms(self._now_ms()) + first_loss = not self._loss_pending + self._loss_pending = True + if self._loss_day is None or day < self._loss_day: + self._loss_day = day + self._accepting = False + if self._idle is not None: + self._idle.clear() + if self._wake is not None: + self._wake.set() + if first_loss: + logger.critical( + "Terminal outcome continuity lost", + extra={"reason": reason, "epoch": self._epoch}, + ) + + def _now_ms(self) -> int: + now = self._clock() + if not _valid_nonnegative_int(now): + raise ValueError("Terminal outcome clock returned an invalid value") + return now + + def _new_task(self) -> asyncio.Task[None]: + task = asyncio.create_task(self._run(), name="terminal-outcome-writer") + task.add_done_callback(self._writer_stopped) + return task + + def _clear_runtime(self) -> None: + self._task = None + self._queue = None + self._idle = None + self._wake = None + self._run_id = None + self._enabled = False + self._accepting = False + self._stopping = False + + async def _ensure_epoch(self, now: int) -> None: + while True: + async with self._session_factory() as session: + current = await _current_epoch(session) + if current is not None: + self._epoch = current.epoch + return + latest_result = await session.exec( + select(TerminalOutcomeEpoch).order_by( + col(TerminalOutcomeEpoch.epoch).desc() + ) + ) + latest = latest_result.first() + next_epoch = 0 if latest is None else latest.epoch + 1 + session.add( + TerminalOutcomeEpoch( + epoch=next_epoch, + coverage_start_day=_utc_day_from_ms(now) + timedelta(days=1), + current_slot=1, + ) + ) + try: + await session.commit() + except IntegrityError: + await session.rollback() + continue + self._epoch = next_epoch + return + + async def _create_run( + self, + now: int, + status: str, + lost_day: date | None = None, + ) -> None: + run_id = uuid.uuid4().hex + async with self._session_factory() as session: + session.add( + TerminalOutcomeWriterRun( + run_id=run_id, + status=status, + started_at_ms=now, + heartbeat_at_ms=now, + flushed_through_ms=now if status == "active" else None, + loss_day=lost_day, + ) + ) + await session.commit() + self._run_id = run_id + + async def _update_run(self, expected: tuple[str, ...], **values: object) -> bool: + if self._run_id is None: + return False + async with self._session_factory() as session: + result = await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeWriterRun) + .where(col(TerminalOutcomeWriterRun.run_id) == self._run_id) + .where(col(TerminalOutcomeWriterRun.status).in_(expected)) + .values(**values) + ) + await session.commit() + return bool(result.rowcount == 1) + + async def _close_coverage(self, closed_day: date) -> None: + async with self._session_factory() as session: + await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeEpoch) + .where(col(TerminalOutcomeEpoch.current_slot) == 1) + .values(coverage_end_day=closed_day, current_slot=None) + ) + await session.commit() + + async def _close_run(self) -> bool: + now = self._now_ms() + return await self._update_run( + ("active",), + status="clean", + heartbeat_at_ms=now, + closed_at_ms=now, + ) + + async def _touch_run(self, *, drained: bool = False) -> bool: + now = self._now_ms() + values = {"heartbeat_at_ms": now} + if drained and not self._loss_pending: + values["flushed_through_ms"] = now + touched = await self._update_run(("active", "degraded"), **values) + if not touched: + self.declare_loss("terminal outcome writer lease was lost") + return touched + + async def _mark_degraded(self, lost_day: date) -> None: + updated = await self._update_run( + ("active", "degraded"), + status="degraded", + loss_day=lost_day, + ) + if not updated: + await self._create_run(self._now_ms(), "degraded", lost_day) + + async def _activate_run(self) -> None: + now = self._now_ms() + updated = await self._update_run( + ("degraded",), + status="active", + heartbeat_at_ms=now, + flushed_through_ms=now, + loss_day=None, + ) + if not updated: + await self._create_run(now, "active") + + async def _claim_stale_runs(self, now: int) -> None: + cutoff = now - self._lease_timeout_ms + async with self._session_factory() as session: + statement = ( + select(TerminalOutcomeWriterRun) + .where(col(TerminalOutcomeWriterRun.status).in_(("active", "degraded"))) + .where(col(TerminalOutcomeWriterRun.heartbeat_at_ms) < cutoff) + ) + if self._run_id is not None: + statement = statement.where( + col(TerminalOutcomeWriterRun.run_id) != self._run_id + ) + stale_runs = (await session.exec(statement)).all() + claimed = 0 + for stale in stale_runs: + checkpoint_day = _utc_day_from_ms( + stale.flushed_through_ms or stale.started_at_ms + ) + lost_day = min(checkpoint_day, stale.loss_day or checkpoint_day) + result = await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeWriterRun) + .where(col(TerminalOutcomeWriterRun.run_id) == stale.run_id) + .where(col(TerminalOutcomeWriterRun.status) == stale.status) + .where( + col(TerminalOutcomeWriterRun.heartbeat_at_ms) + == stale.heartbeat_at_ms + ) + .values( + status="lost", + closed_at_ms=now, + loss_day=lost_day, + ) + ) + claimed += int(result.rowcount == 1) + await session.commit() + if claimed: + logger.critical( + "Stale terminal outcome writer lease detected", + extra={"stale_runs": claimed}, + ) + + async def _stage_rotation( + self, + session: AsyncSession, + current: TerminalOutcomeEpoch, + lost_day: date, + ) -> int | None: + transition = await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeEpoch) + .where(col(TerminalOutcomeEpoch.epoch) == current.epoch) + .where(col(TerminalOutcomeEpoch.current_slot) == 1) + .values( + coverage_end_day=lost_day - timedelta(days=1), + current_slot=None, + ) + ) + if transition.rowcount != 1: + return None + await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeEpoch) + .where(col(TerminalOutcomeEpoch.current_slot).is_(None)) + .where(col(TerminalOutcomeEpoch.coverage_end_day) >= lost_day) + .values( + coverage_end_day=lost_day - timedelta(days=1), + ) + ) + next_epoch = current.epoch + 1 + recovery_day = _utc_day_from_ms(self._now_ms()) + session.add( + TerminalOutcomeEpoch( + epoch=next_epoch, + coverage_start_day=max(lost_day, recovery_day) + timedelta(days=1), + current_slot=1, + ) + ) + return next_epoch + + async def _rotate_epoch(self, lost_day: date) -> None: + while True: + now = self._now_ms() + async with self._session_factory() as session: + current = await _current_epoch(session) + if current is None: + await self._ensure_epoch(now) + continue + next_epoch = await self._stage_rotation(session, current, lost_day) + if next_epoch is None: + await session.rollback() + continue + try: + await session.commit() + except IntegrityError: + await session.rollback() + continue + self._epoch = next_epoch + return + + async def _recover_pending_runs(self) -> None: + while True: + now = self._now_ms() + async with self._session_factory() as session: + pending = ( + await session.exec( + select(TerminalOutcomeWriterRun).where( + col(TerminalOutcomeWriterRun.status) == "lost" + ) + ) + ).all() + loss_days = [ + run.loss_day for run in pending if run.loss_day is not None + ] + if not loss_days: + return + pending_ids = [run.run_id for run in pending] + current = await _current_epoch(session) + if current is None: + await self._ensure_epoch(now) + continue + claimed = await session.exec( # type: ignore[call-overload] + update(TerminalOutcomeWriterRun) + .where(col(TerminalOutcomeWriterRun.status) == "lost") + .where(col(TerminalOutcomeWriterRun.run_id).in_(pending_ids)) + .values(status="recovered") + ) + if claimed.rowcount == 0: + await session.rollback() + continue + next_epoch = await self._stage_rotation( + session, current, min(loss_days) + ) + if next_epoch is None: + await session.rollback() + continue + try: + await session.commit() + except IntegrityError: + await session.rollback() + continue + self._epoch = next_epoch + return + + async def _recover_stale_runs(self, now: int) -> None: + await self._claim_stale_runs(now) + await self._recover_pending_runs() + + async def _recover_unattended_coverage(self, now: int) -> None: + while True: + async with self._session_factory() as session: + current = await _current_epoch(session) + if current is None or current.coverage_start_day > _utc_day_from_ms( + now + ): + return + live = ( + await session.exec( + select(TerminalOutcomeWriterRun.run_id) + .where( + col(TerminalOutcomeWriterRun.status).in_( + ("active", "degraded", "lost") + ) + ) + .limit(1) + ) + ).first() + if live is not None: + return + last_clean_close = ( + await session.exec( + select(TerminalOutcomeWriterRun.closed_at_ms) + .where(col(TerminalOutcomeWriterRun.status) == "clean") + .where(col(TerminalOutcomeWriterRun.closed_at_ms).is_not(None)) + .order_by(col(TerminalOutcomeWriterRun.closed_at_ms).desc()) + .limit(1) + ) + ).first() + # A clean stop proves its drained records, not later process uptime. + lost_day = current.coverage_start_day + if last_clean_close is not None: + lost_day = max(lost_day, _utc_day_from_ms(last_clean_close)) + next_epoch = await self._stage_rotation(session, current, lost_day) + if next_epoch is None: + await session.rollback() + continue + try: + await session.commit() + except IntegrityError: + await session.rollback() + continue + self._epoch = next_epoch + return + + async def _void_served_coverage(self, now: int) -> None: + async with self._session_factory() as session: + current = await _current_epoch(session) + # Coverage that already began cannot include what this worker missed. + if current is not None and current.coverage_start_day <= _utc_day_from_ms(now): + await self._rotate_epoch(current.coverage_start_day) + + async def _reconcile(self, queued: _QueuedOutcome) -> _PersistResult: + try: + async with self._session_factory() as session: + existing = await session.get(TerminalOutcome, queued.outcome_id) + except asyncio.CancelledError: + raise + except Exception: + return _PersistResult.RETRY + if existing is None: + return _PersistResult.RETRY + return ( + _PersistResult.STORED + if _same_outcome(existing, queued) + else _PersistResult.CONFLICT + ) + + async def _persist_once(self, queued: _QueuedOutcome) -> _PersistResult: + try: + async with self._session_factory() as session: + session.add(TerminalOutcome(**asdict(queued))) + await session.commit() + return _PersistResult.STORED + except asyncio.CancelledError: + raise + except Exception: + return await self._reconcile(queued) + + async def _persist(self, queued: _QueuedOutcome) -> _PersistResult | None: + logged = False + while not self._loss_pending: + result = await self._persist_once(queued) + if result is not _PersistResult.RETRY: + return result + if not logged: + logger.error( + "Terminal outcome persistence will be retried", + extra={"outcome_id": queued.outcome_id}, + ) + logged = True + try: + await self._touch_run() + await self._recover_stale_runs(self._now_ms()) + except Exception: + pass + await asyncio.sleep(self._retry_seconds) + self.declare_loss( + "in-flight terminal outcome abandoned after continuity loss", + queued.terminal_day, + ) + return None + + def _discard_queue(self) -> None: + if self._queue is None: + return + while True: + try: + queued = self._queue.get_nowait() + except asyncio.QueueEmpty: + return + else: + self.declare_loss( + "queued terminal outcome discarded after continuity loss", + queued.terminal_day, + ) + self._queue.task_done() + + async def _recover_loss(self) -> None: + self._discard_queue() + rotated_for: date | None = None + while self._loss_pending: + lost_day = self._loss_day or _utc_day_from_ms(self._now_ms()) + try: + await self._recover_stale_runs(self._now_ms()) + await self._mark_degraded(lost_day) + if rotated_for is None or lost_day < rotated_for: + await self._rotate_epoch(lost_day) + rotated_for = lost_day + await self._activate_run() + except asyncio.CancelledError: + raise + except Exception: + logger.critical( + "Terminal outcome continuity recovery failed", exc_info=True + ) + await asyncio.sleep(self._retry_seconds) + continue + if self._loss_day is not None and self._loss_day < lost_day: + continue + self._loss_pending = False + self._loss_day = None + if not self._stopping: + self._accepting = True + + async def _heartbeat(self) -> None: + try: + await self._touch_run( + drained=self._queue is not None and self._queue.empty() + ) + await self._recover_stale_runs(self._now_ms()) + except asyncio.CancelledError: + raise + except Exception: + logger.error("Terminal outcome heartbeat failed", exc_info=True) + + async def _run(self) -> None: + if self._queue is None or self._wake is None: + return + next_heartbeat = time.monotonic() + self._heartbeat_seconds + while True: + try: + if self._loss_pending: + await self._recover_loss() + self._wake.clear() + try: + queued = self._queue.get_nowait() + except asyncio.QueueEmpty: + if self._idle is not None: + self._idle.set() + try: + await asyncio.wait_for( + self._wake.wait(), timeout=self._heartbeat_seconds + ) + except TimeoutError: + await self._heartbeat() + next_heartbeat = time.monotonic() + self._heartbeat_seconds + continue + try: + result = await self._persist(queued) + if result is _PersistResult.CONFLICT: + self.declare_loss( + "conflicting terminal outcome idempotency key", + queued.terminal_day, + ) + finally: + self._queue.task_done() + if time.monotonic() >= next_heartbeat: + await self._heartbeat() + next_heartbeat = time.monotonic() + self._heartbeat_seconds + except asyncio.CancelledError: + raise + except Exception: + logger.critical("Terminal outcome writer failed", exc_info=True) + self.declare_loss("terminal outcome writer failed") + await asyncio.sleep(self._retry_seconds) + + def _writer_stopped(self, task: asyncio.Task[None]) -> None: + if self._stopping or task.cancelled() or not self._enabled: + return + try: + error = task.exception() + except asyncio.CancelledError: + return + self.declare_loss( + "terminal outcome writer stopped" + if error is None + else f"terminal outcome writer stopped: {type(error).__name__}" + ) + self._task = self._new_task() diff --git a/routstr/core/terminal_outcomes.py b/routstr/core/terminal_outcomes.py new file mode 100644 index 00000000..7ea8509a --- /dev/null +++ b/routstr/core/terminal_outcomes.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import time +from dataclasses import dataclass + +from .logging import get_logger +from .terminal_outcome_writer import ( + TerminalOutcomeWriter, + _QueuedOutcome, + _utc_day_from_ms, + _valid_nonnegative_int, +) + +logger = get_logger(__name__) + +# Far above any real request, and low enough that a day's sums stay JSON-safe. +_MAX_TOKENS = 2**31 - 1 +_MAX_REVENUE_MSATS = 2**40 + + +@dataclass(frozen=True) +class TerminalOutcomeContext: + outcome_id: str | None + model_identifier: str | None + + served_model_identifier: str | None = None + pricing_source: str | None = None + input_source: str | None = None + output_source: str | None = None + cache_read_source: str | None = None + cache_creation_source: str | None = None + + +terminal_outcome_writer = TerminalOutcomeWriter() + + +def record_terminal_outcome( + context: TerminalOutcomeContext, + *, + input_tokens: int, + output_tokens: int, + cache_read_input_tokens: int, + cache_creation_input_tokens: int, + revenue_msats: int, + terminal_at_ms: int | None = None, + usage: object = None, +) -> None: + """Submit a settled outcome without awaiting storage or raising. + + ``usage`` is the raw upstream usage. It replaces the token counts and is + parsed here, so a malformed value can only mark a gap. + """ + try: + if usage is not None: + from ..payment.usage import NormalizedUsage, normalize_usage + + counted = normalize_usage(usage) or NormalizedUsage() + input_tokens = counted.input_tokens + output_tokens = counted.output_tokens + cache_read_input_tokens = counted.cache_read_tokens + cache_creation_input_tokens = counted.cache_write_tokens + sources = { + name + "_source": getattr(context, name + "_source") or "missing" + for name in ("input", "output", "cache_read", "cache_creation") + } + tokens = ( + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + ) + timestamp = ( + terminal_at_ms if terminal_at_ms is not None else int(time.time() * 1000) + ) + terminal_day = _utc_day_from_ms(timestamp) + if ( + not context.outcome_id + or any( + source not in {"reported", "estimated", "missing"} + for source in sources.values() + ) + or any(not _valid_nonnegative_int(value, _MAX_TOKENS) for value in tokens) + or not _valid_nonnegative_int(revenue_msats, _MAX_REVENUE_MSATS) + or not _valid_nonnegative_int(timestamp) + ): + terminal_outcome_writer.declare_loss( + "invalid settled terminal outcome", terminal_day + ) + return + terminal_outcome_writer.submit( + _QueuedOutcome( + outcome_id=context.outcome_id, + terminal_at_ms=timestamp, + terminal_day=terminal_day, + model_identifier=context.model_identifier, + served_model_identifier=context.served_model_identifier, + pricing_source=context.pricing_source, + **sources, + input_tokens=input_tokens, + output_tokens=output_tokens, + cache_read_input_tokens=cache_read_input_tokens, + cache_creation_input_tokens=cache_creation_input_tokens, + revenue_msats=revenue_msats, + ) + ) + except BaseException: + try: + terminal_outcome_writer.declare_loss("terminal outcome submission raised") + logger.critical("Terminal outcome submission failed", exc_info=True) + except BaseException: + pass + + +def mark_terminal_outcome_loss(reason: str) -> None: + try: + terminal_outcome_writer.declare_loss(reason) + except BaseException: + pass + + +def cashu_retained_msats(amount: int, unit: str, refund_amount: int = 0) -> int | None: + try: + if not _valid_nonnegative_int(amount) or not _valid_nonnegative_int( + refund_amount + ): + raise ValueError("Cashu amounts must be nonnegative integers") + if refund_amount > amount: + raise ValueError("Cashu refund exceeds redeemed amount") + if unit == "msat": + multiplier = 1 + elif unit == "sat": + multiplier = 1000 + else: + raise ValueError(f"Unsupported Cashu unit: {unit}") + return (amount - refund_amount) * multiplier + except Exception: + mark_terminal_outcome_loss("invalid Cashu retained value") + return None + + +async def start_terminal_outcome_writer(*, serving: bool = False) -> bool: + return await terminal_outcome_writer.start(serving=serving) + + +async def stop_terminal_outcome_writer( + *, timeout: float = 5.0, close_coverage: bool = False +) -> bool: + return await terminal_outcome_writer.stop( + timeout=timeout, close_coverage=close_coverage + ) diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index b4b0ff6b..c718ccff 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -7,7 +7,12 @@ from ..core import get_logger from ..core.settings import settings from .price import sats_usd_price from .rates import coerce_rate, is_usable_rate -from .usage import normalize_usage, parse_token_count +from .usage import ( + UsageFieldPresence, + normalize_usage, + parse_token_count, + usage_field_presence, +) if TYPE_CHECKING: from .models import Model @@ -35,6 +40,12 @@ class CostData(BaseModel): cache_creation_input_tokens: int = 0 cache_read_msats: int = 0 cache_creation_msats: int = 0 + # Settlement-only metadata must not expand the client cost contract. + pricing_source: str = Field(default="missing", exclude=True) + input_source: str = Field(default="missing", exclude=True) + output_source: str = Field(default="missing", exclude=True) + cache_read_source: str = Field(default="missing", exclude=True) + cache_creation_source: str = Field(default="missing", exclude=True) # Actual debit after finalization; None means settlement has not run yet. charged_msats: int | None = None upstream_usd: float = Field(default=0.0, exclude=True) @@ -49,13 +60,17 @@ class CostDataError(BaseModel): code: str -def _empty_cost(cls: type[CostData] = CostData) -> CostData: +def _empty_cost( + cls: type[CostData] = CostData, + usage_presence: UsageFieldPresence | None = None, +) -> CostData: """Build an all-zero cost object — a full refund for an empty response. Shared by the two paths that must not bill: an upstream response with no usage data at all, and one that reports a USD cost but carries zero tokens in every bucket. """ + presence = usage_presence if usage_presence is not None else UsageFieldPresence() return cls( base_msats=0, input_msats=0, @@ -68,14 +83,35 @@ def _empty_cost(cls: type[CostData] = CostData) -> CostData: cache_creation_input_tokens=0, cache_read_msats=0, cache_creation_msats=0, + **presence.sources_dict(), ) +def unpriced_cost( + response_data: dict, usage_presence: UsageFieldPresence | None = None +) -> CostData: + raw_usage = response_data.get("usage") + presence = usage_presence + if presence is None or ( + isinstance(raw_usage, dict) and raw_usage.get("estimated") is True + ): + presence = usage_field_presence(raw_usage) + cost = _empty_cost(MaxCostData, presence) + usage = normalize_usage(raw_usage) + if usage is not None: + cost.input_tokens = usage.input_tokens + cost.output_tokens = usage.output_tokens + cost.cache_read_input_tokens = usage.cache_read_tokens + cost.cache_creation_input_tokens = usage.cache_write_tokens + return cost + + async def calculate_cost( response_data: dict, max_cost: int, model_obj: "Model | None" = None, provider_fee: float | None = None, + usage_presence: UsageFieldPresence | None = None, ) -> CostData | MaxCostData | CostDataError: """Calculate the cost of an API request based on token usage. @@ -91,6 +127,7 @@ async def calculate_cost( pricing already carries the fee baked in). Without it, the fee is re-derived from the response's model string, which yields the best-ranked provider's fee. + usage_presence: Raw presence retained when a stream rebuilt its usage. Returns: Cost data or error information @@ -107,7 +144,16 @@ async def calculate_cost( }, ) - usage = normalize_usage(response_data.get("usage")) + raw_usage = response_data.get("usage") + presence = ( + usage_presence + if usage_presence is not None + else usage_field_presence(raw_usage) + ) + usage = normalize_usage(raw_usage) + + if isinstance(raw_usage, dict) and raw_usage.get("estimated") is True: + presence = usage_field_presence(raw_usage) if usage is None: logger.warning( @@ -124,7 +170,7 @@ async def calculate_cost( else None, }, ) - return _empty_cost(MaxCostData) + return _empty_cost(MaxCostData, presence) usage_data = response_data.get("usage") or {} if not isinstance(usage_data, dict): @@ -158,7 +204,7 @@ async def calculate_cost( else None, }, ) - return _empty_cost() + return _empty_cost(usage_presence=presence) if input_tokens == 0 and output_tokens == 0: logger.warning( "Upstream reported a USD cost but no token counts — " @@ -191,9 +237,11 @@ async def calculate_cost( cache_pricing_rates: tuple[float, float, float, float] | None = None if cache_read_tokens > 0 or cache_creation_tokens > 0: try: - cache_pricing_rates = _get_pricing_rates( + reported_rates = _get_pricing_rates( response_data, model_obj, provider_fee ) + if reported_rates is not None: + cache_pricing_rates = reported_rates[:4] except ValueError: logger.warning( "Cache pricing unavailable for USD cost breakdown; " @@ -221,6 +269,7 @@ async def calculate_cost( response_data, provider_fee, cache_pricing_rates, + presence, ) except Exception as e: logger.warning( @@ -243,8 +292,15 @@ async def calculate_cost( output_rate = float(settings.fixed_per_1k_output_tokens) * 1000.0 cache_read_rate = input_rate cache_creation_rate = input_rate + pricing_source = "fixed" else: - input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates + ( + input_rate, + output_rate, + cache_read_rate, + cache_creation_rate, + pricing_source, + ) = pricing_rates # Truthiness is not the question: `NaN` and a negative rate are both truthy # and sailed past this gate into the token math, while a rate of zero is a @@ -278,6 +334,7 @@ async def calculate_cost( cache_creation_input_tokens=cache_creation_tokens, cache_read_msats=0, cache_creation_msats=0, + **presence.sources_dict(), ) return _calculate_from_tokens( @@ -290,6 +347,8 @@ async def calculate_cost( cache_read_rate, cache_creation_rate, response_data, + presence, + pricing_source, ) @@ -366,14 +425,14 @@ def _get_pricing_rates( response_data: dict, model_obj: "Model | None", provider_fee: float | None, -) -> tuple[float, float, float, float] | None: +) -> tuple[float, float, float, float, str] | None: """Get configured rates, falling back to LiteLLM's model cost map. The served ``model_obj`` (when the caller has it) is billed directly; otherwise the response's model string is resolved through the alias map, which yields the best-ranked candidate rather than the serving one. - Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate). + Returns rates and their source (configured or LiteLLM). ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( @@ -458,7 +517,7 @@ def _get_pricing_rates( "cache_write_price_msats_per_1k": mscw_1k, }, ) - return mspp_1k, mspc_1k, mscr_1k, mscw_1k + return mspp_1k, mspc_1k, mscr_1k, mscw_1k, source def _resolve_provider_fee(model_id: str) -> float: @@ -488,6 +547,7 @@ def _calculate_from_usd_cost( response_data: dict, provider_fee: float | None, pricing_rates: tuple[float, float, float, float] | None = None, + usage_presence: UsageFieldPresence | None = None, ) -> CostData: """Calculate cost from USD figures, deriving input/output split from tokens.""" if provider_fee is None: @@ -566,6 +626,7 @@ def _calculate_from_usd_cost( }, ) + presence = usage_presence if usage_presence is not None else UsageFieldPresence() return CostData( base_msats=0, input_msats=input_msats, @@ -579,6 +640,8 @@ def _calculate_from_usd_cost( cache_read_msats=cache_read_msats, cache_creation_msats=cache_creation_msats, upstream_usd=reported_usd, + pricing_source="reported_usd", + **presence.sources_dict(), ) @@ -592,6 +655,8 @@ def _calculate_from_tokens( cache_read_rate: float, cache_creation_rate: float, response_data: dict, + usage_presence: UsageFieldPresence | None = None, + pricing_source: str = "missing", ) -> CostData: """Calculate cost from token counts using pricing rates.""" calc_input_msats = round(input_tokens / 1000 * input_rate, 3) @@ -636,6 +701,7 @@ def _calculate_from_tokens( visible_output_msats = int(calc_output_msats) visible_input_msats = token_based_cost - visible_output_msats + presence = usage_presence if usage_presence is not None else UsageFieldPresence() return CostData( base_msats=0, input_msats=visible_input_msats, @@ -648,4 +714,6 @@ def _calculate_from_tokens( cache_creation_input_tokens=cache_creation_tokens, cache_read_msats=int(calc_cache_read_msats), cache_creation_msats=int(calc_cache_write_msats), + pricing_source=pricing_source, + **presence.sources_dict(), ) diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 08d675ad..41c61b80 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -38,9 +38,54 @@ names do not collide, so a single union parser is safe; a vendor whose fields would genuinely conflict needs a dedicated branch here. """ +import math +from dataclasses import dataclass +from typing import TypedDict + from pydantic.v1 import BaseModel +class UsageSources(TypedDict): + input_source: str + output_source: str + cache_read_source: str + cache_creation_source: str + + +@dataclass(frozen=True) +class UsageFieldPresence: + """Where each canonical usage count came from: reported, estimated or missing.""" + + input_source: str = "missing" + output_source: str = "missing" + cache_read_source: str = "missing" + cache_creation_source: str = "missing" + + def sources_dict(self) -> UsageSources: + return { + "input_source": self.input_source, + "output_source": self.output_source, + "cache_read_source": self.cache_read_source, + "cache_creation_source": self.cache_creation_source, + } + + def merged(self, other: "UsageFieldPresence") -> "UsageFieldPresence": + def stronger(a: str, b: str) -> str: + for label in ("reported", "estimated"): + if label in (a, b): + return label + return "missing" + + return UsageFieldPresence( + input_source=stronger(self.input_source, other.input_source), + output_source=stronger(self.output_source, other.output_source), + cache_read_source=stronger(self.cache_read_source, other.cache_read_source), + cache_creation_source=stronger( + self.cache_creation_source, other.cache_creation_source + ), + ) + + class NormalizedUsage(BaseModel): """Canonical token usage: input_tokens never includes cached tokens.""" @@ -66,6 +111,69 @@ def parse_token_count(value: object) -> int: return 0 +def _is_parseable_token_count(value: object) -> bool: + """Return whether a value is a usable non-negative token count.""" + if isinstance(value, bool): + return False + if isinstance(value, int): + return value >= 0 + if isinstance(value, float): + return math.isfinite(value) and value >= 0 + if isinstance(value, str): + try: + parsed = float(value) + except ValueError: + return False + return math.isfinite(parsed) and parsed >= 0 + return False + + +def _has_parseable_field(data: object, *fields: str) -> bool: + if not isinstance(data, dict): + return False + return any( + field in data and _is_parseable_token_count(data[field]) for field in fields + ) + + +def usage_field_presence(usage_data: object) -> UsageFieldPresence: + """Label each raw usage count before numeric normalization. + + Explicit zero is reported, while an absent or unusable value is missing. + """ + if not isinstance(usage_data, dict): + return UsageFieldPresence() + + found = "estimated" if usage_data.get("estimated") is True else "reported" + + def source(present: bool) -> str: + return found if present else "missing" + + prompt_details = usage_data.get("prompt_tokens_details") + input_details = usage_data.get("input_tokens_details") + return UsageFieldPresence( + input_source=source( + _has_parseable_field(usage_data, "prompt_tokens", "input_tokens") + ), + output_source=source( + _has_parseable_field(usage_data, "completion_tokens", "output_tokens") + ), + cache_read_source=source( + _has_parseable_field(usage_data, "cache_read_input_tokens") + or _has_parseable_field(prompt_details, "cached_tokens") + or _has_parseable_field(input_details, "cached_tokens") + or _has_parseable_field(usage_data, "prompt_cache_hit_tokens") + ), + cache_creation_source=source( + _has_parseable_field(usage_data, "cache_creation_input_tokens") + or _has_parseable_field( + prompt_details, "cache_creation_tokens", "cache_write_tokens" + ) + or _has_parseable_field(input_details, "cache_write_tokens") + ), + ) + + def _first_token_count(usage_data: dict, *fields: str) -> int: """Return the first positive token count among the given fields.""" for field in fields: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2f65ebd3..d7e7f8fe 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -7,6 +7,7 @@ import traceback import typing import uuid from collections.abc import AsyncGenerator, AsyncIterator, Iterator +from dataclasses import dataclass, replace from typing import Any, Mapping, Self, cast import httpx @@ -40,11 +41,18 @@ from ..core.error_scope import ( ) from ..core.exceptions import UpstreamError from ..core.redaction import redact_org_ids +from ..core.terminal_outcomes import ( + TerminalOutcomeContext, + cashu_retained_msats, + mark_terminal_outcome_loss, + record_terminal_outcome, +) from ..payment.cost_calculation import ( CostData, CostDataError, MaxCostData, calculate_cost, + unpriced_cost, ) from ..payment.helpers import create_error_response from ..payment.models import ( @@ -56,6 +64,7 @@ from ..payment.models import ( list_models, ) from ..payment.price import sats_usd_price +from ..payment.usage import UsageFieldPresence, usage_field_presence from ..wallet import ( SPENT_TOKEN_CODES, classify_redemption_error, @@ -95,6 +104,222 @@ logger = get_logger(__name__) CostMetadata = CostData | MaxCostData | dict[str, Any] +@dataclass +class _TerminalOutcomeState: + context: TerminalOutcomeContext | None + success_marker_seen: bool = False + failure_seen: bool = False + transport_failed: bool = False + usage: dict[str, Any] | None = None + + def observe(self, event: dict[str, Any]) -> None: + event_type = str(event.get("type") or "").lower() + status = str(event.get("status") or "").lower() + nested_response = event.get("response") + if isinstance(nested_response, dict): + status = str(nested_response.get("status") or status).lower() + # Messages report input and output usage in separate events. + for payload in (event.get("message"), event): + if isinstance(payload, dict) and isinstance(payload.get("usage"), dict): + self.usage = {**(self.usage or {}), **payload["usage"]} + if ( + event.get("error") is not None + or event_type in {"error", "response.failed"} + or status in {"cancelled", "failed"} + ): + self.failure_seen = True + return + choices = event.get("choices") + if isinstance(choices, list): + finish_reasons = { + str(choice.get("finish_reason") or "").lower() + for choice in choices + if isinstance(choice, dict) and choice.get("finish_reason") is not None + } + if "error" in finish_reasons: + self.failure_seen = True + return + if finish_reasons - {""}: + self.success_marker_seen = True + # An output-limit truncation is a paid terminal response, like length. + if event_type in { + "response.completed", + "response.incomplete", + "message_stop", + } or status in {"completed", "incomplete"}: + self.success_marker_seen = True + delta = event.get("delta") + if isinstance(delta, dict) and delta.get("stop_reason") not in (None, ""): + self.success_marker_seen = True + + def mark_success(self) -> None: + self.success_marker_seen = True + + def mark_transport_failure(self) -> None: + self.transport_failed = True + + def settlement_context( + self, *, require_success: bool = False + ) -> TerminalOutcomeContext | None: + if self.failure_seen: + return None + if self.transport_failed and not self.success_marker_seen: + return None + if require_success and not self.success_marker_seen: + return None + return self.context + + +def _terminal_outcome_context( + request_id: str | None, model_obj: Model | None +) -> TerminalOutcomeContext: + return TerminalOutcomeContext( + outcome_id=request_id, + model_identifier=(model_obj.canonical_slug or model_obj.id) + if model_obj is not None + else None, + served_model_identifier=(model_obj.forwarded_model_id or model_obj.id) + if model_obj is not None + else None, + ) + + +def _event_usage_presence(event: object) -> UsageFieldPresence: + if not isinstance(event, dict): + return UsageFieldPresence() + presence = usage_field_presence(event.get("usage")) + for key in ("message", "response"): + nested = event.get(key) + if isinstance(nested, dict): + presence = presence.merged(usage_field_presence(nested.get("usage"))) + return presence + + +def _record_x_cashu_terminal_outcome( + context: TerminalOutcomeContext, + cost_data: CostMetadata | None, + *, + amount: int, + unit: str, + refund_amount: int = 0, + usage: object = None, +) -> None: + revenue_msats = cashu_retained_msats(amount, unit, refund_amount) + if revenue_msats is None: + return + metadata = {} + if cost_data is not None: + for name in ( + "input_source", + "output_source", + "cache_read_source", + "cache_creation_source", + "pricing_source", + ): + value = ( + cost_data.get(name) + if isinstance(cost_data, dict) + else getattr(cost_data, name, None) + ) + if isinstance(value, str): + metadata[name] = value + context = replace(context, **metadata) + if usage is not None: + # Stats keep what upstream reported, even where billing did not parse it. + presence = usage_field_presence(usage) + context = replace(context, **presence.sources_dict()) + counted: CostMetadata = cost_data if cost_data is not None else {} + record_terminal_outcome( + context, + input_tokens=int(_cost_field(counted, "input_tokens")), + output_tokens=int(_cost_field(counted, "output_tokens")), + cache_read_input_tokens=int(_cost_field(counted, "cache_read_input_tokens")), + cache_creation_input_tokens=int( + _cost_field(counted, "cache_creation_input_tokens") + ), + revenue_msats=revenue_msats, + usage=usage, + ) + + +async def _track_generic_terminal_stream( + stream: AsyncIterator[bytes], state: _TerminalOutcomeState +) -> AsyncGenerator[bytes, None]: + try: + async for chunk in stream: + yield chunk + state.mark_success() + except BaseException: + state.mark_transport_failure() + raise + + +async def _track_x_cashu_generic_stream( + stream: AsyncIterator[bytes], + state: _TerminalOutcomeState, + *, + amount: int, + unit: str, +) -> AsyncGenerator[bytes, None]: + try: + async for chunk in stream: + yield chunk + state.mark_success() + except BaseException: + state.mark_transport_failure() + raise + finally: + terminal_context = state.settlement_context(require_success=True) + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + None, + amount=amount, + unit=unit, + ) + + +def _observe_terminal_sse_bytes( + state: _TerminalOutcomeState, + buffered: bytes, + chunk: bytes = b"", + *, + final: bool = False, +) -> bytes: + """Observe complete SSE events while preserving a split trailing event.""" + pending = (buffered + chunk).replace(b"\r\n", b"\n") + events: list[bytes] = [] + while b"\n\n" in pending: + event, pending = pending.split(b"\n\n", 1) + events.append(event) + if final and pending.strip(): + events.append(pending) + pending = b"" + + for event in events: + data_lines = [ + line[len(b"data:") :].lstrip(b" ") + for line in event.split(b"\n") + if line.startswith(b"data:") + ] + if not data_lines: + continue + payload = b"\n".join(data_lines) + if payload.strip() == b"[DONE]": + state.mark_success() + continue + try: + parsed = json.loads(payload) + except ValueError: + # Bytes cut inside a character raise UnicodeDecodeError, not JSONDecodeError. + if final: + state.mark_transport_failure() + continue + if isinstance(parsed, dict): + state.observe(parsed) + return pending + + def _cost_field( cost_data: CostMetadata, field: str, default: int | float = 0 ) -> int | float: @@ -1160,6 +1385,7 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, legacy_completion: bool = False, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -1194,6 +1420,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None provider_seen: str | None = None + outcome_state = _TerminalOutcomeState(terminal_outcome) async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1213,6 +1440,9 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context( + require_success=True + ), ) usage_finalized = True except Exception: @@ -1302,11 +1532,13 @@ class BaseUpstreamProvider: if data.strip() == b"[DONE]": done_seen = True + outcome_state.mark_success() return obj = json_codec.loads(data) if isinstance(obj, dict): + outcome_state.observe(obj) usage_estimator.observe(obj) provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): @@ -1358,6 +1590,7 @@ class BaseUpstreamProvider: # mid-event, so ``data`` is incomplete JSON. Emitting it # as a ``data:`` frame would hand the client invalid # JSON (the "unexpected token" parse error). Drop it. + outcome_state.mark_transport_failure() return # Non-JSON data payload (partial fragment already reassembled # by buffering, or a provider control string). Re-prefix each @@ -1401,6 +1634,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) usage_finalized = True except BaseException as e: @@ -1464,6 +1698,7 @@ class BaseUpstreamProvider: yield b"data: [DONE]\n\n" except httpx.RemoteProtocolError as stream_error: + outcome_state.mark_transport_failure() logger.warning( "Upstream stream ended before the response was complete", extra={ @@ -1471,7 +1706,8 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) - except Exception as stream_error: + except BaseException as stream_error: + outcome_state.mark_transport_failure() logger.warning( "Streaming interrupted; finalizing before closing upstream", extra={ @@ -1507,6 +1743,7 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, legacy_completion: bool = False, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> Response: """Handle non-streaming chat completion responses with token usage tracking and cost adjustment. @@ -1532,6 +1769,8 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) + outcome_state = _TerminalOutcomeState(terminal_outcome) + outcome_state.observe(response_json) self._apply_provider_field(response_json) logger.debug( @@ -1565,6 +1804,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) await session.refresh(key) @@ -1655,6 +1895,7 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1680,6 +1921,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None provider_seen: str | None = None + outcome_state = _TerminalOutcomeState(terminal_outcome) async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1699,6 +1941,9 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context( + require_success=True + ), ) usage_finalized = True except Exception: @@ -1774,11 +2019,13 @@ class BaseUpstreamProvider: if data.strip() == b"[DONE]": done_seen = True + outcome_state.mark_success() return obj = json_codec.loads(data) if isinstance(obj, dict): + outcome_state.observe(obj) provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) @@ -1808,6 +2055,7 @@ class BaseUpstreamProvider: # Final flush of a truncated tail: upstream closed # mid-event, so ``data`` is incomplete JSON. Dropping it # avoids handing the client an invalid ``data:`` frame. + outcome_state.mark_transport_failure() return # Re-prefix each line so multi-line ``data`` stays valid SSE # framing for the client. @@ -1845,6 +2093,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) usage_finalized = True except BaseException as e: @@ -1908,6 +2157,7 @@ class BaseUpstreamProvider: yield b"data: [DONE]\n\n" except httpx.RemoteProtocolError as stream_error: + outcome_state.mark_transport_failure() logger.warning( "Upstream Responses API stream ended before the response was complete", extra={ @@ -1915,7 +2165,8 @@ class BaseUpstreamProvider: "key_hash": key.hashed_key[:8] + "...", }, ) - except Exception as stream_error: + except BaseException as stream_error: + outcome_state.mark_transport_failure() logger.warning( "Responses API streaming interrupted; finalizing before closing upstream", extra={ @@ -1950,6 +2201,7 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> Response: """Handle non-streaming Responses API responses with token usage tracking and cost adjustment. @@ -1975,6 +2227,8 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) + outcome_state = _TerminalOutcomeState(terminal_outcome) + outcome_state.observe(response_json) self._apply_provider_field(response_json) logger.debug( @@ -2009,6 +2263,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) await session.refresh(key) @@ -2098,6 +2353,7 @@ class BaseUpstreamProvider: model_obj: Model | None, provider_fee: float | None, reservation_snapshot: ReservationSnapshot, + outcome_state: _TerminalOutcomeState | None = None, ) -> None: """Finalize payment for a generic streaming request.""" async with create_session() as session: @@ -2121,6 +2377,11 @@ class BaseUpstreamProvider: model_obj=model_obj, provider_fee=provider_fee, reservation_snapshot=reservation_snapshot, + terminal_outcome=outcome_state.settlement_context( + require_success=True + ) + if outcome_state is not None + else None, ) logger.debug( "Finalized generic streaming payment", @@ -2149,6 +2410,7 @@ class BaseUpstreamProvider: provider_fee: float | None, reservation_snapshot: ReservationSnapshot, finalizer: PersistentStreamFinalizer | None = None, + outcome_state: _TerminalOutcomeState | None = None, ) -> AsyncGenerator[bytes, None]: """Relay an opaque stream and settle it even if the caller disconnects.""" if finalizer is None: @@ -2161,12 +2423,16 @@ class BaseUpstreamProvider: model_obj, provider_fee, reservation_snapshot, + outcome_state, ), response, ) ) + chunks = response.aiter_bytes() + if outcome_state is not None: + chunks = _track_generic_terminal_stream(chunks, outcome_state) try: - async for chunk in response.aiter_bytes(): + async for chunk in chunks: yield chunk finally: await finalizer.run() @@ -2180,7 +2446,9 @@ class BaseUpstreamProvider: model_obj: Model | None, provider_fee: float | None, reservation_snapshot: ReservationSnapshot, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> ClosingStreamingResponse: + outcome_state = _TerminalOutcomeState(terminal_outcome) finalizer = PersistentStreamFinalizer( lambda: finalize_and_close_stream( lambda: self._finalize_generic_streaming_payment( @@ -2190,6 +2458,7 @@ class BaseUpstreamProvider: model_obj, provider_fee, reservation_snapshot, + outcome_state, ), response, ) @@ -2203,6 +2472,7 @@ class BaseUpstreamProvider: provider_fee, reservation_snapshot, finalizer, + outcome_state, ) return ClosingStreamingResponse( stream, @@ -2220,13 +2490,18 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse: usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_finalized = False last_model_seen: str | None = None provider_seen: str | None = None + usage_presence = UsageFieldPresence() + outcome_state = _TerminalOutcomeState(terminal_outcome) - async def finalize_without_usage() -> bytes | None: + async def finalize_without_usage( + *, require_success: bool = False + ) -> bytes | None: nonlocal usage_finalized if usage_finalized: return None @@ -2244,6 +2519,11 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context( + require_success=require_success + ), + usage_presence=usage_presence, + terminal_usage=outcome_state.usage, ) usage_finalized = True return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() @@ -2265,7 +2545,7 @@ class BaseUpstreamProvider: async def finalize_db_only() -> None: if not usage_finalized: - await finalize_without_usage() + await finalize_without_usage(require_success=True) stream_finalizer = PersistentStreamFinalizer( lambda: finalize_and_close_stream(finalize_db_only, response) @@ -2274,7 +2554,7 @@ class BaseUpstreamProvider: async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - nonlocal usage_finalized, last_model_seen, provider_seen + nonlocal usage_finalized, last_model_seen, provider_seen, usage_presence stored_chunks: list[bytes] = [] input_tokens: int = 0 output_tokens: int = 0 @@ -2283,6 +2563,7 @@ class BaseUpstreamProvider: total_cost: float = 0.0 input_cost: float = 0.0 output_cost: float = 0.0 + terminal_sse_buffer = b"" def _coerce_usd(value: object) -> float: if value is None or isinstance(value, bool): @@ -2316,6 +2597,9 @@ class BaseUpstreamProvider: try: async for chunk in response.aiter_bytes(): stored_chunks.append(chunk) + terminal_sse_buffer = _observe_terminal_sse_bytes( + outcome_state, terminal_sse_buffer, chunk + ) try: decoded_chunk = chunk.decode("utf-8", errors="ignore") modified_lines = [] @@ -2325,6 +2609,9 @@ class BaseUpstreamProvider: try: data = json.loads(line[6:]) if isinstance(data, dict): + usage_presence = usage_presence.merged( + _event_usage_presence(data) + ) usage_estimator.observe(data) msg = data.get("message", {}) if msg and msg.get("model"): @@ -2425,6 +2712,9 @@ class BaseUpstreamProvider: except Exception: yield chunk + _observe_terminal_sse_bytes( + outcome_state, terminal_sse_buffer, final=True + ) usage_data = { "input_tokens": input_tokens, "output_tokens": output_tokens, @@ -2462,6 +2752,9 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), + usage_presence=usage_presence, + terminal_usage=outcome_state.usage, ) self.inject_cost_metadata( @@ -2495,10 +2788,12 @@ class BaseUpstreamProvider: yield maybe_cost_event except httpx.ReadError: + outcome_state.mark_transport_failure() if not usage_finalized: await finalize_without_usage() # Upstream dropped the connection mid-stream; response already started, swallow silently - except Exception: + except BaseException: + outcome_state.mark_transport_failure() if not usage_finalized: await finalize_without_usage() raise @@ -2527,10 +2822,13 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> Response: try: content = await response.aread() response_json = json.loads(content) + outcome_state = _TerminalOutcomeState(terminal_outcome) + outcome_state.observe(response_json) if requested_model: if "model" in response_json: @@ -2560,6 +2858,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2654,6 +2953,7 @@ class BaseUpstreamProvider: max_cost_for_model: int, model_obj: Model, reservation_snapshot: ReservationSnapshot | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> Response | StreamingResponse: """Translate /v1/messages to upstream chat/completions via litellm. @@ -2676,9 +2976,12 @@ class BaseUpstreamProvider: model_obj, reservation_snapshot, request_body, + terminal_outcome, ) response_json = messages_dispatch.coerce_litellm_payload(result) + outcome_state = _TerminalOutcomeState(terminal_outcome) + outcome_state.observe(response_json) if requested_model and "model" in response_json: response_json["model"] = requested_model if not isinstance(response_json.get("usage"), dict): @@ -2696,6 +2999,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), ) self.inject_cost_metadata(response_json, cost_data, key) @@ -2744,6 +3048,10 @@ class BaseUpstreamProvider: ) response_json = messages_dispatch.coerce_litellm_payload(result) + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) + ) + outcome_state.observe(response_json) self._apply_provider_field(response_json) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -2761,19 +3069,28 @@ class BaseUpstreamProvider: self._fold_cache_into_input_tokens(response_json["usage"]) response_headers: dict[str, str] = {} + refund_amount_sent = 0 if cost_data: _inject_cost_response_headers(response_headers, cost_data) refund_amount = messages_dispatch.compute_refund( amount, unit, cost_data.total_msats ) if refund_amount > 0: - refund_token = await self.send_refund( - refund_amount, - unit, - mint, - request_id=request_id, - ) + try: + refund_token = await self.send_refund( + refund_amount, + unit, + mint, + request_id=request_id, + ) + except BaseException: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss( + "X-Cashu LiteLLM refund commit ambiguous" + ) + raise response_headers["X-Cashu"] = refund_token + refund_amount_sent = refund_amount logger.info( "Refund processed for non-streaming /v1/messages via litellm", extra={ @@ -2783,6 +3100,16 @@ class BaseUpstreamProvider: }, ) + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + ) + return Response( content=json.dumps(response_json).encode(), status_code=200, @@ -2801,6 +3128,7 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse: """Re-emit a litellm Anthropic-event iterator as live SSE bytes with cost reconciliation appended at end of stream.""" @@ -2808,8 +3136,12 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_finalized = False last_model_seen: str | None = None + usage_presence = UsageFieldPresence() + outcome_state = _TerminalOutcomeState(terminal_outcome) - async def finalize_without_usage() -> bytes | None: + async def finalize_without_usage( + *, require_success: bool = False + ) -> bytes | None: nonlocal usage_finalized if usage_finalized: return None @@ -2839,6 +3171,10 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context( + require_success=require_success + ), + usage_presence=usage_presence, ) usage_finalized = True return ( @@ -2863,14 +3199,14 @@ class BaseUpstreamProvider: async def finalize_stream() -> None: try: if not usage_finalized: - await finalize_without_usage() + await finalize_without_usage(require_success=True) finally: await aclose_if_needed(iterator) stream_finalizer = PersistentStreamFinalizer(finalize_stream) async def stream_with_cost() -> AsyncGenerator[bytes, None]: - nonlocal usage_finalized, last_model_seen + nonlocal usage_finalized, last_model_seen, usage_presence input_tokens = 0 output_tokens = 0 cache_read_input_tokens = 0 @@ -2883,6 +3219,10 @@ class BaseUpstreamProvider: async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model ): + outcome_state.observe(annotated.event) + usage_presence = usage_presence.merged( + _event_usage_presence(annotated.event) + ) usage_estimator.observe(annotated.event) if annotated.model: last_model_seen = annotated.model @@ -2944,6 +3284,8 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=outcome_state.settlement_context(), + usage_presence=usage_presence, ) self.inject_cost_metadata( combined_data, cost_data, fresh_key @@ -2976,7 +3318,8 @@ class BaseUpstreamProvider: if cost_event is not None: yield cost_event - except Exception: + except BaseException: + outcome_state.mark_transport_failure() if not usage_finalized: await finalize_without_usage() raise @@ -3020,13 +3363,21 @@ class BaseUpstreamProvider: output_tokens = 0 cache_read_input_tokens = 0 cache_creation_input_tokens = 0 + usage_presence = UsageFieldPresence() total_cost = 0.0 input_cost = 0.0 output_cost = 0.0 + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) + ) async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model ): + outcome_state.observe(annotated.event) + usage_presence = usage_presence.merged( + _event_usage_presence(annotated.event) + ) if annotated.model: last_model_seen = annotated.model # See _stream_litellm_messages for why this is max() not +=. @@ -3048,6 +3399,8 @@ class BaseUpstreamProvider: "Cache-Control": "no-cache", "Connection": "keep-alive", } + refund_amount_sent = 0 + settlement_failed = False if ( input_tokens == 0 @@ -3092,7 +3445,10 @@ class BaseUpstreamProvider: } try: cost_data = await self.get_x_cashu_cost( - response_data, max_cost_for_model, model_obj + response_data, + max_cost_for_model, + model_obj, + usage_presence, ) if cost_data: refund_amount = messages_dispatch.compute_refund( @@ -3106,6 +3462,7 @@ class BaseUpstreamProvider: request_id=request_id, ) response_headers["X-Cashu"] = refund_token + refund_amount_sent = refund_amount logger.info( "Refund processed for streaming /v1/messages via litellm", extra={ @@ -3114,7 +3471,18 @@ class BaseUpstreamProvider: "model": last_model_seen, }, ) + except asyncio.CancelledError: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss( + "X-Cashu LiteLLM stream settlement cancelled" + ) + raise except Exception as exc: + settlement_failed = True + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss( + "X-Cashu LiteLLM stream settlement failed" + ) logger.error( "Error calculating cost for streaming /v1/messages", extra={ @@ -3125,6 +3493,17 @@ class BaseUpstreamProvider: }, ) + if not settlement_failed: + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + replace(terminal_context, **usage_presence.sources_dict()), + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + ) + if cost_data: _inject_cost_response_headers(response_headers, cost_data) for index, annotated in enumerate(buffered): @@ -3182,6 +3561,9 @@ class BaseUpstreamProvider: """ completion_path = _openai_completion_path(path) path = self.normalize_request_path(path, model_obj) + terminal_outcome = _terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj + ) if ( path.endswith("messages/count_tokens") @@ -3201,6 +3583,7 @@ class BaseUpstreamProvider: max_cost_for_model=max_cost_for_model, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + terminal_outcome=terminal_outcome, ) url = self.build_request_url(path, model_obj) @@ -3337,6 +3720,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=request_body, + terminal_outcome=terminal_outcome, ) response_handoff.handoff() return result @@ -3353,6 +3737,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=request_body, + terminal_outcome=terminal_outcome, ) finally: await response_handoff.close() @@ -3370,6 +3755,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=request_body, + terminal_outcome=None, ) finally: await response_handoff.close() @@ -3418,6 +3804,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, request_body=request_body, legacy_completion=completion_path == "completions", + terminal_outcome=terminal_outcome, ) response_handoff.handoff() return result @@ -3435,6 +3822,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, request_body=request_body, legacy_completion=completion_path == "completions", + terminal_outcome=terminal_outcome, ) finally: await response_handoff.close() @@ -3459,6 +3847,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=terminal_outcome, ) response_handoff.handoff() return result @@ -3583,6 +3972,9 @@ class BaseUpstreamProvider: Response or StreamingResponse from upstream with cost tracking """ path = self.normalize_request_path(path, model_obj) + terminal_outcome = _terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj + ) url = self.build_request_url(path, model_obj) original_model_id = ( @@ -3705,6 +4097,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=transformed_body, + terminal_outcome=terminal_outcome, ) response_handoff.handoff() return result @@ -3720,6 +4113,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=transformed_body, + terminal_outcome=terminal_outcome, ) finally: await response_handoff.close() @@ -3744,6 +4138,7 @@ class BaseUpstreamProvider: model_obj, self.provider_fee, reservation_snapshot, + terminal_outcome=terminal_outcome, ) response_handoff.handoff() return result @@ -3938,6 +4333,7 @@ class BaseUpstreamProvider: response_data: dict, max_cost_for_model: int, model_obj: Model | None, + usage_presence: UsageFieldPresence | None = None, ) -> MaxCostData | CostData | None: """Calculate cost for X-Cashu payment based on response data. @@ -3962,6 +4358,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, self.provider_fee, + usage_presence, ): case MaxCostData() as cost: logger.debug( @@ -3990,9 +4387,7 @@ class BaseUpstreamProvider: "error_code": error.code, }, ) - return MaxCostData( - base_msats=0, input_msats=0, output_msats=0, total_msats=0 - ) + return unpriced_cost(response_data, usage_presence) return None async def send_refund( @@ -4075,6 +4470,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, + record_outcome: bool = True, ) -> StreamingResponse: """Handle streaming response for X-Cashu payment, calculating refund if needed. @@ -4107,7 +4503,16 @@ class BaseUpstreamProvider: model = None cost_data: CostData | MaxCostData | None = None usage_estimator = MissingUsageEstimator(request_body, model_obj) + refund_amount_sent = 0 + settlement_failed = False + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) if record_outcome else None + ) + # Stats observe both SSE prefix forms; billing keeps its existing parse. + _observe_terminal_sse_bytes( + outcome_state, b"", content_str.encode(), final=True + ) lines = content_str.strip().split("\n") for line in lines: if line.startswith("data: "): @@ -4199,6 +4604,7 @@ class BaseUpstreamProvider: request_id=request_id, ) response_headers["X-Cashu"] = refund_token + refund_amount_sent = refund_amount logger.info( "Refund processed for streaming response", @@ -4224,7 +4630,14 @@ class BaseUpstreamProvider: # extractUsageFromResponseHeaders can populate # inputMsats/outputMsats/totalMsats for x-cashu requests. _inject_cost_response_headers(response_headers, cost_data) + except asyncio.CancelledError: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss("X-Cashu streaming settlement cancelled") + raise except Exception as e: + settlement_failed = True + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss("X-Cashu streaming settlement failed") logger.error( "Error calculating cost for streaming response", extra={ @@ -4236,6 +4649,18 @@ class BaseUpstreamProvider: }, ) + if not settlement_failed: + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + usage=outcome_state.usage or usage_data, + ) + provider_seen: str | None = None for i, line in enumerate(lines): if line.startswith("data: "): @@ -4278,6 +4703,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, + record_outcome: bool = True, ) -> Response: """Handle non-streaming response for X-Cashu payment, calculating refund if needed. @@ -4298,6 +4724,12 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) + if record_outcome + else None + ) + outcome_state.observe(response_json) self._apply_provider_field(response_json) _apply_estimated_usage( response_json, request_body, model_obj, amount, unit, "chat" @@ -4370,12 +4802,17 @@ class BaseUpstreamProvider: ) if refund_amount > 0: - refund_token = await self.send_refund( - refund_amount, - unit, - mint, - request_id=request_id, - ) + try: + refund_token = await self.send_refund( + refund_amount, + unit, + mint, + request_id=request_id, + ) + except BaseException: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss("X-Cashu refund commit ambiguous") + raise response_headers["X-Cashu"] = refund_token logger.info( @@ -4389,6 +4826,16 @@ class BaseUpstreamProvider: }, ) + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + cost_data, + amount=amount, + unit=unit, + refund_amount=max(0, refund_amount), + ) + return Response( content=json.dumps(response_json), status_code=response.status_code, @@ -4445,6 +4892,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, + record_outcome: bool = True, ) -> StreamingResponse | Response: """Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming. @@ -4492,6 +4940,7 @@ class BaseUpstreamProvider: request_id=request_id, model_obj=model_obj, request_body=request_body, + record_outcome=record_outcome, ) else: return await self.handle_x_cashu_non_streaming_response( @@ -4504,6 +4953,7 @@ class BaseUpstreamProvider: request_id=request_id, model_obj=model_obj, request_body=request_body, + record_outcome=record_outcome, ) except Exception as e: @@ -4707,6 +5157,7 @@ class BaseUpstreamProvider: request_id=getattr(request.state, "request_id", None), model_obj=model_obj, request_body=request_body, + record_outcome=not path.endswith("messages/count_tokens"), ) if isinstance(result, StreamingResponse) and not response.is_closed: return attach_upstream_stream_owner(result, response, client) @@ -4718,8 +5169,22 @@ class BaseUpstreamProvider: extra={"path": path, "status_code": response.status_code}, ) + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj + ) + ) return ClosingStreamingResponse( - OwnedUpstreamStream(response.aiter_bytes(), response, client), + OwnedUpstreamStream( + _track_x_cashu_generic_stream( + response.aiter_bytes(), + outcome_state, + amount=amount, + unit=unit, + ), + response, + client, + ), status_code=response.status_code, headers=dict(response.headers), ) @@ -5023,8 +5488,22 @@ class BaseUpstreamProvider: extra={"path": path, "status_code": response.status_code}, ) + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj + ) + ) return ClosingStreamingResponse( - OwnedUpstreamStream(response.aiter_bytes(), response, client), + OwnedUpstreamStream( + _track_x_cashu_generic_stream( + response.aiter_bytes(), + outcome_state, + amount=amount, + unit=unit, + ), + response, + client, + ), status_code=response.status_code, headers=dict(response.headers), ) @@ -5179,9 +5658,15 @@ class BaseUpstreamProvider: reasoning_tokens = 0 cost_data: CostData | MaxCostData | None = None usage_estimator = MissingUsageEstimator(request_body, model_obj) + refund_amount_sent = 0 + settlement_failed = False + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) + ) for _fields, data in events: if data.strip() == "[DONE]": + outcome_state.mark_success() continue try: data_json = json.loads(data) @@ -5190,6 +5675,7 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue usage_estimator.observe(data_json) + outcome_state.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. @@ -5270,6 +5756,7 @@ class BaseUpstreamProvider: request_id=request_id, ) response_headers["X-Cashu"] = refund_token + refund_amount_sent = refund_amount logger.info( "Refund processed for streaming Responses API response", @@ -5295,7 +5782,16 @@ class BaseUpstreamProvider: # extractUsageFromResponseHeaders can populate # inputMsats/outputMsats/totalMsats for x-cashu requests. _inject_cost_response_headers(response_headers, cost_data) + except asyncio.CancelledError: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss( + "X-Cashu Responses stream settlement cancelled" + ) + raise except Exception as e: + settlement_failed = True + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss("X-Cashu Responses stream settlement failed") logger.error( "Error calculating cost for streaming Responses API response", extra={ @@ -5307,6 +5803,18 @@ class BaseUpstreamProvider: }, ) + if not settlement_failed: + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + cost_data, + amount=amount, + unit=unit, + refund_amount=refund_amount_sent, + usage=usage_data, + ) + provider_seen: str | None = None for i, (fields, data) in enumerate(events): if data.strip() == "[DONE]": @@ -5358,6 +5866,10 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) + outcome_state = _TerminalOutcomeState( + _terminal_outcome_context(request_id, model_obj) + ) + outcome_state.observe(response_json) self._apply_provider_field(response_json) _apply_estimated_usage( response_json, request_body, model_obj, amount, unit, "responses" @@ -5427,12 +5939,19 @@ class BaseUpstreamProvider: ) if refund_amount > 0: - refund_token = await self.send_refund( - refund_amount, - unit, - mint, - request_id=request_id, - ) + try: + refund_token = await self.send_refund( + refund_amount, + unit, + mint, + request_id=request_id, + ) + except BaseException: + if outcome_state.settlement_context() is not None: + mark_terminal_outcome_loss( + "X-Cashu Responses refund commit ambiguous" + ) + raise response_headers["X-Cashu"] = refund_token logger.info( @@ -5446,6 +5965,17 @@ class BaseUpstreamProvider: }, ) + terminal_context = outcome_state.settlement_context() + if terminal_context is not None: + _record_x_cashu_terminal_outcome( + terminal_context, + cost_data, + amount=amount, + unit=unit, + refund_amount=max(0, refund_amount), + usage=response_json.get("usage"), + ) + return Response( content=json.dumps(response_json), status_code=response.status_code, diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index c1ee8cfc..6f58b570 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -4,8 +4,8 @@ import json import math import time import traceback -from dataclasses import dataclass, field -from typing import AsyncIterator, Awaitable, Mapping +from dataclasses import dataclass, field, replace +from typing import AsyncIterator, Awaitable, Mapping, TypedDict from urllib.parse import urlsplit, urlunsplit from fastapi import Request @@ -41,13 +41,21 @@ from ..core.error_scope import ( ) from ..core.exceptions import EhbpTimeoutError, UpstreamError from ..core.settings import settings +from ..core.terminal_outcomes import ( + TerminalOutcomeContext, + cashu_retained_msats, + mark_terminal_outcome_loss, + record_terminal_outcome, +) from ..payment.cost_calculation import ( CostData, MaxCostData, calculate_cost, + unpriced_cost, ) from ..payment.helpers import create_error_response from ..payment.models import Model +from ..payment.usage import UsageFieldPresence, usage_field_presence from ..wallet import ( SPENT_TOKEN_CODES, classify_redemption_error, @@ -143,7 +151,9 @@ _PROXY_ONLY_HEADERS = frozenset( TINFOIL_MODEL_PREFIX = "tinfoil-" -def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: +def parse_tinfoil_usage_metrics( + header_value: str | None, *, allow_partial: bool = False +) -> dict | None: """Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict. The header format is:: @@ -194,7 +204,7 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: prompt = int_parts.get("prompt") completion = int_parts.get("completion") - if prompt is None or completion is None: + if (prompt is None or completion is None) and not allow_partial: logger.warning( "Failed to parse X-Tinfoil-Usage-Metrics header", extra={ @@ -204,10 +214,11 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: ) return None - result: dict[str, int | float | str] = { - "prompt_tokens": prompt, - "completion_tokens": completion, - } + result: dict[str, int | float | str] = {} + if prompt is not None: + result["prompt_tokens"] = prompt + if completion is not None: + result["completion_tokens"] = completion if "total" in int_parts: result["total_tokens"] = int_parts["total"] if "cached_prompt_tokens" in int_parts: @@ -221,6 +232,45 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: return result +def _tinfoil_usage_presence(header_value: str | None) -> UsageFieldPresence: + return usage_field_presence( + parse_tinfoil_usage_metrics(header_value, allow_partial=True) + ) + + +def _context_with_presence( + context: TerminalOutcomeContext, + presence: UsageFieldPresence, +) -> TerminalOutcomeContext: + return replace(context, **presence.sources_dict()) + + +def _context_with_served_model( + context: TerminalOutcomeContext, + cost_info: dict, +) -> TerminalOutcomeContext: + identifier = cost_info.pop("actual_model_identifier", None) + served_identifier = cost_info.pop("served_model_identifier", None) + context = replace( + context, + served_model_identifier=served_identifier or context.served_model_identifier, + pricing_source=cost_info.get("pricing_source", "missing"), + ) + unresolved = bool(cost_info.pop("actual_model_unresolved", False)) + if identifier or unresolved: + return replace(context, model_identifier=identifier) + return context + + +def _cost_info_presence(cost_info: Mapping[str, object]) -> UsageFieldPresence: + return UsageFieldPresence( + input_source=str(cost_info.get("input_source", "missing")), + output_source=str(cost_info.get("output_source", "missing")), + cache_read_source=str(cost_info.get("cache_read_source", "missing")), + cache_creation_source=str(cost_info.get("cache_creation_source", "missing")), + ) + + def _get_header_case_insensitive( headers: Mapping[str, str], header_name: str ) -> str | None: @@ -361,6 +411,13 @@ def _prepare_ehbp_upstream_headers( return {**_strip_proxy_headers(headers, profile), **dict(target_headers)} +class _CostTokenCounts(TypedDict): + input_tokens: int + output_tokens: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + + def _build_cost_info( total_msats: int, input_tokens: int = 0, @@ -373,6 +430,11 @@ def _build_cost_info( cache_creation_msats: int = 0, total_usd: float = 0.0, actual_model: str | None = None, + actual_model_identifier: str | None = None, + actual_model_unresolved: bool = False, + usage_presence: UsageFieldPresence | None = None, + served_model_identifier: str | None = None, + pricing_source: str = "missing", ) -> dict: """Build a cost-info dict with token counts and per-token-type costs. @@ -380,7 +442,8 @@ def _build_cost_info( one), it is included in the returned dict so callers can use it for billing finalization and logging. """ - result: dict[str, int | float | str | None] = { + presence = usage_presence if usage_presence is not None else UsageFieldPresence() + result: dict[str, object] = { "total_msats": total_msats, "input_tokens": input_tokens, "output_tokens": output_tokens, @@ -392,9 +455,17 @@ def _build_cost_info( "cache_read_msats": cache_read_msats, "cache_creation_msats": cache_creation_msats, "total_usd": total_usd, + **presence.sources_dict(), + "pricing_source": pricing_source, } + if served_model_identifier: + result["served_model_identifier"] = served_model_identifier if actual_model: result["actual_model"] = actual_model + if actual_model_identifier: + result["actual_model_identifier"] = actual_model_identifier + if actual_model_unresolved: + result["actual_model_unresolved"] = True return result @@ -437,9 +508,20 @@ async def _compute_ehbp_actual_cost( ``total_tokens``, ``input_msats``, and ``output_msats`` (and optionally ``actual_model``). """ + presence = _tinfoil_usage_presence(usage_header) usage_dict = parse_tinfoil_usage_metrics(usage_header) + unpriced = unpriced_cost( + {"usage": parse_tinfoil_usage_metrics(usage_header, allow_partial=True)}, + presence, + ) + fallback_tokens: _CostTokenCounts = { + "input_tokens": unpriced.input_tokens, + "output_tokens": unpriced.output_tokens, + "cache_read_input_tokens": unpriced.cache_read_input_tokens, + "cache_creation_input_tokens": unpriced.cache_creation_input_tokens, + } if usage_dict is None: - return _build_cost_info(0) + return _build_cost_info(0, usage_presence=presence, **fallback_tokens) # The enclave may serve a different model than the one requested (e.g. # due to failover). The usage-metrics header's ``model=`` carries @@ -450,6 +532,11 @@ async def _compute_ehbp_actual_cost( # from the expected upstream ID do we treat it as a real mismatch and # look up the actual model's pricing. actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type] + served_model_identifier = ( + actual_model or model_obj.forwarded_model_id or model_obj.id + ) + actual_model_identifier: str | None = None + actual_model_unresolved = False pricing_model_id = model_obj.id # Bill the model we actually routed to. Passing only the model *string* # to calculate_cost makes it re-derive pricing from the global alias map, @@ -483,10 +570,9 @@ async def _compute_ehbp_actual_cost( # model — instead of the cheaper cross-provider model the bare id # would resolve to in the global map. namespaced_served = actual_model - if ( - expected_upstream_model.startswith(TINFOIL_MODEL_PREFIX) - and not actual_model.startswith(TINFOIL_MODEL_PREFIX) - ): + if expected_upstream_model.startswith( + TINFOIL_MODEL_PREFIX + ) and not actual_model.startswith(TINFOIL_MODEL_PREFIX): namespaced_served = TINFOIL_MODEL_PREFIX + actual_model actual_model_obj = get_model_instance(namespaced_served) @@ -502,6 +588,7 @@ async def _compute_ehbp_actual_cost( "actual_model": actual_model, }, ) + actual_model_unresolved = True actual_model = None else: resolved_upstream_model = ( @@ -521,6 +608,9 @@ async def _compute_ehbp_actual_cost( ) pricing_model_id = actual_model_obj.id pricing_model_obj = actual_model_obj + actual_model_identifier = ( + actual_model_obj.canonical_slug or actual_model_obj.id + ) else: # A different registry/client alias resolved to the same # upstream model; retain the requested model's pricing. @@ -544,7 +634,15 @@ async def _compute_ehbp_actual_cost( "usage": usage_dict, }, ) - return _build_cost_info(0, actual_model=actual_model) + return _build_cost_info( + 0, + **fallback_tokens, + actual_model=actual_model, + actual_model_identifier=actual_model_identifier, + served_model_identifier=served_model_identifier, + actual_model_unresolved=actual_model_unresolved, + usage_presence=presence, + ) if isinstance(cost, MaxCostData): logger.warning( @@ -557,7 +655,15 @@ async def _compute_ehbp_actual_cost( "cost_total_msats": cost.total_msats, }, ) - return _build_cost_info(0, actual_model=actual_model) + return _build_cost_info( + 0, + **fallback_tokens, + actual_model=actual_model, + actual_model_identifier=actual_model_identifier, + served_model_identifier=served_model_identifier, + actual_model_unresolved=actual_model_unresolved, + usage_presence=presence, + ) if isinstance(cost, CostData): actual = max(int(cost.total_msats), int(settings.min_request_msat)) clamped = min(actual, max_cost_for_model) @@ -582,7 +688,12 @@ async def _compute_ehbp_actual_cost( cache_read_msats=cost.cache_read_msats, cache_creation_msats=cost.cache_creation_msats, total_usd=cost.total_usd, + pricing_source=cost.pricing_source, actual_model=actual_model, + actual_model_identifier=actual_model_identifier, + served_model_identifier=served_model_identifier, + actual_model_unresolved=actual_model_unresolved, + usage_presence=presence, ) # CostDataError logger.warning( @@ -592,7 +703,15 @@ async def _compute_ehbp_actual_cost( "error": getattr(cost, "message", str(cost)), }, ) - return _build_cost_info(0, actual_model=actual_model) + return _build_cost_info( + 0, + **fallback_tokens, + actual_model=actual_model, + actual_model_identifier=actual_model_identifier, + served_model_identifier=served_model_identifier, + actual_model_unresolved=actual_model_unresolved, + usage_presence=presence, + ) def _extract_usage_from_response( @@ -687,6 +806,7 @@ async def finalize_ehbp_actual_cost_payment( model_id: str, cost_info: dict, reservation_snapshot: ReservationSnapshot | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> int: """Finalize an EHBP bearer request using clamped provider usage metrics.""" reservation = reservation_snapshot or await get_reservation_snapshot(key, session) @@ -721,7 +841,24 @@ async def finalize_ehbp_actual_cost_payment( await _release_failed_ehbp_charge(reservation, session) return 0 - await session.commit() + try: + await session.commit() + except BaseException: + if terminal_outcome is not None: + mark_terminal_outcome_loss("ehbp_commit_ambiguous") + raise + if terminal_outcome is not None: + record_terminal_outcome( + _context_with_presence( + _context_with_served_model(terminal_outcome, cost_info), + _cost_info_presence(cost_info), + ), + input_tokens=cost_info.get("input_tokens", 0), + output_tokens=cost_info.get("output_tokens", 0), + cache_read_input_tokens=cost_info.get("cache_read_input_tokens", 0), + cache_creation_input_tokens=cost_info.get("cache_creation_input_tokens", 0), + revenue_msats=total_cost_msats, + ) await _stop_reservation_heartbeat(reservation.release_id) await session.refresh(key) @@ -772,6 +909,8 @@ async def finalize_ehbp_max_cost_payment( max_cost_for_model: int, model_id: str, reservation_snapshot: ReservationSnapshot | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, + usage_data: dict | None = None, ) -> int: """Release an unmeasured EHBP request without charging its reservation. @@ -782,7 +921,29 @@ async def finalize_ehbp_max_cost_payment( reservation = reservation_snapshot or await get_reservation_snapshot(key, session) await _validate_reservation_snapshot(key, reservation, session) key_log_hash = key.hashed_key[:8] + "..." - await release_reservation(reservation, session, reservation.reserved_msats) + try: + released = await release_reservation( + reservation, + session, + reservation.reserved_msats, + idempotent_success=False, + ) + except BaseException: + if terminal_outcome is not None: + mark_terminal_outcome_loss("ehbp_release_commit_ambiguous") + raise + if not released: + await _stop_reservation_heartbeat(reservation.release_id) + if released and terminal_outcome is not None: + usage = unpriced_cost({"usage": usage_data}) + record_terminal_outcome( + terminal_outcome, + input_tokens=usage.input_tokens, + output_tokens=usage.output_tokens, + cache_read_input_tokens=usage.cache_read_input_tokens, + cache_creation_input_tokens=usage.cache_creation_input_tokens, + revenue_msats=0, + ) logger.warning( "Released unmeasured EHBP reservation without charging max cost", extra={ @@ -841,6 +1002,11 @@ async def forward_ehbp_request( trailer (streaming). Usage is captured from both response headers and HTTP trailers via an h11-based client (httpx silently discards trailers). """ + terminal_outcome = TerminalOutcomeContext( + outcome_id=getattr(request.state, "request_id", None), + model_identifier=model_obj.canonical_slug or model_obj.id, + served_model_identifier=model_obj.forwarded_model_id or model_obj.id, + ) target = upstream.get_ehbp_forwarding_target(path, model_obj) # type: ignore[attr-defined] provider_type = getattr(upstream, "provider_type", "unknown") @@ -929,6 +1095,10 @@ async def forward_ehbp_request( usage_header = _extract_usage_from_response( resp.headers, resp.trailers, usage_header_name ) + terminal_outcome = _context_with_presence( + terminal_outcome, + _tinfoil_usage_presence(usage_header), + ) usage_dict = parse_tinfoil_usage_metrics(usage_header) usage_source = ( "header" @@ -967,6 +1137,7 @@ async def forward_ehbp_request( usage_header, model_obj, max_cost_for_model ) billing_model = cost_info.pop("actual_model", None) or model_obj.id + terminal_outcome = _context_with_served_model(terminal_outcome, cost_info) computed_msats = int(cost_info["total_msats"]) charged_msats = await _record_ehbp_settlement( finalize_ehbp_actual_cost_payment( @@ -976,6 +1147,7 @@ async def forward_ehbp_request( billing_model, cost_info, reservation_snapshot, + terminal_outcome, ), key=key, model_id=billing_model, @@ -1006,6 +1178,10 @@ async def forward_ehbp_request( max_cost_for_model, model_obj.id, reservation_snapshot, + terminal_outcome, + usage_data=parse_tinfoil_usage_metrics( + usage_header, allow_partial=True + ), ), key=key, model_id=model_obj.id, @@ -1105,6 +1281,11 @@ async def forward_ehbp_x_cashu_request( client because httpx silently discards them. """ request_id = getattr(request.state, "request_id", None) + terminal_outcome = TerminalOutcomeContext( + outcome_id=request_id, + model_identifier=model_obj.canonical_slug or model_obj.id, + served_model_identifier=model_obj.forwarded_model_id or model_obj.id, + ) amount = 0 unit = "msat" mint: str | None = None @@ -1233,8 +1414,13 @@ async def forward_ehbp_x_cashu_request( cost_info = await _compute_ehbp_actual_cost( usage_header, model_obj, max_cost_for_model ) + terminal_outcome = _context_with_presence( + terminal_outcome, + _cost_info_presence(cost_info), + ) actual_cost_msats = cost_info["total_msats"] - actual_model = cost_info.get("actual_model") + actual_model = cost_info.pop("actual_model", None) + terminal_outcome = _context_with_served_model(terminal_outcome, cost_info) billing_model = actual_model or model_obj.id refund_amount = amount - _msats_to_unit_amount(actual_cost_msats, unit) logger.info( @@ -1268,9 +1454,32 @@ async def forward_ehbp_x_cashu_request( # opaque encrypted blobs, cost can only go into response headers. _inject_cost_response_headers(response_headers, cost_info) + persisted_refund_amount = 0 if refund_amount > 0: - response_headers["X-Cashu"] = await send_cashu_refund( - refund_amount, unit, mint, request_id + try: + response_headers["X-Cashu"] = await send_cashu_refund( + refund_amount, unit, mint, request_id + ) + except BaseException: + mark_terminal_outcome_loss("ehbp_x_cashu_refund_commit_ambiguous") + raise + persisted_refund_amount = refund_amount + + revenue_msats = cashu_retained_msats( + amount, + unit, + refund_amount=persisted_refund_amount, + ) + if revenue_msats is not None: + record_terminal_outcome( + terminal_outcome, + input_tokens=cost_info.get("input_tokens", 0), + output_tokens=cost_info.get("output_tokens", 0), + cache_read_input_tokens=cost_info.get("cache_read_input_tokens", 0), + cache_creation_input_tokens=cost_info.get( + "cache_creation_input_tokens", 0 + ), + revenue_msats=revenue_msats, ) async def _stream_body_xcashu() -> AsyncIterator[bytes]: diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index f82b3248..86691e9a 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -9,7 +9,7 @@ request body, and echo the fallback's model id to the client. import json from typing import Any, AsyncGenerator -from unittest.mock import patch +from unittest.mock import MagicMock, patch import httpx import pytest @@ -165,6 +165,7 @@ async def test_failover_serve_billed_at_serving_providers_rate( at the winner's (0.001/0.002) it would be 2_000 msats. """ sent_requests: list[httpx.Request] = [] + record_outcome = MagicMock() # Patch the network transport (not AsyncClient.send) so the in-process # ASGI test client is untouched and only the proxy's upstream hop is mocked. @@ -185,6 +186,7 @@ async def test_failover_serve_billed_at_serving_providers_rate( "routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005, ), + patch("routstr.auth.record_terminal_outcome", record_outcome), ): response = await authenticated_client.post( "/v1/chat/completions", @@ -213,6 +215,11 @@ async def test_failover_serve_billed_at_serving_providers_rate( # Billed at the serving provider's rate: 1000/1000*5000 + 500/1000*10000. assert payload["cost"]["total_msats"] == 10_000 + record_outcome.assert_called_once() + terminal_outcome = record_outcome.call_args.args[0] + serving_model = dual_provider_maps[1].get_cached_models()[0] + assert terminal_outcome.model_identifier == serving_model.id + assert record_outcome.call_args.kwargs["revenue_msats"] == 10_000 # The fallback's larger max-cost envelope requires a replacement # reservation. The failed candidate is released, the serving candidate is diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 65ad5091..1208e3b0 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -46,8 +46,8 @@ async def test_openai_cache_subtraction() -> None: "completion_tokens": 100, "prompt_tokens_details": { "cached_tokens": 1000 # ← Extracted separately - } - } + }, + }, } result = await calculate_cost(response, max_cost=100000) @@ -55,6 +55,9 @@ async def test_openai_cache_subtraction() -> None: assert result.input_tokens == 1000 # 2000 - 1000 assert result.cache_read_input_tokens == 1000 assert result.output_tokens == 100 + assert result.input_source == result.output_source == "reported" + assert result.cache_read_source == "reported" + assert result.cache_creation_source == "missing" # ============================================================================ @@ -776,6 +779,8 @@ async def test_missing_usage_block(mock_fixed_pricing: None) -> None: assert result.input_tokens == 0 assert result.cache_read_input_tokens == 0 assert result.output_tokens == 0 + assert result.input_source == result.output_source == "missing" + assert result.cache_read_source == result.cache_creation_source == "missing" # ============================================================================ @@ -790,3 +795,61 @@ async def test_null_usage_block(mock_fixed_pricing: None) -> None: assert isinstance(result, MaxCostData) assert result.input_tokens == 0 assert result.cache_read_input_tokens == 0 + + +@pytest.mark.asyncio +async def test_explicit_zero_presence_survives_cost_calculation( + mock_fixed_pricing: None, +) -> None: + response = { + "model": "gpt-4", + "usage": { + "input_tokens": 0, + "output_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + } + + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + assert result.input_source == result.output_source == "reported" + assert result.cache_read_source == result.cache_creation_source == "reported" + assert "input_source" not in result.dict() + assert "pricing_source" not in result.dict() + + +@pytest.mark.asyncio +async def test_estimated_usage_cost_is_not_reported(mock_fixed_pricing: None) -> None: + response = { + "model": "gpt-4", + "usage": { + "input_tokens": 12, + "output_tokens": 3, + "estimated": True, + }, + } + + result = await calculate_cost(response, max_cost=100000) + + assert isinstance(result, CostData) + assert result.input_tokens == 12 + assert result.output_tokens == 3 + assert result.input_source == result.output_source == "estimated" + assert result.cache_read_source == result.cache_creation_source == "missing" + assert result.pricing_source == "fixed" + + +def test_unpriced_estimate_keeps_estimated_provenance() -> None: + from routstr.payment.cost_calculation import unpriced_cost + from routstr.payment.usage import UsageFieldPresence + + result = unpriced_cost( + {"usage": {"input_tokens": 12, "output_tokens": 3, "estimated": True}}, + UsageFieldPresence(), + ) + assert result.total_msats == 0 + assert result.input_tokens == 12 and result.output_tokens == 3 + assert result.input_source == result.output_source == "estimated" + assert result.pricing_source == "missing" diff --git a/tests/unit/test_ehbp_finalize_payment.py b/tests/unit/test_ehbp_finalize_payment.py index caf0a5aa..cece6e55 100644 --- a/tests/unit/test_ehbp_finalize_payment.py +++ b/tests/unit/test_ehbp_finalize_payment.py @@ -2,6 +2,7 @@ from __future__ import annotations import logging from contextlib import contextmanager +from types import SimpleNamespace from typing import Any, AsyncGenerator, Iterator from unittest.mock import AsyncMock, MagicMock @@ -12,13 +13,19 @@ from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession import routstr.auth as auth_module +import routstr.core.terminal_outcomes as outcomes_module from routstr.auth import get_reservation_snapshot, pay_for_request from routstr.core.db import ApiKey, ReservationRelease +from routstr.core.terminal_outcomes import TerminalOutcomeContext from routstr.upstream.ehbp import ( + EHBPForwardingTarget, + _context_with_served_model, _inject_cost_response_headers, finalize_ehbp_actual_cost_payment, finalize_ehbp_max_cost_payment, + forward_ehbp_x_cashu_request, ) +from routstr.upstream.tinfoil_trailer import TrailerResponse def _make_engine() -> AsyncEngine: @@ -29,6 +36,16 @@ def _make_engine() -> AsyncEngine: ) +def test_unresolved_served_model_does_not_publish_requested_identity() -> None: + context = TerminalOutcomeContext("ehbp-unknown-model", "requested/model") + cost_info = {"actual_model_unresolved": True} + + resolved = _context_with_served_model(context, cost_info) + + assert resolved.model_identifier is None + assert cost_info == {} + + @pytest.fixture async def session( monkeypatch: pytest.MonkeyPatch, @@ -77,12 +94,19 @@ def _fail_nth_api_key_update( @pytest.mark.asyncio async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve( session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, ) -> None: key = ApiKey(hashed_key="ehbp-actual", balance=10_000) session.add(key) await session.commit() await pay_for_request(key, 3_000, session) reservation = await get_reservation_snapshot(key, session) + record_outcome = MagicMock() + monkeypatch.setattr("routstr.upstream.ehbp.record_terminal_outcome", record_outcome) + terminal_outcome = TerminalOutcomeContext( + outcome_id="ehbp-actual-outcome", + model_identifier="tinfoil/model", + ) charged = await finalize_ehbp_actual_cost_payment( key, @@ -95,8 +119,13 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve "output_tokens": 20, "input_msats": 500, "output_msats": 700, + "input_source": "reported", + "output_source": "reported", + "cache_read_source": "reported", + "cache_creation_source": "missing", }, reservation_snapshot=reservation, + terminal_outcome=terminal_outcome, ) assert charged == 1_200 @@ -106,6 +135,22 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve assert updated.reserved_balance == 0 assert updated.reserved_at is None assert updated.total_spent == 1_200 + record_outcome.assert_called_once_with( + TerminalOutcomeContext( + outcome_id="ehbp-actual-outcome", + model_identifier="tinfoil/model", + pricing_source="missing", + input_source="reported", + output_source="reported", + cache_read_source="reported", + cache_creation_source="missing", + ), + input_tokens=10, + output_tokens=20, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=1_200, + ) @contextmanager @@ -216,12 +261,19 @@ async def test_finalize_actual_cost_payment_logs_zero_cache_when_absent( @pytest.mark.asyncio async def test_unmeasured_ehbp_releases_reservation( session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, ) -> None: key = ApiKey(hashed_key="ehbp-key", balance=10_000) session.add(key) await session.commit() await pay_for_request(key, 3_000, session) reservation = await get_reservation_snapshot(key, session) + record_outcome = MagicMock() + monkeypatch.setattr("routstr.upstream.ehbp.record_terminal_outcome", record_outcome) + terminal_outcome = TerminalOutcomeContext( + outcome_id="ehbp-unmeasured-outcome", + model_identifier="tinfoil/model", + ) charged = await finalize_ehbp_max_cost_payment( key, @@ -229,6 +281,7 @@ async def test_unmeasured_ehbp_releases_reservation( max_cost_for_model=3_000, model_id="tinfoil/model", reservation_snapshot=reservation, + terminal_outcome=terminal_outcome, ) assert charged == 0 @@ -238,6 +291,17 @@ async def test_unmeasured_ehbp_releases_reservation( assert updated.reserved_balance == 0 assert updated.reserved_at is None assert updated.total_spent == 0 + record_outcome.assert_called_once_with( + TerminalOutcomeContext( + outcome_id="ehbp-unmeasured-outcome", + model_identifier="tinfoil/model", + ), + input_tokens=0, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=0, + ) @pytest.mark.asyncio @@ -324,3 +388,101 @@ def test_zero_debit_ehbp_headers_preserve_computed_cost() -> None: assert headers["X-Routstr-Cost-Msats"] == "0" assert headers["X-Routstr-Computed-Cost-Msats"] == "1500" + + +@pytest.mark.asyncio +async def test_x_cashu_ledger_failure_cannot_trigger_full_refund( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class FailingWriter: + def __init__(self) -> None: + self.losses: list[str] = [] + self.submissions: list[object] = [] + + def submit(self, outcome: object) -> bool: + self.submissions.append(outcome) + raise RuntimeError("ledger unavailable") + + def declare_loss(self, reason: str, lost_day: object = None) -> None: + self.losses.append(reason) + + writer = FailingWriter() + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + monkeypatch.setattr( + "routstr.upstream.ehbp.recieve_token", + AsyncMock(return_value=(10, "sat", "https://mint.example")), + ) + monkeypatch.setattr("routstr.upstream.ehbp.store_cashu_transaction", AsyncMock()) + monkeypatch.setattr( + "routstr.upstream.ehbp.forward_with_trailer", + AsyncMock( + return_value=TrailerResponse( + status_code=200, + headers=[], + body=b"encrypted-response", + ) + ), + ) + monkeypatch.setattr( + "routstr.upstream.ehbp._compute_ehbp_actual_cost", + AsyncMock( + return_value={ + "total_msats": 1_999, + "input_tokens": 10, + "output_tokens": 5, + "input_msats": 1_000, + "output_msats": 999, + "input_source": "reported", + "output_source": "reported", + "cache_read_source": "missing", + "cache_creation_source": "missing", + "actual_model": "served-model", + "actual_model_identifier": "served/canonical-model", + } + ), + ) + send_refund = AsyncMock(return_value="cashu-refund") + monkeypatch.setattr("routstr.upstream.ehbp.send_cashu_refund", send_refund) + + request = MagicMock() + request.state = SimpleNamespace(request_id="ehbp-xcashu-ledger-failure") + request.headers = {} + request.method = "POST" + request.query_params = {} + request.body = AsyncMock(return_value=b"encrypted-request") + upstream = MagicMock() + upstream.provider_type = "tinfoil" + upstream.prepare_headers.return_value = {} + upstream.get_ehbp_forwarding_target.return_value = EHBPForwardingTarget( + url="https://enclave.tinfoil.sh/v1/chat/completions" + ) + upstream.get_confidential_inference_profile.return_value = None + upstream.prepare_params.return_value = {} + model = MagicMock() + model.id = "tinfoil/model" + model.canonical_slug = "author/model" + + response = await forward_ehbp_x_cashu_request( + request=request, + x_cashu_token="cashu-input", + path="v1/chat/completions", + max_cost_for_model=10_000, + model_obj=model, + upstream=upstream, + ) + + assert response.status_code == 200 + assert response.headers["x-cashu"] == "cashu-refund" + send_refund.assert_awaited_once_with( + 8, + "sat", + "https://mint.example", + "ehbp-xcashu-ledger-failure", + ) + assert writer.losses == ["terminal outcome submission raised"] + assert len(writer.submissions) == 1 + submission = writer.submissions[0] + assert getattr(submission, "model_identifier") == "served/canonical-model" + assert getattr(submission, "revenue_msats") == 2_000 + assert getattr(submission, "input_source") == "reported" + assert getattr(submission, "output_source") == "reported" diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 47b3ccd9..c9b6910d 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -578,7 +578,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: "role": "assistant", "model": "openai/gpt-4o-mini", "content": [], - "usage": {"input_tokens": 5, "output_tokens": 0}, + "usage": {"input_tokens": 0, "output_tokens": 0}, }, } yield { @@ -595,7 +595,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: yield { "type": "message_delta", "delta": {"stop_reason": "end_turn"}, - "usage": {"output_tokens": 7}, + "usage": {"output_tokens": 0}, } yield {"type": "message_stop"} @@ -617,10 +617,13 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: model_obj: Any = None, provider_fee: Any = None, reservation_snapshot: Any = None, + terminal_outcome: Any = None, + usage_presence: Any = None, ) -> dict: captured_cost_call["combined_data"] = combined_data captured_cost_call["max_cost"] = max_cost captured_cost_call["reservation_snapshot"] = reservation_snapshot + captured_cost_call["usage_presence"] = usage_presence return fake_cost fake_session = MagicMock() @@ -647,6 +650,11 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: "routstr.upstream.base.create_session", new=lambda: FakeSessionCtx(), ), + patch("routstr.upstream.count_tokens._count_with_litellm", return_value=0), + patch( + "routstr.upstream.count_tokens._count_text_with_litellm", + return_value=0, + ), ): result = await provider._forward_messages_via_litellm( request_body=body, @@ -674,10 +682,13 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None: assert "event: cost" in joined combined = captured_cost_call["combined_data"] - assert combined["usage"]["input_tokens"] == 5 - assert combined["usage"]["output_tokens"] == 7 + assert combined["usage"]["input_tokens"] == 0 + assert combined["usage"]["output_tokens"] == 0 + assert combined["usage"]["estimated"] is True assert combined["model"] == "openai/gpt-4o-mini" assert captured_cost_call["reservation_snapshot"] is reservation + assert captured_cost_call["usage_presence"].input_source == "reported" + assert captured_cost_call["usage_presence"].output_source == "reported" @pytest.mark.asyncio @@ -721,9 +732,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: model_obj: Any = None, provider_fee: Any = None, reservation_snapshot: Any = None, + terminal_outcome: Any = None, + usage_presence: Any = None, ) -> dict: captured["combined_data"] = combined_data captured["reservation_snapshot"] = reservation_snapshot + captured["usage_presence"] = usage_presence return fake_cost fake_session = MagicMock() @@ -781,6 +795,8 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None: assert combined["usage"]["output_tokens"] == 4 assert combined["model"] == "openai/gpt-4o-mini" assert captured["reservation_snapshot"] is reservation + assert captured["usage_presence"].input_source == "reported" + assert captured["usage_presence"].output_source == "reported" # --------------------------------------------------------------------------- @@ -938,7 +954,7 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost), - ), + ) as mock_get_cost, patch.object( provider, "send_refund", @@ -963,6 +979,11 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None refund_call = mock_refund.await_args assert refund_call is not None assert refund_call.args[0] == 3_500 + get_cost_call = mock_get_cost.await_args + assert get_cost_call is not None + usage_presence = get_cost_call.args[3] + assert usage_presence.input_source == "reported" + assert usage_presence.output_source == "reported" emitted: list[bytes] = [] async for chunk in result.body_iterator: diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 10c5a3e8..65516744 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -1,8 +1,9 @@ import asyncio import json from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager from typing import cast -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import ANY, AsyncMock, MagicMock, patch import httpx import pytest @@ -21,9 +22,15 @@ from routstr.auth import ( release_reservation, ) from routstr.core.db import ApiKey, ReservationRelease +from routstr.core.terminal_outcomes import TerminalOutcomeContext from routstr.payment.cost_calculation import MaxCostData from routstr.payment.models import Architecture, Model, Pricing -from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.base import ( + BaseUpstreamProvider, + _observe_terminal_sse_bytes, + _TerminalOutcomeState, + _track_generic_terminal_stream, +) async def _engine() -> AsyncEngine: @@ -33,6 +40,25 @@ async def _engine() -> AsyncEngine: return engine +@pytest.mark.asyncio +async def test_generic_terminal_outcome_requires_consumed_clean_eof() -> None: + context = TerminalOutcomeContext( + outcome_id="generic-terminal", + model_identifier="test-model", + ) + state = _TerminalOutcomeState(context) + + assert state.settlement_context(require_success=True) is None + + async def chunks() -> AsyncGenerator[bytes, None]: + yield b"complete" + + assert [ + chunk async for chunk in _track_generic_terminal_stream(chunks(), state) + ] == [b"complete"] + assert state.settlement_context(require_success=True) is context + + @pytest.mark.asyncio async def test_release_reservation_is_durable_and_idempotent() -> None: engine = await _engine() @@ -131,15 +157,23 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None: input_msats=0, output_msats=0, total_msats=500, + input_source="reported", + cache_read_source="reported", ) async with AsyncSession(engine, expire_on_commit=False) as session: session.add(key) await session.commit() await pay_for_request(key, 500, session) snapshot = await get_reservation_snapshot(key, session) + terminal_outcome = TerminalOutcomeContext( + outcome_id="post-commit-refresh-failure", + model_identifier="test-model", + ) + record_outcome = MagicMock() with ( patch("routstr.auth.calculate_cost", AsyncMock(return_value=cost)), + patch("routstr.auth.record_terminal_outcome", record_outcome), patch.object( session, "refresh", @@ -147,7 +181,14 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None: ), ): with pytest.raises(SQLAlchemyError, match="post-commit refresh failed"): - await adjust_payment_for_tokens(key, {}, session, 500) + await adjust_payment_for_tokens( + key, + {}, + session, + 500, + reservation_snapshot=snapshot, + terminal_outcome=terminal_outcome, + ) await session.rollback() assert await release_reservation(snapshot, session, 500) is False @@ -156,6 +197,23 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None: assert charged_key is not None assert (charged_key.balance, charged_key.reserved_balance) == (500, 0) assert record is not None and record.status == "charged" + record_outcome.assert_called_once_with( + TerminalOutcomeContext( + outcome_id="post-commit-refresh-failure", + model_identifier="test-model", + pricing_source="missing", + input_source="reported", + output_source="missing", + cache_read_source="reported", + cache_creation_source="missing", + ), + input_tokens=0, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=500, + usage=None, + ) await engine.dispose() @@ -265,6 +323,7 @@ async def test_generic_stream_completion_settles_and_closes_once() -> None: None, provider.provider_fee, reservation, + None, ) response.aclose.assert_awaited_once_with() @@ -299,6 +358,7 @@ async def test_generic_stream_abort_settles_and_closes_once() -> None: None, provider.provider_fee, reservation, + None, ) response.aclose.assert_awaited_once_with() @@ -362,6 +422,7 @@ async def test_streaming_response_closes_iterator_when_downstream_send_is_cancel None, provider.provider_fee, reservation, + ANY, ) upstream_response.aclose.assert_awaited_once_with() @@ -417,6 +478,7 @@ async def test_generic_stream_settles_when_response_start_fails() -> None: None, provider.provider_fee, reservation, + ANY, ) upstream_response.aclose.assert_awaited_once_with() @@ -614,16 +676,23 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure( @pytest.mark.asyncio @pytest.mark.parametrize("api", ["chat", "responses"]) @pytest.mark.parametrize("finalization_fails", [False, True]) +@pytest.mark.parametrize("terminal_marker_seen", [False, True]) async def test_partial_remote_protocol_error_finalizes_and_closes_once( api: str, finalization_fails: bool, + terminal_marker_seen: bool, ) -> None: provider = BaseUpstreamProvider( base_url="https://api.example.com", api_key="test-key" ) async def aiter_bytes() -> AsyncGenerator[bytes, None]: - yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n' + if terminal_marker_seen and api == "chat": + yield b'data: {"model":"test","choices":[{"finish_reason":"stop"}]}\n\n' + elif terminal_marker_seen: + yield b'data: {"type":"response.completed","response":{"status":"completed"}}\n\n' + else: + yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n' raise httpx.RemoteProtocolError("incomplete chunked read") upstream_response = MagicMock( @@ -651,6 +720,10 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( billing_key_hash=key.hashed_key, reserved_msats=500, ) + terminal_outcome = TerminalOutcomeContext( + outcome_id=f"{api}-partial-outcome", + model_identifier="test-model", + ) release = AsyncMock(return_value=True) with ( @@ -664,6 +737,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( key=key, max_cost_for_model=500, reservation_snapshot=snapshot, + terminal_outcome=terminal_outcome, ) else: response = await provider.handle_streaming_responses_completion( @@ -671,12 +745,17 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( key=key, max_cost_for_model=500, reservation_snapshot=snapshot, + terminal_outcome=terminal_outcome, ) emitted = bytearray() async for chunk in response.body_iterator: emitted.extend(chunk.encode() if isinstance(chunk, str) else bytes(chunk)) adjust.assert_awaited_once() + assert adjust.await_args is not None + assert adjust.await_args.kwargs["terminal_outcome"] is ( + terminal_outcome if terminal_marker_seen else None + ) if finalization_fails: session.rollback.assert_awaited_once() release.assert_awaited_once_with(snapshot, session, 500) @@ -686,6 +765,58 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once( assert b"[DONE]" not in emitted +@pytest.mark.asyncio +@pytest.mark.parametrize("api", ["chat", "responses", "messages"]) +async def test_stream_closed_before_first_chunk_is_not_a_completed_outcome( + api: str, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aclose = AsyncMock() + key = MagicMock(spec=ApiKey) + key.hashed_key = f"{api}-never-started" + key.balance = 10_000 + session = MagicMock() + session.get = AsyncMock(return_value=key) + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0}) + snapshot = ReservationSnapshot( + release_id=f"{api}-never-started-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + terminal_outcome = TerminalOutcomeContext( + outcome_id=f"{api}-never-started", model_identifier="test-model" + ) + handler = getattr(provider, f"handle_streaming_{api}_completion") + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await handler( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + terminal_outcome=terminal_outcome, + ) + # The client leaves before the first byte, so only the finalizer runs. + await cast(AsyncGenerator[bytes, None], response.body_iterator).aclose() + + adjust.assert_awaited_once() + assert adjust.await_args is not None + assert adjust.await_args.kwargs["terminal_outcome"] is None + upstream_response.aclose.assert_awaited_once() + + @pytest.mark.asyncio @pytest.mark.parametrize("api", ["chat", "responses"]) async def test_partial_stream_closes_when_billing_db_is_down( @@ -989,6 +1120,215 @@ async def test_gemini_messages_finalizes_when_response_start_fails() -> None: assert upstream_stream.close_count == 1 +@pytest.mark.asyncio +async def test_native_messages_split_error_event_is_not_counted() -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + key = MagicMock(spec=ApiKey) + key.hashed_key = "messages-split-error" + session = MagicMock() + session.get = AsyncMock(return_value=key) + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0}) + terminal_outcome = TerminalOutcomeContext( + outcome_id="messages-split-error", + model_identifier="test-model", + ) + + async def native_chunks() -> AsyncGenerator[bytes, None]: + yield b'event: error\ndata: {"type":"error",' + yield b'"error":{"message":"failed"}}\n\n' + + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = native_chunks + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await provider.handle_streaming_messages_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=ReservationSnapshot( + release_id="messages-split-error-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ), + terminal_outcome=terminal_outcome, + ) + async for _ in response.body_iterator: + pass + + adjust.assert_awaited_once() + assert adjust.await_args is not None + assert adjust.await_args.kwargs["terminal_outcome"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "split_frames,start_usage,delta_usage,charged,tokens,sources", + [ + ( + False, + {"input_tokens": 10, "output_tokens": 0}, + {"output_tokens": 5}, + 15, + (10, 5), + ("reported", "reported"), + ), + ( + True, + {"input_tokens": 10, "output_tokens": 0}, + {"output_tokens": 5}, + 3, + (10, 5), + ("reported", "reported"), + ), + (False, {}, {}, 3, (3, 0), ("estimated", "estimated")), + (True, {}, {"output_tokens": 5}, 3, (3, 5), ("estimated", "reported")), + ], +) +async def test_native_messages_stats_ignore_network_chunk_boundaries( + split_frames: bool, + start_usage: dict[str, int], + delta_usage: dict[str, int], + charged: int, + tokens: tuple[int, int], + sources: tuple[str, str], +) -> None: + engine = await _engine() + + @asynccontextmanager + async def sessions() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + frames = [ + f'event: message_start\ndata: {{"type":"message_start","message":{{"model":"test-model","usage":{json.dumps(start_usage)}}}}}\n\n'.encode(), + f'event: message_delta\ndata: {{"type":"message_delta","delta":{{"stop_reason":"end_turn"}},"usage":{json.dumps(delta_usage)}}}\n\n'.encode(), + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', + ] + + async def chunks() -> AsyncGenerator[bytes, None]: + for frame in frames: + if split_frames: + boundary = frame.index(b"data: ") + 17 + yield frame[:boundary] + yield frame[boundary:] + else: + yield frame + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = chunks + writer = MagicMock() + provider = BaseUpstreamProvider("https://unused.example", "test-key") + try: + async with sessions() as session: + key = ApiKey(hashed_key="messages-stats", balance=1_000) + session.add(key) + await session.commit() + await pay_for_request(key, 100, session) + reservation = await get_reservation_snapshot(key, session) + + with ( + patch("routstr.upstream.base.create_session", sessions), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + patch("routstr.core.terminal_outcomes.terminal_outcome_writer", writer), + patch( + "routstr.payment.cost_calculation._get_pricing_rates", + return_value=(1_000.0, 1_000.0, 1_000.0, 1_000.0, "configured"), + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005 + ), + patch("routstr.auth.ROUTSTR_FEE_PERCENT", 0), + patch("routstr.upstream.count_tokens._count_with_litellm", return_value=3), + patch( + "routstr.upstream.count_tokens._count_text_with_litellm", return_value=0 + ), + ): + response = await provider.handle_streaming_messages_completion( + upstream_response, + key, + 100, + reservation_snapshot=reservation, + terminal_outcome=TerminalOutcomeContext( + "messages-request", "test-model" + ), + ) + async for _ in response.body_iterator: + pass + + # Billing keeps its existing parser and charge; stats retain reported usage. + async with sessions() as session: + stored_key = await session.get(ApiKey, key.hashed_key) + assert stored_key is not None + assert stored_key.balance == 1_000 - charged + assert stored_key.reserved_balance == 0 + writer.submit.assert_called_once() + outcome = writer.submit.call_args.args[0] + assert outcome.revenue_msats == charged + assert (outcome.input_tokens, outcome.output_tokens) == tokens + assert (outcome.input_source, outcome.output_source) == sources + writer.declare_loss.assert_not_called() + finally: + await engine.dispose() + + +def test_anthropic_stop_reason_is_terminal_before_transport_failure() -> None: + terminal_outcome = TerminalOutcomeContext( + outcome_id="messages-stop-reason", + model_identifier="test-model", + ) + state = _TerminalOutcomeState(terminal_outcome) + + state.observe( + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"output_tokens": 2}, + } + ) + state.mark_transport_failure() + + assert state.settlement_context() is terminal_outcome + + +def test_responses_output_limit_is_terminal_before_transport_failure() -> None: + terminal_outcome = TerminalOutcomeContext( + outcome_id="responses-incomplete", + model_identifier="test-model", + ) + state = _TerminalOutcomeState(terminal_outcome) + + state.observe({"type": "response.incomplete", "response": {"status": "incomplete"}}) + state.mark_transport_failure() + + assert state.settlement_context() is terminal_outcome + + +def test_stream_cut_inside_a_character_does_not_raise() -> None: + state = _TerminalOutcomeState( + TerminalOutcomeContext(outcome_id="cut", model_identifier="test-model") + ) + tail = 'data: {"delta":{"text":"日本'.encode()[:-1] + + assert _observe_terminal_sse_bytes(state, b"", tail, final=True) == b"" + assert state.settlement_context() is None + + @pytest.mark.asyncio async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None: engine = await _engine() @@ -1019,9 +1359,10 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> @pytest.mark.asyncio -async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() -> ( - None -): +@pytest.mark.parametrize("terminal_marker_seen", [False, True]) +async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat( + terminal_marker_seen: bool, +) -> None: """A client abort releases the hold after charging only estimated usage. Starlette closes the response generator (``aclose``) on disconnect. The @@ -1045,7 +1386,10 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() async def aiter_bytes() -> AsyncGenerator[bytes, None]: # A live stream that never sends a usage chunk or [DONE]; the client # disconnects after the first delta. - yield b'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n' + finish_reason = b',"finish_reason":"stop"' if terminal_marker_seen else b"" + yield ( + b'data: {"choices":[{"delta":{"content":"hi"}' + finish_reason + b"}]}\n\n" + ) yield b'data: {"choices":[{"delta":{"content":" there"}}]}\n\n' upstream_response = MagicMock( @@ -1073,6 +1417,11 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() {"model": model.id, "messages": [{"role": "user", "content": "hi"}]} ).encode() + terminal_outcome = TerminalOutcomeContext( + outcome_id="disconnect-outcome", + model_identifier=model.id, + ) + record_outcome = MagicMock() try: with ( patch( @@ -1083,6 +1432,7 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() "routstr.upstream.base.adjust_payment_for_tokens", auth_module.adjust_payment_for_tokens, ), + patch("routstr.auth.record_terminal_outcome", record_outcome), patch("routstr.upstream.count_tokens._count_with_litellm", return_value=3), patch( "routstr.upstream.count_tokens._count_text_with_litellm", @@ -1100,6 +1450,7 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() model_obj=model, reservation_snapshot=snapshot, request_body=request_body, + terminal_outcome=terminal_outcome, ) iterator = cast(AsyncGenerator[bytes, None], response.body_iterator) await iterator.__anext__() # first chunk reaches the client @@ -1118,6 +1469,19 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() # 3 input tokens × 10 msats + 2 output tokens × 20 msats = 70 msats. assert final_key.total_spent == 70 assert final_key.balance == 930 + if terminal_marker_seen: + record_outcome.assert_called_once() + recorded_context = record_outcome.call_args.args[0] + assert recorded_context.outcome_id == terminal_outcome.outcome_id + assert recorded_context.model_identifier == terminal_outcome.model_identifier + assert recorded_context.input_source == "estimated" + assert recorded_context.output_source == "estimated" + assert recorded_context.cache_read_source == "missing" + assert recorded_context.cache_creation_source == "missing" + assert record_outcome.call_args.kwargs["revenue_msats"] == 70 + else: + # Billing settles, but an early disconnect is not a completed outcome. + record_outcome.assert_not_called() # The heartbeat is gone — no forever-renewing task on an abandoned request. assert snapshot.release_id not in auth_module._reservation_heartbeats await engine.dispose() diff --git a/tests/unit/test_streaming_sse_providers.py b/tests/unit/test_streaming_sse_providers.py index afbf22ba..37d090e5 100644 --- a/tests/unit/test_streaming_sse_providers.py +++ b/tests/unit/test_streaming_sse_providers.py @@ -20,12 +20,13 @@ comment ever reaches the client. That invariant is exactly what the buggy import json from collections.abc import AsyncGenerator -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest from routstr.auth import ReservationSnapshot from routstr.core.db import ApiKey +from routstr.core.terminal_outcomes import TerminalOutcomeContext from routstr.upstream import base from routstr.upstream.base import BaseUpstreamProvider @@ -43,7 +44,10 @@ def _make_response(chunks: list[bytes]) -> MagicMock: async def _drive( - chunks: list[bytes], requested_model: str | None = None + chunks: list[bytes], + requested_model: str | None = None, + terminal_outcome: TerminalOutcomeContext | None = None, + adjustment: AsyncMock | None = None, ) -> list[bytes]: """Run the real streaming generator over ``chunks`` and collect output bytes.""" provider = BaseUpstreamProvider( @@ -54,7 +58,7 @@ async def _drive( key.hashed_key = "test_hash" key.balance = 1000 - base.adjust_payment_for_tokens = AsyncMock( + adjustment = adjustment or AsyncMock( return_value={"total_usd": 0.1, "total_msats": 100} ) mock_session = MagicMock() @@ -62,27 +66,30 @@ async def _drive( mock_ctx = MagicMock() mock_ctx.__aenter__ = AsyncMock(return_value=mock_session) mock_ctx.__aexit__ = AsyncMock(return_value=None) - base.create_session = MagicMock(return_value=mock_ctx) - - streaming_response = await provider.handle_streaming_chat_completion( - response=_make_response(chunks), - key=key, - max_cost_for_model=100, - requested_model=requested_model, - reservation_snapshot=ReservationSnapshot( - release_id="test-release", - key_hash="test_hash", - billing_key_hash="test_hash", - reserved_msats=100, - ), - ) - out: list[bytes] = [] - async for chunk in streaming_response.body_iterator: - if isinstance(chunk, str): - out.append(chunk.encode()) - else: - out.append(bytes(chunk)) + with ( + patch.object(base, "adjust_payment_for_tokens", adjustment), + patch.object(base, "create_session", MagicMock(return_value=mock_ctx)), + ): + streaming_response = await provider.handle_streaming_chat_completion( + response=_make_response(chunks), + key=key, + max_cost_for_model=100, + requested_model=requested_model, + terminal_outcome=terminal_outcome, + reservation_snapshot=ReservationSnapshot( + release_id="test-release", + key_hash="test_hash", + billing_key_hash="test_hash", + reserved_msats=100, + ), + ) + + async for chunk in streaming_response.body_iterator: + if isinstance(chunk, str): + out.append(chunk.encode()) + else: + out.append(bytes(chunk)) return out @@ -405,7 +412,16 @@ async def test_truncated_json_tail_on_connection_close() -> None: b'data: {"id":"x","choices":[{"delta":{"content":"ok"}}]}\n\n', b'data: {"id":"x","choices":[{"delta":{"con', # connection dies here ] - out = await _drive(chunks) + terminal_outcome = TerminalOutcomeContext( + outcome_id="truncated-stream", + model_identifier="test-model", + ) + adjustment = AsyncMock(return_value={"total_usd": 0.1, "total_msats": 100}) + out = await _drive( + chunks, + terminal_outcome=terminal_outcome, + adjustment=adjustment, + ) objs = _assert_clean(out) # raises if the partial tail leaked as a data frame contents = [ c["delta"]["content"] @@ -417,3 +433,5 @@ async def test_truncated_json_tail_on_connection_close() -> None: # entirely (no second delta), and _assert_clean above guarantees nothing # non-JSON ever reached the client. assert contents == ["ok"] + assert adjustment.await_args is not None + assert adjustment.await_args.kwargs["terminal_outcome"] is None diff --git a/tests/unit/test_terminal_outcome_migration.py b/tests/unit/test_terminal_outcome_migration.py new file mode 100644 index 00000000..785ab3f3 --- /dev/null +++ b/tests/unit/test_terminal_outcome_migration.py @@ -0,0 +1,190 @@ +from __future__ import annotations + +import os +import sqlite3 +import subprocess +import sys +from pathlib import Path + +import pytest + +REVISION = "c8e4a1f2b3d5" +PREVIOUS_REVISION = "e4c7a1b9d520" + + +def _run_alembic(root: Path, database_url: str, command: str, revision: str) -> None: + env = os.environ.copy() + env["DATABASE_URL"] = database_url + subprocess.run( + [sys.executable, "-m", "alembic", command, revision], + cwd=root, + env=env, + check=True, + capture_output=True, + text=True, + ) + + +def _table_names(connection: sqlite3.Connection) -> set[str]: + return { + row[0] + for row in connection.execute( + "SELECT name FROM sqlite_master WHERE type = 'table'" + ) + } + + +def _columns(connection: sqlite3.Connection, table: str) -> dict[str, tuple[str, int]]: + return { + row[1]: (row[2], row[3]) + for row in connection.execute(f"PRAGMA table_info({table})") + } + + +def test_terminal_outcome_migration_round_trips(tmp_path: Path) -> None: + root = Path(__file__).resolve().parents[2] + database_path = tmp_path / "terminal-outcomes.db" + database_url = f"sqlite+aiosqlite:///{database_path}" + ledger_tables = { + "terminal_outcomes", + "terminal_outcome_epochs", + "terminal_outcome_writer_runs", + } + + _run_alembic(root, database_url, "upgrade", PREVIOUS_REVISION) + with sqlite3.connect(database_path) as connection: + assert not ledger_tables & _table_names(connection) + + _run_alembic(root, database_url, "upgrade", REVISION) + with sqlite3.connect(database_path) as connection: + assert ledger_tables <= _table_names(connection) + + outcome_columns = _columns(connection, "terminal_outcomes") + assert set(outcome_columns) == { + "outcome_id", + "terminal_at_ms", + "terminal_day", + "model_identifier", + "served_model_identifier", + "pricing_source", + "input_source", + "output_source", + "cache_read_source", + "cache_creation_source", + "input_tokens", + "output_tokens", + "cache_read_input_tokens", + "cache_creation_input_tokens", + "revenue_msats", + } + for column in ( + "terminal_at_ms", + "input_tokens", + "output_tokens", + "cache_read_input_tokens", + "cache_creation_input_tokens", + "revenue_msats", + ): + assert outcome_columns[column] == ("BIGINT", 1) + assert outcome_columns["terminal_day"] == ("DATE", 1) + assert outcome_columns["model_identifier"] == ("VARCHAR", 0) + outcome_index = connection.execute( + "PRAGMA index_info(ix_terminal_outcomes_terminal_day_terminal_at_ms)" + ).fetchall() + assert [row[2] for row in outcome_index] == [ + "terminal_day", + "terminal_at_ms", + ] + + epoch_columns = _columns(connection, "terminal_outcome_epochs") + assert set(epoch_columns) == { + "epoch", + "coverage_start_day", + "coverage_end_day", + "current_slot", + } + assert epoch_columns["epoch"] == ("BIGINT", 1) + assert epoch_columns["coverage_start_day"] == ("DATE", 1) + assert epoch_columns["coverage_end_day"] == ("DATE", 0) + + run_columns = _columns(connection, "terminal_outcome_writer_runs") + assert set(run_columns) == { + "run_id", + "status", + "started_at_ms", + "heartbeat_at_ms", + "flushed_through_ms", + "closed_at_ms", + "loss_day", + } + for column in ("started_at_ms", "heartbeat_at_ms"): + assert run_columns[column] == ("BIGINT", 1) + assert run_columns["closed_at_ms"] == ("BIGINT", 0) + run_index = connection.execute( + "PRAGMA index_info(ix_terminal_outcome_writer_runs_status_heartbeat)" + ).fetchall() + assert [row[2] for row in run_index] == ["status", "heartbeat_at_ms"] + + outcome_sql = connection.execute( + "SELECT sql FROM sqlite_master WHERE name = 'terminal_outcomes'" + ).fetchone() + epoch_sql = connection.execute( + "SELECT sql FROM sqlite_master WHERE name = 'terminal_outcome_epochs'" + ).fetchone() + run_sql = connection.execute( + "SELECT sql FROM sqlite_master WHERE name = 'terminal_outcome_writer_runs'" + ).fetchone() + assert outcome_sql is not None and outcome_sql[0].count("CHECK") == 1 + assert epoch_sql is not None and epoch_sql[0].count("CHECK") == 3 + assert run_sql is not None and run_sql[0].count("CHECK") == 4 + + connection.execute( + "INSERT INTO terminal_outcome_epochs VALUES (0, '2026-09-01', NULL, 1)" + ) + connection.execute( + "INSERT INTO terminal_outcome_writer_runs VALUES " + "('run-1', 'active', 1, 1, NULL, NULL, NULL)" + ) + connection.execute( + "INSERT INTO terminal_outcome_writer_runs VALUES " + "('run-2', 'active', 1, 1, NULL, NULL, NULL)" + ) + connection.execute( + "INSERT INTO terminal_outcomes " + "(outcome_id, terminal_at_ms, terminal_day, model_identifier, " + "input_tokens, output_tokens, cache_read_input_tokens, cache_creation_input_tokens, revenue_msats) VALUES " + "('request-1', 1, '2026-08-31', 'author/model', " + "10, 5, 0, 0, 1999)" + ) + with pytest.raises(sqlite3.IntegrityError): + connection.execute( + "INSERT INTO terminal_outcomes " + "(outcome_id, terminal_at_ms, terminal_day, model_identifier, " + "input_tokens, output_tokens, cache_read_input_tokens, cache_creation_input_tokens, revenue_msats) VALUES " + "('request-invalid', 2, '2026-08-31', 'author/model', " + "-1, 0, 0, 0, 0)" + ) + with pytest.raises(sqlite3.IntegrityError): + connection.execute( + "INSERT INTO terminal_outcome_epochs VALUES (1, '2026-09-02', NULL, 1)" + ) + with pytest.raises(sqlite3.IntegrityError): + connection.execute( + "INSERT INTO terminal_outcome_writer_runs VALUES " + "('run-invalid', 'unknown', 1, 1, NULL, NULL, NULL)" + ) + connection.commit() + + _run_alembic(root, database_url, "downgrade", PREVIOUS_REVISION) + with sqlite3.connect(database_path) as connection: + tables = _table_names(connection) + assert not ledger_tables & tables + assert "api_keys" in tables + + _run_alembic(root, database_url, "upgrade", REVISION) + with sqlite3.connect(database_path) as connection: + assert ledger_tables <= _table_names(connection) + for table in ledger_tables: + assert connection.execute(f"SELECT COUNT(*) FROM {table}").fetchone() == ( + 0, + ) diff --git a/tests/unit/test_terminal_outcomes.py b/tests/unit/test_terminal_outcomes.py new file mode 100644 index 00000000..b1506cae --- /dev/null +++ b/tests/unit/test_terminal_outcomes.py @@ -0,0 +1,790 @@ +from __future__ import annotations + +import asyncio +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +from dataclasses import dataclass +from datetime import UTC, date, datetime, timedelta +from pathlib import Path + +import pytest +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel, col, select +from sqlmodel.ext.asyncio.session import AsyncSession + +import routstr.core.terminal_outcomes as outcomes_module +from routstr.core.db import ( + TerminalOutcome, + TerminalOutcomeEpoch, + TerminalOutcomeWriterRun, +) +from routstr.core.terminal_outcome_writer import ( + SessionFactory, + TerminalOutcomeWriter, + _PersistResult, + _QueuedOutcome, +) +from routstr.core.terminal_outcomes import ( + TerminalOutcomeContext, + cashu_retained_msats, + record_terminal_outcome, +) + + +@dataclass +class MutableClock: + value: int + + def __call__(self) -> int: + return self.value + + +@pytest.fixture +async def ledger( + tmp_path: Path, +) -> AsyncGenerator[tuple[AsyncEngine, SessionFactory], None]: + engine = create_async_engine( + f"sqlite+aiosqlite:///{tmp_path / 'terminal-outcomes.db'}" + ) + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + + @asynccontextmanager + async def sessions() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + yield engine, sessions + await engine.dispose() + + +def _timestamp(day: date, hour: int = 12) -> int: + return int( + datetime(day.year, day.month, day.day, hour, tzinfo=UTC).timestamp() * 1000 + ) + + +def _record( + outcome_id: str, + timestamp_ms: int, + *, + revenue_msats: int = 2000, +) -> None: + record_terminal_outcome( + TerminalOutcomeContext(outcome_id, "author/model"), + input_tokens=10, + output_tokens=5, + cache_read_input_tokens=2, + cache_creation_input_tokens=1, + revenue_msats=revenue_msats, + terminal_at_ms=timestamp_ms, + ) + + +async def test_writer_persists_one_immutable_outcome_and_closes_own_run( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + day = date(2026, 8, 31) + timestamp = _timestamp(day) + writer = TerminalOutcomeWriter( + session_factory=sessions, + retry_seconds=0.01, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=MutableClock(timestamp), + ) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + + assert await writer.start() + _record("request-1", timestamp) + _record("request-1", timestamp) + assert await writer.flush(timeout=1) + + async with sessions() as session: + stored = (await session.exec(select(TerminalOutcome))).all() + epochs = (await session.exec(select(TerminalOutcomeEpoch))).all() + runs = (await session.exec(select(TerminalOutcomeWriterRun))).all() + assert len(stored) == 1 + assert stored[0].terminal_day == day + assert stored[0].input_source == "missing" + assert stored[0].revenue_msats == 2000 + assert len(epochs) == 1 + assert epochs[0].coverage_start_day == day + timedelta(days=1) + assert len(runs) == 1 and runs[0].status == "active" + + assert await writer.stop(timeout=1) + async with sessions() as session: + run = (await session.exec(select(TerminalOutcomeWriterRun))).one() + assert run.status == "clean" + assert run.closed_at_ms == timestamp + + +async def test_writer_allocates_after_latest_closed_epoch( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + day = date(2026, 8, 31) + async with sessions() as session: + session.add( + TerminalOutcomeEpoch( + epoch=0, + coverage_start_day=day - timedelta(days=2), + coverage_end_day=day - timedelta(days=1), + current_slot=None, + ) + ) + await session.commit() + writer = TerminalOutcomeWriter( + session_factory=sessions, + clock=MutableClock(_timestamp(day)), + ) + + assert await writer.start() + assert await writer.stop(timeout=1) + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert [ + (epoch.epoch, epoch.coverage_start_day, epoch.current_slot) for epoch in epochs + ] == [ + (0, day - timedelta(days=2), None), + (1, day + timedelta(days=1), 1), + ] + + +async def test_conflicting_duplicate_rotates_epoch_without_overwriting( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + day = date(2026, 8, 31) + timestamp = _timestamp(day) + writer = TerminalOutcomeWriter( + session_factory=sessions, + retry_seconds=0.01, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=MutableClock(timestamp), + ) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + + assert await writer.start() + _record("request-conflict", timestamp, revenue_msats=1000) + assert await writer.flush(timeout=1) + _record("request-conflict", timestamp, revenue_msats=2000) + assert await writer.flush(timeout=1) + + async with sessions() as session: + stored = (await session.exec(select(TerminalOutcome))).one() + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert stored.revenue_msats == 1000 + assert [epoch.epoch for epoch in epochs] == [0, 1] + assert epochs[0].coverage_end_day == day - timedelta(days=1) + assert epochs[1].coverage_start_day == day + timedelta(days=1) + assert epochs[1].current_slot == 1 + assert await writer.stop(timeout=1) + + +async def test_temporary_insert_failure_retries_retained_item_without_rotation( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + day = date(2026, 8, 31) + timestamp = _timestamp(day) + writer = TerminalOutcomeWriter( + session_factory=sessions, + retry_seconds=0.01, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=MutableClock(timestamp), + ) + original = writer._persist_once + attempts = 0 + + async def fail_once( + queued: _QueuedOutcome, + ) -> _PersistResult: + nonlocal attempts + attempts += 1 + if attempts == 1: + return _PersistResult.RETRY + return await original(queued) + + monkeypatch.setattr(writer, "_persist_once", fail_once) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + + assert await writer.start() + _record("request-retry", timestamp) + assert await writer.flush(timeout=1) + async with sessions() as session: + assert len((await session.exec(select(TerminalOutcome))).all()) == 1 + assert len((await session.exec(select(TerminalOutcomeEpoch))).all()) == 1 + assert attempts == 2 + assert await writer.stop(timeout=1) + + +async def test_queue_overflow_is_nonblocking_and_rotates_epoch( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + queued_day = date(2026, 8, 31) + overflow_day = queued_day + timedelta(days=1) + clock = MutableClock(_timestamp(queued_day)) + entered = asyncio.Event() + release = asyncio.Event() + writer = TerminalOutcomeWriter( + session_factory=sessions, + queue_size=1, + retry_seconds=0.01, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=clock, + ) + + async def blocked_insert( + queued: _QueuedOutcome, + ) -> _PersistResult: + entered.set() + await release.wait() + return _PersistResult.STORED + + monkeypatch.setattr(writer, "_persist_once", blocked_insert) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + + assert await writer.start() + _record("request-in-flight", _timestamp(queued_day)) + await asyncio.wait_for(entered.wait(), timeout=1) + _record("request-queued", _timestamp(queued_day)) + clock.value = _timestamp(overflow_day) + _record("request-overflow", clock.value) + assert writer.loss_pending + release.set() + assert await writer.flush(timeout=1) + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert [epoch.epoch for epoch in epochs] == [0, 1] + assert epochs[0].coverage_end_day == queued_day - timedelta(days=1) + assert epochs[1].coverage_start_day == overflow_day + timedelta(days=1) + assert await writer.stop(timeout=1) + + +async def test_live_workers_do_not_rotate_and_clean_stop_closes_only_owner( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + timestamp = _timestamp(date(2026, 8, 31)) + first = TerminalOutcomeWriter( + session_factory=sessions, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=MutableClock(timestamp), + ) + second = TerminalOutcomeWriter( + session_factory=sessions, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + clock=MutableClock(timestamp), + ) + assert await first.start() + assert await second.start() + async with sessions() as session: + assert len((await session.exec(select(TerminalOutcomeEpoch))).all()) == 1 + assert len((await session.exec(select(TerminalOutcomeWriterRun))).all()) == 2 + + first_run_id = first._run_id + second_run_id = second._run_id + assert await first.stop(timeout=1) + async with sessions() as session: + first_run = await session.get(TerminalOutcomeWriterRun, first_run_id) + second_run = await session.get(TerminalOutcomeWriterRun, second_run_id) + assert first_run is not None and first_run.status == "clean" + assert second_run is not None and second_run.status == "active" + assert await second.stop(timeout=1) + + +async def test_worker_that_starts_collecting_late_voids_the_day_it_missed( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + day = date(2026, 8, 31) + first = TerminalOutcomeWriter( + session_factory=sessions, clock=MutableClock(_timestamp(day, 23)) + ) + # Another worker learns of the opt-in just after midnight, while serving. + late = TerminalOutcomeWriter( + session_factory=sessions, + lease_timeout_seconds=7200, + clock=MutableClock(_timestamp(day + timedelta(days=1), 0)), + ) + assert await first.start(serving=True) + assert await late.start(serving=True) + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert [(row.coverage_start_day, row.coverage_end_day) for row in epochs] == [ + (day + timedelta(days=1), day), + (day + timedelta(days=2), None), + ] + assert await first.stop(timeout=1) + assert await late.stop(timeout=1) + + +async def test_stale_run_rotates_once_with_concurrent_recovery( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + lost_day = date(2026, 8, 30) + recovery_day = date(2026, 8, 31) + lost_at = _timestamp(lost_day) + recovered_at = _timestamp(recovery_day) + async with sessions() as session: + session.add( + TerminalOutcomeEpoch( + epoch=0, + coverage_start_day=lost_day + timedelta(days=1), + current_slot=1, + ) + ) + session.add( + TerminalOutcomeWriterRun( + run_id="lost-run", + status="lost", + started_at_ms=lost_at, + heartbeat_at_ms=lost_at, + closed_at_ms=lost_at, + loss_day=lost_day, + ) + ) + await session.commit() + + first = TerminalOutcomeWriter( + session_factory=sessions, clock=MutableClock(recovered_at) + ) + second = TerminalOutcomeWriter( + session_factory=sessions, clock=MutableClock(recovered_at) + ) + await asyncio.gather(first._recover_pending_runs(), second._recover_pending_runs()) + + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + lost_run = await session.get(TerminalOutcomeWriterRun, "lost-run") + assert [epoch.epoch for epoch in epochs] == [0, 1] + assert epochs[0].coverage_end_day == lost_day - timedelta(days=1) + assert epochs[1].coverage_start_day == recovery_day + timedelta(days=1) + assert lost_run is not None and lost_run.status == "recovered" + + +async def test_recovery_does_not_claim_a_late_earlier_loss( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + selected_day = date(2026, 8, 30) + late_day = date(2026, 8, 29) + recovered_at = _timestamp(date(2026, 8, 31)) + async with sessions() as session: + session.add( + TerminalOutcomeEpoch( + epoch=0, + coverage_start_day=selected_day, + current_slot=1, + ) + ) + session.add( + TerminalOutcomeWriterRun( + run_id="selected-loss", + status="lost", + started_at_ms=_timestamp(selected_day), + heartbeat_at_ms=_timestamp(selected_day), + closed_at_ms=_timestamp(selected_day), + loss_day=selected_day, + ) + ) + await session.commit() + + writer = TerminalOutcomeWriter( + session_factory=sessions, clock=MutableClock(recovered_at) + ) + original_stage_rotation = writer._stage_rotation + + async def inject_late_loss( + session: AsyncSession, + current: TerminalOutcomeEpoch, + lost_day: date, + ) -> int | None: + next_epoch = await original_stage_rotation(session, current, lost_day) + session.add( + TerminalOutcomeWriterRun( + run_id="late-loss", + status="lost", + started_at_ms=_timestamp(late_day), + heartbeat_at_ms=_timestamp(late_day), + closed_at_ms=_timestamp(late_day), + loss_day=late_day, + ) + ) + return next_epoch + + monkeypatch.setattr(writer, "_stage_rotation", inject_late_loss) + await writer._recover_pending_runs() + + async with sessions() as session: + selected = await session.get(TerminalOutcomeWriterRun, "selected-loss") + late = await session.get(TerminalOutcomeWriterRun, "late-loss") + assert selected is not None and selected.status == "recovered" + assert late is not None and late.status == "lost" + + +def test_record_wrapper_never_raises_on_invalid_or_failed_submission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class StubWriter: + def __init__(self) -> None: + self.losses: list[str] = [] + self.submissions: list[_QueuedOutcome] = [] + + def declare_loss(self, reason: str, lost_day: date | None = None) -> None: + self.losses.append(reason) + + def submit(self, outcome: _QueuedOutcome) -> bool: + self.submissions.append(outcome) + if outcome.outcome_id == "request-submit-error": + raise RuntimeError("submission failed") + return True + + writer = StubWriter() + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + record_terminal_outcome( + TerminalOutcomeContext("request-invalid", "author/model"), + input_tokens=-1, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=0, + ) + record_terminal_outcome( + TerminalOutcomeContext("request-unstorable", "author/model"), + input_tokens=10**30, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=0, + ) + record_terminal_outcome( + TerminalOutcomeContext("request-unreadable-usage", "author/model"), + input_tokens=0, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=105, + usage={"input_tokens": 100, "input_tokens_details": {"cached_tokens": "1e309"}}, + ) + record_terminal_outcome( + TerminalOutcomeContext("request-unknown-model", None), + input_tokens=1, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=0, + ) + record_terminal_outcome( + TerminalOutcomeContext("request-submit-error", "author/model"), + input_tokens=1, + output_tokens=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=0, + ) + assert writer.losses == [ + "invalid settled terminal outcome", + "invalid settled terminal outcome", + "terminal outcome submission raised", + "terminal outcome submission raised", + ] + assert writer.submissions[0].model_identifier is None + + +def test_cashu_retained_msats_uses_exact_persisted_units( + monkeypatch: pytest.MonkeyPatch, +) -> None: + losses: list[str] = [] + monkeypatch.setattr(outcomes_module, "mark_terminal_outcome_loss", losses.append) + assert cashu_retained_msats(10, "sat", 8) == 2000 + assert cashu_retained_msats(2500, "msat", 501) == 1999 + assert cashu_retained_msats(1, "sat", 2) is None + assert losses == ["invalid Cashu retained value"] + + +async def test_collection_pause_excludes_disabled_days_after_restart( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + day = date(2026, 9, 1) + clock = MutableClock(_timestamp(day)) + writer = TerminalOutcomeWriter(session_factory=sessions, clock=clock) + assert await writer.start() + clock.value = _timestamp(day + timedelta(days=3)) + assert await writer.stop(timeout=1, close_coverage=True) + clock.value = _timestamp(day + timedelta(days=6)) + assert await writer.start() + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert epochs[0].coverage_start_day == day + timedelta(days=1) + assert epochs[0].coverage_end_day == day + timedelta(days=2) + assert epochs[1].coverage_start_day == day + timedelta(days=7) + assert await writer.stop(timeout=1) + + +@pytest.mark.parametrize("charge", [0, 1200]) +async def test_encrypted_settlement_persists_once_after_successful_debit( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, + charge: int, +) -> None: + from routstr.auth import get_reservation_snapshot, pay_for_request + from routstr.core.db import ApiKey + from routstr.upstream.ehbp import finalize_ehbp_actual_cost_payment + + _, sessions = ledger + writer = TerminalOutcomeWriter(session_factory=sessions) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0) + assert await writer.start() + async with sessions() as session: + key = ApiKey(hashed_key=f"settlement-{charge}", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3000, session) + reservation = await get_reservation_snapshot(key, session) + context = TerminalOutcomeContext( + outcome_id=f"request-{charge}", + model_identifier="canonical/model", + served_model_identifier="provider-model-v2", + ) + cost_info = { + "total_msats": charge, + "input_tokens": 10, + "output_tokens": 20, + "input_source": "reported", + "output_source": "reported", + "cache_read_source": "missing", + "cache_creation_source": "missing", + "pricing_source": "configured", + } + args = ( + key, + session, + 3000, + "provider-model-v2", + cost_info, + reservation, + context, + ) + assert await finalize_ehbp_actual_cost_payment(*args) == charge + with pytest.raises(RuntimeError, match="reservation record does not match"): + await finalize_ehbp_actual_cost_payment(*args) + assert key.balance == 10_000 - charge + assert key.reserved_balance == 0 + assert await writer.flush(timeout=1) + async with sessions() as session: + stored = (await session.exec(select(TerminalOutcome))).one() + assert stored.revenue_msats == charge + assert stored.input_tokens == 10 and stored.output_tokens == 20 + assert stored.model_identifier == "canonical/model" + assert stored.served_model_identifier == "provider-model-v2" + assert stored.input_source == stored.output_source == "reported" + assert stored.cache_read_source == stored.cache_creation_source == "missing" + assert stored.pricing_source == "configured" + assert await writer.stop(timeout=1) + + +async def test_writer_failure_does_not_fail_settlement_and_marks_gap( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.auth import get_reservation_snapshot, pay_for_request + from routstr.core.db import ApiKey + from routstr.upstream.ehbp import finalize_ehbp_actual_cost_payment + + _, sessions = ledger + clock = MutableClock(_timestamp(date(2026, 9, 18))) + writer = TerminalOutcomeWriter(session_factory=sessions, clock=clock) + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0) + assert await writer.start() + + def unavailable(_: _QueuedOutcome) -> bool: + raise OSError("synthetic stats storage failure") + + monkeypatch.setattr(writer, "submit", unavailable) + async with sessions() as session: + key = ApiKey(hashed_key="settlement-during-stats-failure", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3000, session) + reservation = await get_reservation_snapshot(key, session) + charged = await finalize_ehbp_actual_cost_payment( + key, + session, + 3000, + "model", + {"total_msats": 1200}, + reservation, + TerminalOutcomeContext("request-storage-failure", "model"), + ) + assert charged == 1200 + assert key.balance == 8800 and key.reserved_balance == 0 + assert await writer.flush() + async with sessions() as session: + assert not (await session.exec(select(TerminalOutcome))).all() + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert len(epochs) == 2 + assert epochs[0].coverage_end_day == date(2026, 9, 17) + assert epochs[1].coverage_start_day == date(2026, 9, 19) + assert await writer.stop(timeout=1) + + +async def test_unclean_restart_preserves_days_before_last_durable_flush( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + first_day = date(2026, 9, 1) + last_day = date(2026, 9, 10) + clock = MutableClock(_timestamp(first_day)) + writer = TerminalOutcomeWriter( + session_factory=sessions, + clock=clock, + heartbeat_seconds=0.05, + lease_timeout_seconds=1, + ) + assert await writer.start() + clock.value = _timestamp(last_day) + assert await writer.flush(timeout=1) + assert writer._task is not None + writer._task.cancel() + with pytest.raises(asyncio.CancelledError): + await writer._task + assert not await writer.stop(timeout=1) + clock.value += 2000 + assert await writer.start() + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + try: + assert epochs[0].coverage_end_day == last_day - timedelta(days=1) + assert epochs[1].coverage_start_day == last_day + timedelta(days=1) + finally: + assert await writer.stop(timeout=1) + + +@pytest.mark.parametrize("fresh_process", [False, True]) +async def test_failed_writer_start_cannot_backfill_missed_days_as_zero( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, + fresh_process: bool, +) -> None: + _, sessions = ledger + day = date(2026, 9, 1) + clock = MutableClock(_timestamp(day)) + writer = TerminalOutcomeWriter(session_factory=sessions, clock=clock) + original = writer._create_run + + async def unavailable(*args: object) -> None: + raise OSError("synthetic stats startup failure") + + monkeypatch.setattr(writer, "_create_run", unavailable) + assert not await writer.start() + assert writer.loss_pending + monkeypatch.setattr(writer, "_create_run", original) + if fresh_process: + writer = TerminalOutcomeWriter(session_factory=sessions, clock=clock) + clock.value = _timestamp(day + timedelta(days=3)) + assert await writer.start() + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + assert epochs[0].coverage_end_day == ( + day if fresh_process else day - timedelta(days=1) + ) + assert epochs[1].coverage_start_day == day + timedelta(days=4) + assert await writer.stop(timeout=1) + + +async def test_disable_closes_coverage_after_background_rotation_is_stopped( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + writer = TerminalOutcomeWriter(session_factory=sessions) + assert await writer.start() + original = writer._close_coverage + + async def close_without_racing_writer(day: date) -> None: + assert not writer.running + await original(day) + + monkeypatch.setattr(writer, "_close_coverage", close_without_racing_writer) + assert await writer.stop(timeout=1, close_coverage=True) + + +async def test_restart_alongside_live_writer_preserves_continuous_coverage( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + first_day = date(2026, 9, 1) + clock = MutableClock(_timestamp(first_day)) + writers = [ + TerminalOutcomeWriter(session_factory=sessions, clock=clock) for _ in range(3) + ] + assert await writers[0].start() + assert await writers[1].start() + clock.value = _timestamp(first_day + timedelta(days=3)) + assert await writers[0].flush(timeout=1) + assert await writers[1].flush(timeout=1) + assert await writers[0].stop(timeout=1) + assert await writers[2].start() + try: + async with sessions() as session: + epochs = (await session.exec(select(TerminalOutcomeEpoch))).all() + assert len(epochs) == 1 + assert epochs[0].coverage_start_day == first_day + timedelta(days=1) + assert epochs[0].coverage_end_day is None + finally: + assert await writers[1].stop(timeout=1) + assert await writers[2].stop(timeout=1) diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 14715f58..f9e15bd0 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -273,6 +273,8 @@ class TestComputeEhbpActualCost: assert result["total_msats"] == 0 assert result["input_tokens"] == 0 assert result["output_tokens"] == 0 + assert result["input_source"] == "missing" + assert result["output_source"] == "missing" @pytest.mark.asyncio async def test_usage_parsed_and_clamped(self) -> None: @@ -308,6 +310,10 @@ class TestComputeEhbpActualCost: assert result["total_tokens"] == 109 assert result["input_msats"] == 10 assert result["output_msats"] == 20 + assert result["input_source"] == "reported" + assert result["output_source"] == "reported" + assert result["cache_read_source"] == "missing" + assert result["cache_creation_source"] == "missing" @pytest.mark.asyncio async def test_cache_fields_propagated(self) -> None: @@ -350,8 +356,11 @@ class TestComputeEhbpActualCost: assert result["total_usd"] == 0.0003 @pytest.mark.asyncio + @pytest.mark.parametrize("input_count,output_count", [(0, 0), (10, 5)]) async def test_unpriceable_usage_does_not_charge_authorization_ceiling( self, + input_count: int, + output_count: int, ) -> None: model_obj = MagicMock() model_obj.id = "llama3-3-70b" @@ -372,13 +381,35 @@ class TestComputeEhbpActualCost: output_tokens=0, ) result = await _compute_ehbp_actual_cost( - "prompt=0,completion=0", + f"prompt={input_count},completion={output_count}", model_obj, 50_000, ) assert result["total_msats"] == 0 - assert result["input_tokens"] == 0 - assert result["output_tokens"] == 0 + assert result["input_tokens"] == input_count + assert result["output_tokens"] == output_count + assert result["input_source"] == "reported" + assert result["output_source"] == "reported" + + @pytest.mark.asyncio + async def test_partially_malformed_usage_preserves_independent_presence( + self, + ) -> None: + model_obj = MagicMock() + model_obj.id = "llama3-3-70b" + model_obj.forwarded_model_id = "llama3-3-70b" + + result = await _compute_ehbp_actual_cost( + "prompt=7,completion=not-a-number", + model_obj, + 50_000, + ) + + assert result["total_msats"] == 0 + assert result["input_tokens"] == 7 + assert result["output_tokens"] == 0 + assert result["input_source"] == "reported" + assert result["output_source"] == "missing" @pytest.mark.asyncio async def test_model_match_no_actual_model_key(self) -> None: @@ -456,6 +487,7 @@ class TestComputeEhbpActualCost: actual_model_obj = MagicMock() actual_model_obj.id = "tinfoil-llama3-3-70b" # client-facing of actual actual_model_obj.forwarded_model_id = "llama3-3-70b" + actual_model_obj.canonical_slug = "meta/llama-3.3-70b" with ( patch( @@ -485,6 +517,7 @@ class TestComputeEhbpActualCost: 100_000, ) assert result["actual_model"] == "llama3-3-70b" + assert result["actual_model_identifier"] == "meta/llama-3.3-70b" assert result["total_msats"] == 60 # calculate_cost called with the actual model's client-facing ID call_args = mock_calc.call_args @@ -748,6 +781,7 @@ class TestComputeEhbpActualCost: 100_000, ) assert "actual_model" not in result + assert result["actual_model_unresolved"] is True # calculate_cost called with the requested model (fallback) call_args = mock_calc.call_args assert call_args[0][0]["model"] == "gpt-oss-120b" diff --git a/tests/unit/test_usage_normalization.py b/tests/unit/test_usage_normalization.py index 66a44dc4..ae3f910c 100644 --- a/tests/unit/test_usage_normalization.py +++ b/tests/unit/test_usage_normalization.py @@ -15,7 +15,12 @@ os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") import pytest -from routstr.payment.usage import NormalizedUsage, normalize_usage +from routstr.payment.usage import ( + NormalizedUsage, + UsageFieldPresence, + normalize_usage, + usage_field_presence, +) # ============================================================================ # The union parser: one canonical shape for all known dialects @@ -190,3 +195,69 @@ def test_normalize_usage_never_negative() -> None: assert result is not None assert result.input_tokens == 0 assert result.cache_read_tokens == 150 + + +def test_usage_presence_counts_explicit_zero_across_dialects() -> None: + presence = usage_field_presence( + { + "prompt_tokens": 0, + "completion_tokens": "0", + "prompt_tokens_details": { + "cached_tokens": 0.0, + "cache_write_tokens": "0", + }, + } + ) + + assert presence == UsageFieldPresence( + input_source="reported", + output_source="reported", + cache_read_source="reported", + cache_creation_source="reported", + ) + + +def test_usage_presence_rejects_missing_and_unparseable_fields() -> None: + presence = usage_field_presence( + { + "input_tokens": None, + "output_tokens": "not-a-number", + "cache_read_input_tokens": -1, + "cache_creation_input_tokens": False, + } + ) + + assert presence == UsageFieldPresence() + + +def test_usage_presence_recognizes_responses_cache_details() -> None: + presence = usage_field_presence( + { + "input_tokens": 12, + "output_tokens": 3, + "input_tokens_details": {"cached_tokens": 0}, + } + ) + + assert presence == UsageFieldPresence( + input_source="reported", + output_source="reported", + cache_read_source="reported", + ) + + +def test_locally_estimated_usage_is_not_reported() -> None: + presence = usage_field_presence( + { + "input_tokens": 12, + "output_tokens": 3, + "estimated": True, + } + ) + + assert presence.sources_dict() == { + "input_source": "estimated", + "output_source": "estimated", + "cache_read_source": "missing", + "cache_creation_source": "missing", + } diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py index 901cc2f6..a7c19a63 100644 --- a/tests/unit/test_x_cashu_cost_sats.py +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -1,6 +1,6 @@ import json import os -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -8,8 +8,12 @@ import pytest os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") os.environ.setdefault("UPSTREAM_API_KEY", "test") +from routstr.core.terminal_outcomes import TerminalOutcomeContext # noqa: E402 from routstr.payment.cost_calculation import CostData # noqa: E402 -from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 +from routstr.upstream.base import ( # noqa: E402 + BaseUpstreamProvider, + _record_x_cashu_terminal_outcome, +) def _make_provider() -> BaseUpstreamProvider: @@ -29,9 +33,35 @@ def _make_cost_data(total_msats: int = 5000) -> CostData: total_usd=0.00025, input_tokens=100, output_tokens=50, + input_source="reported", + output_source="reported", ) +def test_zero_usage_x_cashu_preserves_captured_presence() -> None: + record = MagicMock() + context = TerminalOutcomeContext( + outcome_id="zero-usage", + model_identifier="author/model", + input_source="reported", + output_source="reported", + cache_read_source="missing", + cache_creation_source="missing", + ) + + with patch("routstr.upstream.base.record_terminal_outcome", record): + _record_x_cashu_terminal_outcome( + context, + None, + amount=10, + unit="sat", + ) + + recorded_context = record.call_args.args[0] + assert recorded_context.input_source == recorded_context.output_source == "reported" + assert record.call_args.kwargs["revenue_msats"] == 10_000 + + # --------------------------------------------------------------------------- # Non-streaming (chat completions) # --------------------------------------------------------------------------- @@ -80,24 +110,47 @@ async def test_non_streaming_includes_cost_sats() -> None: async def test_non_streaming_cost_sats_value_rounds_down() -> None: provider = _make_provider() cost_data = _make_cost_data(total_msats=1999) + settlement_order: list[str] = [] + + async def persist_refund(*args: object, **kwargs: object) -> str: + settlement_order.append("refund") + return "cashuA_refund_token" + + def record_outcome(*args: object, **kwargs: object) -> None: + settlement_order.append("record") + + send_refund = AsyncMock(side_effect=persist_refund) + record = MagicMock(side_effect=record_outcome) response_body = {"model": "gpt-4o", "usage": {"prompt_tokens": 10}} content_str = json.dumps(response_body) with ( patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)), - patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")), + patch.object(provider, "send_refund", new=send_refund), + patch("routstr.upstream.base.record_terminal_outcome", record), ): response = await provider.handle_x_cashu_non_streaming_response( content_str=content_str, response=_make_httpx_response(), - amount=10000, - unit="msat", + amount=10, + unit="sat", max_cost_for_model=10000, + request_id="sat-rounding-request", ) body = json.loads(response.body) assert body["usage"]["cost_sats"] == 1 # 1999 // 1000 + send_refund.assert_awaited_once_with( + 8, "sat", None, request_id="sat-rounding-request" + ) + record.assert_called_once() + recorded_context = record.call_args.args[0] + assert recorded_context.input_source == recorded_context.output_source == "reported" + assert recorded_context.cache_read_source == "missing" + assert recorded_context.cache_creation_source == "missing" + assert record.call_args.kwargs["revenue_msats"] == 2000 + assert settlement_order == ["refund", "record"] @pytest.mark.asyncio @@ -218,3 +271,110 @@ async def test_streaming_non_usage_chunks_unmodified() -> None: regular_line_data = json.loads(lines[0][6:]) # regular chunk should not have cost_sats injected assert "cost_sats" not in regular_line_data.get("usage", {}) + + +@pytest.mark.asyncio +async def test_streaming_no_space_error_event_is_not_recorded() -> None: + provider = _make_provider() + record = MagicMock() + + with patch("routstr.upstream.base.record_terminal_outcome", record): + await provider.handle_x_cashu_streaming_response( + content_str='data:{"type":"error","error":{"message":"failed"}}\n\n', + response=_make_httpx_response(), + amount=10, + unit="sat", + max_cost_for_model=10_000, + request_id="no-space-error", + ) + + record.assert_not_called() + + +@pytest.mark.asyncio +async def test_streaming_no_space_lines_keep_existing_billing_parse() -> None: + provider = _make_provider() + cost = AsyncMock(return_value=None) + body = 'data:{"usage":{"prompt_tokens":100,"completion_tokens":50}}\n\n' + + with patch.object(provider, "get_x_cashu_cost", new=cost): + response = await provider.handle_x_cashu_streaming_response( + content_str=body, + response=_make_httpx_response(), + amount=10, + unit="sat", + max_cost_for_model=10_000, + request_id="no-space-usage", + ) + + assert cost.await_args is not None + assert cost.await_args.args[0]["usage"] is None + assert "".join(await _collect_streaming(response)).strip() == body.strip() + + +@pytest.mark.asyncio +async def test_streaming_no_space_usage_is_still_recorded_as_reported() -> None: + provider = _make_provider() + record = MagicMock() + body = 'data:{"usage":{"prompt_tokens":100,"completion_tokens":5}}\n\ndata:[DONE]\n\n' + + with ( + patch("routstr.upstream.base.record_terminal_outcome", record), + patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=None)), + ): + await provider.handle_x_cashu_streaming_response( + content_str=body, + response=_make_httpx_response(), + amount=10, + unit="sat", + max_cost_for_model=10_000, + request_id="no-space-recorded", + ) + + context = record.call_args.args[0] + assert (context.input_source, context.output_source) == ("reported", "reported") + assert record.call_args.kwargs["usage"] == { + "prompt_tokens": 100, + "completion_tokens": 5, + } + + +@pytest.mark.asyncio +async def test_native_messages_stream_keeps_input_usage_in_stats() -> None: + provider = _make_provider() + body = ( + 'event: message_start\ndata: {"type":"message_start","message":{"model":"test-model","usage":{"input_tokens":10,"output_tokens":0}}}\n\n' + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}\n\n' + 'event: message_stop\ndata: {"type":"message_stop"}\n\n' + ) + writer = MagicMock() + send_refund = AsyncMock(return_value="cashu-refund") + with ( + patch("routstr.core.terminal_outcomes.terminal_outcome_writer", writer), + patch.object(provider, "send_refund", send_refund), + patch( + "routstr.payment.cost_calculation._get_pricing_rates", + return_value=(1_000.0, 1_000.0, 1_000.0, 1_000.0, "configured"), + ), + patch("routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005), + ): + response = await provider.handle_x_cashu_streaming_response( + content_str=body, + response=_make_httpx_response(), + amount=100, + unit="msat", + max_cost_for_model=100, + request_id="messages-request", + ) + await _collect_streaming(response) + + send_refund.assert_awaited_once_with( + 95, "msat", None, request_id="messages-request" + ) + assert response.headers["x-cashu"] == "cashu-refund" + writer.submit.assert_called_once() + outcome = writer.submit.call_args.args[0] + assert outcome.revenue_msats == 5 + assert (outcome.input_tokens, outcome.output_tokens) == (10, 5) + assert outcome.input_source == outcome.output_source == "reported" + writer.declare_loss.assert_not_called() diff --git a/tests/unit/test_x_cashu_missing_usage.py b/tests/unit/test_x_cashu_missing_usage.py index b2510d00..ad0188cd 100644 --- a/tests/unit/test_x_cashu_missing_usage.py +++ b/tests/unit/test_x_cashu_missing_usage.py @@ -302,3 +302,26 @@ async def test_full_refund_is_logged_with_model_provider_and_body( assert logged["refund_amount"] == 10_000 assert logged["unit"] == "msat" assert "unpriced-model" in logged["response_body_preview"] + + +@pytest.mark.asyncio +async def test_pricing_failure_keeps_reported_usage_for_zero_charge_stats() -> None: + from routstr.upstream.base import BaseUpstreamProvider + + provider = BaseUpstreamProvider("https://unused.example/v1", "unused", 1.0) + with patch( + "routstr.payment.cost_calculation._get_pricing_rates", + side_effect=ValueError("No pricing for model"), + ): + cost = await provider.get_x_cashu_cost( + { + "model": "unpriced", + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + }, + 10_000, + None, + ) + assert cost is not None and cost.total_msats == 0 + assert (cost.input_tokens, cost.output_tokens) == (10, 5) + assert cost.input_source == cost.output_source == "reported" + assert cost.pricing_source == "missing" diff --git a/tests/unit/test_x_cashu_responses_streaming_sse.py b/tests/unit/test_x_cashu_responses_streaming_sse.py index 08154e40..0dcd7ec8 100644 --- a/tests/unit/test_x_cashu_responses_streaming_sse.py +++ b/tests/unit/test_x_cashu_responses_streaming_sse.py @@ -5,10 +5,11 @@ multi-line ``data:`` payloads and a ``[DONE]`` sentinel. Canonical Responses API usage arrives nested under ``response`` on ``response.completed``. """ +import asyncio import json import os from typing import Any -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest @@ -219,3 +220,35 @@ async def test_malformed_events_do_not_retain_whole_token() -> None: body = await _collect(response) assert b"\\n" not in body assert body.endswith(b"\n\n") + + +@pytest.mark.asyncio +async def test_cancelled_refund_marks_continuity_loss() -> None: + provider = _make_provider() + mark_loss = MagicMock() + record_outcome = MagicMock() + with ( + patch.object( + provider, + "get_x_cashu_cost", + new=AsyncMock(return_value=_make_cost_data(4000)), + ), + patch.object( + provider, + "send_refund", + new=AsyncMock(side_effect=asyncio.CancelledError()), + ), + patch("routstr.upstream.base.mark_terminal_outcome_loss", mark_loss), + patch("routstr.upstream.base.record_terminal_outcome", record_outcome), + ): + with pytest.raises(asyncio.CancelledError): + await provider.handle_x_cashu_responses_completion( + response=_sse_response(_canonical_chunks()), + amount=10_000, + unit="msat", + max_cost_for_model=9_000, + request_id="cancelled-refund", + ) + + mark_loss.assert_called_once_with("X-Cashu Responses stream settlement cancelled") + record_outcome.assert_not_called()