feat: record completed request usage after settlement

This commit is contained in:
Ashen
2026-10-01 11:21:53 +05:30
parent 5b16ea978e
commit 5c6871ad6e
22 changed files with 4185 additions and 128 deletions
@@ -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")
+67 -10
View File
@@ -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
+111 -1
View File
@@ -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)
+756
View File
@@ -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()
+150
View File
@@ -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
)
+78 -10
View File
@@ -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(),
)
+108
View File
@@ -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:
+565 -35
View File
File diff suppressed because it is too large Load Diff
+231 -22
View File
@@ -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=<name>`` 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]:
+8 -1
View File
@@ -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
+65 -2
View File
@@ -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"
+162
View File
@@ -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"
+26 -5
View File
@@ -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:
@@ -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()
+42 -24
View File
@@ -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
@@ -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,
)
+790
View File
@@ -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)
+37 -3
View File
@@ -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"
+72 -1
View File
@@ -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",
}
+165 -5
View File
@@ -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()
+23
View File
@@ -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"
@@ -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()