mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
feat: record completed request usage after settlement
This commit is contained in:
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
)
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+231
-22
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user