Merge branch 'main' into fix/streaming-billing-finalization

# Conflicts:
#	routstr/auth.py
#	routstr/proxy.py
#	routstr/upstream/base.py
#	tests/integration/test_balance_negative_on_cost_overrun.py
#	tests/integration/test_free_response_stale_reservation.py
This commit is contained in:
9qeklajc
2026-07-23 00:03:01 +02:00
18 changed files with 1324 additions and 222 deletions
+11 -7
View File
@@ -89,7 +89,9 @@ def create_model_mappings(
overrides_by_key: dict[tuple[str, int], tuple],
disabled_model_keys: set[tuple[str, int]],
) -> tuple[
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
dict[str, "Model"],
dict[str, list[tuple["Model", "BaseUpstreamProvider"]]],
dict[str, "Model"],
]:
"""Create optimal model mappings based on cost and provider preferences.
@@ -97,7 +99,9 @@ def create_model_mappings(
and creates three mappings based on cost optimization:
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias)
2. provider_map: alias -> List[(Model, UpstreamProvider)] (sorted candidate
list for each alias; each provider is paired with ITS OWN model so
failover can forward and bill the candidate that actually serves)
3. unique_models: base_id -> Model (unique models without provider prefixes)
The algorithm:
@@ -327,7 +331,7 @@ def create_model_mappings(
# Sort candidates and build final maps
model_instances: dict[str, "Model"] = {}
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
provider_map: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is.
@@ -374,13 +378,13 @@ def create_model_mappings(
best_model, best_provider = items[0]
model_instances[alias] = best_model
provider_map[alias] = [p for _, p in items]
provider_map[alias] = list(items)
# Log provider distribution (using top provider for stats)
provider_counts: dict[str, int] = {}
for providers in provider_map.values():
if providers:
provider = providers[0]
for candidate_list in provider_map.values():
if candidate_list:
provider = candidate_list[0][1]
provider_name = getattr(provider, "upstream_name", "unknown")
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
+58 -35
View File
@@ -7,7 +7,7 @@ import uuid
from contextvars import ContextVar
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
from typing import TYPE_CHECKING, Optional
from fastapi import HTTPException
from sqlalchemy import case, inspect
@@ -35,12 +35,17 @@ from .wallet import (
deserialize_token_from_string,
)
if TYPE_CHECKING:
from .payment.models import Model
logger = get_logger(__name__)
payments_logger = get_logger("routstr.payments")
# Routstr platform fee constants
ROUTSTR_FEE_PERCENT: float = 2.1
ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
ROUTSTR_LN_ADDRESS: str = (
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
)
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
@@ -63,6 +68,7 @@ def _clear_current_reservation(snapshot: ReservationSnapshot) -> None:
if current is not None and current.release_id == snapshot.release_id:
_current_reservation.set(None)
# TODO: implement prepaid api key (not like it was before)
# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None)
# PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats
@@ -691,11 +697,21 @@ async def pay_for_request(
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if child_result.rowcount == 0:
# Build the error before rollback expires ORM attributes.
limit_message = (
f"Balance limit exceeded: {key.balance_limit} mSats limit. "
f"{key.total_spent} already spent ({key.reserved_balance} reserved), "
f"{cost_per_request} required for this request."
)
# The parent reservation update already ran in this transaction.
# Roll it back before failover code attempts to restore the previous
# reservation; otherwise that later commit can persist both updates.
await session.rollback()
raise HTTPException(
status_code=402,
detail={
"error": {
"message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.",
"message": limit_message,
"type": "insufficient_quota",
"code": "balance_limit_exceeded",
}
@@ -860,13 +876,8 @@ async def _transition_reservation_to_released(
.where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(
col(ReservationRelease.billing_key_hash)
== snapshot.billing_key_hash
)
.where(
col(ReservationRelease.reserved_msats) == snapshot.reserved_msats
)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
.values(status="released")
)
transition_result = await session.exec(transition) # type: ignore[call-overload]
@@ -883,8 +894,7 @@ async def _transition_reservation_to_released(
)
values: dict[str, object] = {
"reserved_balance": col(ApiKey.reserved_balance)
- snapshot.reserved_msats,
"reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats,
"reserved_at": case(
(
col(ApiKey.reserved_balance) - snapshot.reserved_msats > 0,
@@ -951,13 +961,8 @@ async def _claim_reservation_for_charge(
.where(col(ReservationRelease.id) == snapshot.release_id)
.where(col(ReservationRelease.status) == "active")
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
.where(
col(ReservationRelease.billing_key_hash)
== snapshot.billing_key_hash
)
.where(
col(ReservationRelease.reserved_msats) == snapshot.reserved_msats
)
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
.values(status="charged")
)
result = await session.exec(statement) # type: ignore[call-overload]
@@ -974,6 +979,8 @@ async def adjust_payment_for_tokens(
response_data: dict,
session: AsyncSession,
deducted_max_cost: int,
model_obj: "Model | None" = None,
provider_fee: float | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> dict:
"""
@@ -981,6 +988,10 @@ async def adjust_payment_for_tokens(
This is called after the initial payment and the upstream request is complete.
Returns cost data to be included in the response.
``model_obj`` is the model that actually served the request; it is passed
through to ``calculate_cost`` so billing uses the serving candidate's
pricing instead of re-deriving it from the response's model string.
The response's usage object is normalized with the default union parser in
``calculate_cost``.
"""
@@ -1043,7 +1054,9 @@ async def adjust_payment_for_tokens(
extra={"error": str(e), "fee_msats": fee_msats},
)
calculated_cost = await calculate_cost(response_data, deducted_max_cost)
calculated_cost = await calculate_cost(
response_data, deducted_max_cost, model_obj, provider_fee
)
if not isinstance(calculated_cost, CostDataError):
if not await _claim_reservation_for_charge(reservation, session):
# A prior charge or release already owns this reservation. Returning
@@ -1079,8 +1092,10 @@ async def adjust_payment_for_tokens(
)
safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
@@ -1098,8 +1113,10 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
child_stmt = (
@@ -1209,8 +1226,10 @@ async def adjust_payment_for_tokens(
)
exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
@@ -1228,8 +1247,10 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
child_stmt = (
@@ -1384,8 +1405,10 @@ async def adjust_payment_for_tokens(
)
refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
@@ -1403,8 +1426,10 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
(
col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost,
),
else_=0,
)
child_stmt = (
@@ -1579,9 +1604,7 @@ async def periodic_dead_key_prune() -> None:
try:
async with create_session() as session:
await prune_dead_api_keys(
session, settings.dead_key_min_age_seconds
)
await prune_dead_api_keys(session, settings.dead_key_min_age_seconds)
except asyncio.CancelledError:
break
except Exception as e:
+35 -4
View File
@@ -1,4 +1,5 @@
import math
from typing import TYPE_CHECKING
from pydantic.v1 import BaseModel
@@ -7,6 +8,9 @@ from ..core.settings import settings
from .price import sats_usd_price
from .usage import normalize_usage, parse_token_count
if TYPE_CHECKING:
from .models import Model
__all__ = [
"CostData",
"CostDataError",
@@ -66,12 +70,23 @@ def _empty_cost(cls: type[CostData] = CostData) -> CostData:
async def calculate_cost(
response_data: dict,
max_cost: int,
model_obj: "Model | None" = None,
provider_fee: float | None = None,
) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage.
Args:
response_data: Response data containing usage information
max_cost: Maximum cost in millisats
model_obj: The model that actually served the request. When given,
its pricing is billed directly; without it, pricing is re-derived
from the response's model string via the alias map, which resolves
to the best-ranked candidate — not necessarily the serving one.
provider_fee: The serving provider's fee multiplier, applied on the
USD-cost path and the litellm pricing fallback (configured model
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.
Returns:
Cost data or error information
@@ -177,6 +192,7 @@ async def calculate_cost(
cache_creation_tokens,
output_tokens,
response_data,
provider_fee,
)
except Exception as e:
logger.warning(
@@ -190,7 +206,7 @@ async def calculate_cost(
# Fall back to token-based pricing
try:
pricing_rates = _get_pricing_rates(response_data)
pricing_rates = _get_pricing_rates(response_data, model_obj, provider_fee)
except ValueError as e:
return CostDataError(message=str(e), code="pricing_error")
@@ -307,9 +323,15 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
def _get_pricing_rates(
response_data: dict,
model_obj: "Model | None",
provider_fee: float | None,
) -> tuple[float, float, float, float] | 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).
``None`` means configured fixed pricing should be used by the caller.
"""
@@ -323,7 +345,13 @@ def _get_pricing_rates(
from .models import litellm_cost_entry
response_model = response_data.get("model", "")
model_obj = get_model_instance(response_model)
if model_obj is None:
logger.warning(
"Settling without routed model identity — re-deriving pricing "
"from the response's model string via the alias map",
extra={"response_model": response_model},
)
model_obj = get_model_instance(response_model)
if model_obj and model_obj.sats_pricing:
try:
@@ -360,7 +388,8 @@ def _get_pricing_rates(
if input_usd <= 0 or output_usd <= 0:
raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}")
provider_fee = _resolve_provider_fee(response_model)
if provider_fee is None:
provider_fee = _resolve_provider_fee(response_model)
usd_per_sat = sats_usd_price()
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
@@ -421,9 +450,11 @@ def _calculate_from_usd_cost(
cache_creation_tokens: int,
output_tokens: int,
response_data: dict,
provider_fee: float | None,
) -> CostData:
"""Calculate cost from USD figures, deriving input/output split from tokens."""
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
if provider_fee is None:
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
usd_cost = usd_cost * provider_fee
input_usd = input_usd * provider_fee
output_usd = output_usd * provider_fee
+78 -42
View File
@@ -43,10 +43,9 @@ logger = get_logger(__name__)
proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
_model_instances: dict[str, Model] = {} # All aliases -> Model
_provider_map: dict[
str, list[BaseUpstreamProvider]
] = {} # All aliases -> List[Provider]
str, list[tuple[Model, BaseUpstreamProvider]]
] = {} # All aliases -> sorted [(candidate Model, its Provider)]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
@@ -78,32 +77,44 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
return _upstreams
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
def get_candidates(
model_id: str,
) -> list[tuple[Model, BaseUpstreamProvider]] | None:
"""Get the sorted (model, provider) candidate list for a model ID.
Each provider is paired with its own model for the alias, so routing can
forward and bill the candidate that actually serves. Version suffixes
(e.g. ``-20251222``) are stripped as a retry when the exact ID is
unknown, since upstreams may return a specific version of a base model
we track.
"""
if not model_id:
return None
model_id_lower = model_id.lower()
# Try exact match first
if model := _model_instances.get(model_id_lower):
return model
if candidates := _provider_map.get(model_id_lower):
return candidates
# Try stripping common version suffixes (e.g., -20251222)
# This handles cases where upstream returns a specific version
# but we only track the base model name.
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower:
if model := _model_instances.get(base_model_id):
return model
if candidates := _provider_map.get(base_model_id):
return candidates
return None
def get_model_instance(model_id: str) -> Model | None:
"""Get the best-ranked Model instance for a model ID."""
candidates = get_candidates(model_id)
return candidates[0][0] if candidates else None
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
"""Get UpstreamProvider list for model ID from global cache."""
return _provider_map.get(model_id.lower())
"""Get the sorted UpstreamProvider list for a model ID."""
candidates = get_candidates(model_id)
return [provider for _, provider in candidates] if candidates else None
def get_unique_models() -> list[Model]:
@@ -143,7 +154,7 @@ async def refresh_model_maps() -> None:
"""Refresh global model and provider maps using the cost-based algorithm."""
from sqlalchemy.orm import selectinload
global _model_instances, _provider_map, _unique_models
global _provider_map, _unique_models
async with create_session() as session:
# Fetch all providers with their models in a single logical operation
@@ -166,7 +177,7 @@ async def refresh_model_maps() -> None:
else:
disabled_model_keys.add(model_key)
_model_instances, _provider_map, _unique_models = create_model_mappings(
_, _provider_map, _unique_models = create_model_mappings(
upstreams=_upstreams,
overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys,
@@ -289,25 +300,20 @@ async def proxy(
"upstream_error", "All upstreams failed", 502, request=request
)
model_obj = get_model_instance(model_id)
candidates = get_candidates(model_id)
if not model_obj:
if not candidates:
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
upstreams = get_provider_for_model(model_id)
if not upstreams:
return create_error_response(
"invalid_model",
f"No provider found for model '{model_id}'",
400,
request=request,
)
if is_ehbp:
upstreams = [upstream for upstream in upstreams if upstream.supports_ehbp]
if not upstreams:
candidates = [
(model, upstream)
for model, upstream in candidates
if upstream.supports_ehbp
]
if not candidates:
return create_error_response(
"unsupported_request",
f"No EHBP-capable provider found for model '{model_id}'",
@@ -315,9 +321,10 @@ async def proxy(
request=request,
)
# todo figure out cost calculation since fallback provider is usually not the same price
# Use first provider for initial checks/cost calculation
# primary_upstream = upstreams[0]
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried.
model_obj = candidates[0][0]
_max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
@@ -332,7 +339,7 @@ async def proxy(
if x_cashu := headers.get("x-cashu", None):
last_error = None
for i, upstream in enumerate(upstreams):
for i, (model_obj, upstream) in enumerate(candidates):
try:
if is_ehbp:
if not upstream.supports_ehbp:
@@ -370,7 +377,7 @@ async def proxy(
"status_code": e.status_code,
},
)
if i == len(upstreams) - 1:
if i == len(candidates) - 1:
last_error = e
continue
@@ -397,12 +404,12 @@ async def proxy(
logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None
for i, upstream in enumerate(upstreams):
for i, (_, upstream) in enumerate(candidates):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if response.status_code in [502, 429] and i < len(upstreams) - 1:
if response.status_code in [502, 429] and i < len(candidates) - 1:
error_message = ""
try:
if hasattr(response, "body"):
@@ -432,7 +439,7 @@ async def proxy(
return response
except UpstreamError as e:
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
if i == len(upstreams) - 1:
if i == len(candidates) - 1:
last_error_response = create_upstream_error_response(e, request)
continue
return last_error_response or create_error_response(
@@ -449,7 +456,35 @@ async def proxy(
# the reactive retry can never loop unboundedly.
already_stripped: set[str] = set()
for i, upstream in enumerate(upstreams):
for i, (model_obj, upstream) in enumerate(candidates):
if i > 0 and request_body_dict:
# The reservation was sized to the previous candidate's envelope;
# settlement bills the serving candidate, so a pricier fallback
# must be re-reserved at its own max cost before it is tried. A
# candidate whose envelope the key cannot cover is rejected, just
# as it would be had it been ranked first.
candidate_max = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
candidate_max = await calculate_discounted_max_cost(
candidate_max, request_body_dict, model_obj=model_obj
)
candidate_max = max(candidate_max, settings.min_request_msat)
if candidate_max > max_cost_for_model:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
try:
await pay_for_request(key, candidate_max, session)
except HTTPException:
if i == len(candidates) - 1:
raise
await pay_for_request(key, max_cost_for_model, session)
reservation_snapshot = await get_reservation_snapshot(key, session)
continue
reservation_snapshot = await get_reservation_snapshot(key, session)
max_cost_for_model = candidate_max
headers = upstream.prepare_headers(dict(request.headers))
try:
@@ -488,6 +523,7 @@ async def proxy(
max_cost_for_model,
session,
model_obj,
reservation_snapshot,
)
else:
response = await upstream.forward_request(
@@ -554,7 +590,7 @@ async def proxy(
if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
should_retry = response.status_code in [502, 429, 400, 401, 403, 404]
if should_retry and i < len(upstreams) - 1:
if should_retry and i < len(candidates) - 1:
error_message = ""
try:
if hasattr(response, "body"):
@@ -638,12 +674,12 @@ async def proxy(
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
"retry": i < len(upstreams) - 1,
"retry": i < len(candidates) - 1,
},
)
# If this was the last provider
if i == len(upstreams) - 1:
if i == len(candidates) - 1:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
+162 -86
View File
@@ -70,8 +70,7 @@ logger = get_logger(__name__)
def _is_json_content_type(content_type: str | None) -> bool:
"""Return True when the upstream response should be parsed as JSON.
"""
"""Return True when the upstream response should be parsed as JSON."""
if not content_type:
return False
main = content_type.split(";", 1)[0].strip().lower()
@@ -234,9 +233,7 @@ class BaseUpstreamProvider:
pass
if "prompt_tokens" in usage:
try:
usage["prompt_tokens"] = (
int(usage.get("prompt_tokens") or 0) + extra
)
usage["prompt_tokens"] = int(usage.get("prompt_tokens") or 0) + extra
except (TypeError, ValueError):
pass
@@ -738,7 +735,9 @@ class BaseUpstreamProvider:
and rate_limit.retry_after_seconds is not None
and "retry-after" not in {k.lower() for k in headers}
):
headers["Retry-After"] = str(max(1, math.ceil(rate_limit.retry_after_seconds)))
headers["Retry-After"] = str(
max(1, math.ceil(rate_limit.retry_after_seconds))
)
if is_json_body:
if not content_type:
@@ -804,6 +803,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
background_tasks: BackgroundTasks,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> StreamingResponse:
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
@@ -818,9 +818,7 @@ class BaseUpstreamProvider:
"""
if reservation_snapshot is None:
async with create_session() as snapshot_session:
snapshot_key = await snapshot_session.get(
key.__class__, key.hashed_key
)
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
if snapshot_key is None:
raise RuntimeError("Billing key disappeared before streaming")
reservation_snapshot = await get_reservation_snapshot(
@@ -858,6 +856,8 @@ class BaseUpstreamProvider:
{"model": last_model_seen or "unknown", "usage": None},
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
@@ -983,9 +983,7 @@ class BaseUpstreamProvider:
# line so multi-line ``data`` stays valid SSE framing - a bare
# second line would otherwise reach the client without its
# ``data:`` field and break naive parsers.
body = b"".join(
b"data: " + ln + b"\n" for ln in data.split(b"\n")
)
body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n"))
yield prefix + body + b"\n"
try:
@@ -1031,6 +1029,8 @@ class BaseUpstreamProvider:
adjustment_input,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
@@ -1076,24 +1076,18 @@ class BaseUpstreamProvider:
if usage_chunk_data is None:
if not hasattr(self, "_current_stream_id"):
self._current_stream_id = (
f"chatcmpl-{uuid.uuid4()}"
)
self._current_stream_id = f"chatcmpl-{uuid.uuid4()}"
usage_chunk_data = {
"id": self._current_stream_id,
"object": "chat.completion.chunk",
"model": last_model_seen or "unknown",
"choices": [],
"usage": {
"prompt_tokens": cost_data.get(
"input_tokens", 0
),
"prompt_tokens": cost_data.get("input_tokens", 0),
"completion_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
}
@@ -1148,6 +1142,7 @@ class BaseUpstreamProvider:
session: AsyncSession,
deducted_max_cost: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response:
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
@@ -1195,6 +1190,9 @@ class BaseUpstreamProvider:
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
await session.refresh(key)
@@ -1290,6 +1288,8 @@ class BaseUpstreamProvider:
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
@@ -1333,6 +1333,9 @@ class BaseUpstreamProvider:
{"model": last_model_seen or "unknown", "usage": None},
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except Exception:
@@ -1417,9 +1420,7 @@ class BaseUpstreamProvider:
return
# Re-prefix each line so multi-line ``data`` stays valid SSE
# framing for the client.
body = b"".join(
b"data: " + ln + b"\n" for ln in data.split(b"\n")
)
body = b"".join(b"data: " + ln + b"\n" for ln in data.split(b"\n"))
yield prefix + body + b"\n"
try:
@@ -1462,6 +1463,9 @@ class BaseUpstreamProvider:
adjustment_input,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
except Exception as e:
@@ -1494,22 +1498,14 @@ class BaseUpstreamProvider:
"output_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
},
"usage": {
"input_tokens": cost_data.get(
"input_tokens", 0
),
"output_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
"input_tokens": cost_data.get("input_tokens", 0),
"output_tokens": cost_data.get("output_tokens", 0),
"total_tokens": cost_data.get("input_tokens", 0)
+ cost_data.get("output_tokens", 0),
},
}
@@ -1525,9 +1521,9 @@ class BaseUpstreamProvider:
usage_chunk_data["response"]["usage"]["cost"] = (
cost_data.get("total_usd", 0.0)
)
usage_chunk_data["response"]["usage"][
"cost_sats"
] = sats_cost
usage_chunk_data["response"]["usage"]["cost_sats"] = (
sats_cost
)
usage_chunk_data["response"]["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
@@ -1580,6 +1576,8 @@ class BaseUpstreamProvider:
session: AsyncSession,
deducted_max_cost: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response:
"""Handle non-streaming Responses API responses with token usage tracking and cost adjustment.
@@ -1629,6 +1627,9 @@ class BaseUpstreamProvider:
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
await session.refresh(key)
@@ -1719,7 +1720,13 @@ class BaseUpstreamProvider:
raise
async def _finalize_generic_streaming_payment(
self, key_hash: str, max_cost: int, path: str
self,
key_hash: str,
max_cost: int,
path: str,
model_obj: Model | None,
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
) -> None:
"""Background task to finalize payment for generic streaming requests."""
async with create_session() as session:
@@ -1733,11 +1740,16 @@ class BaseUpstreamProvider:
try:
# Finalize with "unknown" model and no usage to release reservation/charge max cost
# (no routed identity here by design: the None usage settles at
# MaxCostData before any pricing lookup can happen).
await adjust_payment_for_tokens(
key,
{"model": "unknown", "usage": None},
session,
max_cost,
model_obj=model_obj,
provider_fee=provider_fee,
reservation_snapshot=reservation_snapshot,
)
logger.debug(
"Finalized generic streaming payment in background",
@@ -1762,6 +1774,8 @@ class BaseUpstreamProvider:
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> StreamingResponse:
async def stream_with_cost(
max_cost_for_model: int,
@@ -1804,9 +1818,7 @@ class BaseUpstreamProvider:
_coerce_usd(cd.get("output_cost")),
)
for field in ("total_cost", "cost"):
total_cost = max(
total_cost, _coerce_usd(usage_or_root.get(field))
)
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
@@ -1827,6 +1839,9 @@ class BaseUpstreamProvider:
fallback,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
@@ -1850,9 +1865,7 @@ class BaseUpstreamProvider:
if msg and msg.get("model"):
last_model_seen = str(msg.get("model"))
provider_added = (
"provider" not in data
)
provider_added = "provider" not in data
self._apply_provider_field(data)
if requested_model:
@@ -1978,6 +1991,9 @@ class BaseUpstreamProvider:
combined_data,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(
@@ -2025,6 +2041,8 @@ class BaseUpstreamProvider:
deducted_max_cost: int,
path: str,
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response:
try:
content = await response.aread()
@@ -2049,6 +2067,9 @@ class BaseUpstreamProvider:
response_json,
session,
deducted_max_cost,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -2096,9 +2117,7 @@ class BaseUpstreamProvider:
async def _aggregate_anthropic_events_to_message(
self, iterator: AsyncIterator[Any]
) -> dict:
return await messages_dispatch.aggregate_anthropic_events_to_message(
iterator
)
return await messages_dispatch.aggregate_anthropic_events_to_message(iterator)
async def _dispatch_anthropic_messages(
self,
@@ -2124,6 +2143,7 @@ class BaseUpstreamProvider:
session: AsyncSession,
max_cost_for_model: int,
model_obj: Model,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response | StreamingResponse:
"""Translate /v1/messages to upstream chat/completions via litellm.
@@ -2143,6 +2163,8 @@ class BaseUpstreamProvider:
key,
max_cost_for_model,
requested_model,
model_obj,
reservation_snapshot,
)
response_json = messages_dispatch.coerce_litellm_payload(result)
@@ -2154,6 +2176,9 @@ class BaseUpstreamProvider:
response_json,
session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(response_json, cost_data, key)
@@ -2193,6 +2218,7 @@ class BaseUpstreamProvider:
requested_model,
mint,
request_id,
model_obj,
)
response_json = messages_dispatch.coerce_litellm_payload(result)
@@ -2200,10 +2226,14 @@ class BaseUpstreamProvider:
if requested_model and "model" in response_json:
response_json["model"] = requested_model
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json and isinstance(
response_json["usage"], dict
if (
cost_data
and "usage" in response_json
and isinstance(response_json["usage"], dict)
):
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
self._fold_cache_into_input_tokens(response_json["usage"])
@@ -2245,6 +2275,8 @@ class BaseUpstreamProvider:
key: ApiKey,
max_cost_for_model: int,
requested_model: str | None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
) -> StreamingResponse:
"""Re-emit a litellm Anthropic-event iterator as live SSE bytes
with cost reconciliation appended at end of stream."""
@@ -2279,9 +2311,7 @@ class BaseUpstreamProvider:
},
)
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
usage_finalized = True
return None
@@ -2295,11 +2325,13 @@ class BaseUpstreamProvider:
fallback,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
usage_finalized = True
return (
f"event: cost\ndata: "
f"{json.dumps({'cost': cost_data})}\n\n"
f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n"
).encode()
except Exception:
usage_finalized = True
@@ -2338,9 +2370,7 @@ class BaseUpstreamProvider:
or total_cost > 0
):
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
rebuilt_usage: dict = {
@@ -2368,6 +2398,9 @@ class BaseUpstreamProvider:
combined_data,
new_session,
max_cost_for_model,
model_obj,
self.provider_fee,
reservation_snapshot,
)
self.inject_cost_metadata(
combined_data, cost_data, fresh_key
@@ -2405,6 +2438,7 @@ class BaseUpstreamProvider:
requested_model: str | None,
mint: str | None,
request_id: str | None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Buffer a litellm stream end-to-end, compute cost, then replay.
@@ -2496,7 +2530,7 @@ class BaseUpstreamProvider:
}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
refund_amount = messages_dispatch.compute_refund(
@@ -2511,8 +2545,7 @@ class BaseUpstreamProvider:
)
response_headers["X-Cashu"] = refund_token
logger.info(
"Refund processed for streaming /v1/messages "
"via litellm",
"Refund processed for streaming /v1/messages via litellm",
extra={
"refund_amount": refund_amount,
"unit": unit,
@@ -2585,6 +2618,7 @@ class BaseUpstreamProvider:
session=session,
max_cost_for_model=max_cost_for_model,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
url = self.build_request_url(path, model_obj)
@@ -2714,6 +2748,8 @@ class BaseUpstreamProvider:
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -2730,6 +2766,8 @@ class BaseUpstreamProvider:
max_cost_for_model,
path,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
finally:
await response.aclose()
@@ -2745,6 +2783,8 @@ class BaseUpstreamProvider:
max_cost_for_model,
path,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
finally:
await response.aclose()
@@ -2794,6 +2834,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
background_tasks,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
result.background = background_tasks
@@ -2808,11 +2849,16 @@ class BaseUpstreamProvider:
session,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
finally:
await response.aclose()
await client.aclose()
if reservation_snapshot is None:
reservation_snapshot = await get_reservation_snapshot(key, session)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
@@ -2821,6 +2867,9 @@ class BaseUpstreamProvider:
key.hashed_key,
max_cost_for_model,
path,
model_obj,
self.provider_fee,
reservation_snapshot,
)
logger.debug(
@@ -2895,7 +2944,9 @@ class BaseUpstreamProvider:
supports_ehbp: bool = False
def get_confidential_inference_profile(self) -> "ConfidentialInferenceProfile | None":
def get_confidential_inference_profile(
self,
) -> "ConfidentialInferenceProfile | None":
"""Return provider policy for encrypted/confidential inference forwarding."""
return None
@@ -2923,6 +2974,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
session: AsyncSession,
model_obj: Model,
reservation_snapshot: ReservationSnapshot | None = None,
) -> Response | StreamingResponse:
"""Forward authenticated Responses API request to upstream service with cost tracking.
@@ -3060,6 +3112,8 @@ class BaseUpstreamProvider:
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -3075,11 +3129,16 @@ class BaseUpstreamProvider:
session,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
finally:
await response.aclose()
await client.aclose()
if reservation_snapshot is None:
reservation_snapshot = await get_reservation_snapshot(key, session)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
@@ -3088,6 +3147,9 @@ class BaseUpstreamProvider:
key.hashed_key,
max_cost_for_model,
path,
model_obj,
self.provider_fee,
reservation_snapshot,
)
logger.debug(
@@ -3251,13 +3313,19 @@ class BaseUpstreamProvider:
)
async def get_x_cashu_cost(
self, response_data: dict, max_cost_for_model: int
self,
response_data: dict,
max_cost_for_model: int,
model_obj: Model | None,
) -> MaxCostData | CostData | None:
"""Calculate cost for X-Cashu payment based on response data.
Args:
response_data: Response data containing model and usage information
max_cost_for_model: Maximum cost for the model
model_obj: The model that actually served the request; billed
directly instead of re-deriving pricing from the upstream's
echoed model string
Returns:
Cost data object (MaxCostData or CostData) or None if calculation fails
@@ -3271,6 +3339,8 @@ class BaseUpstreamProvider:
match await calculate_cost(
response_data,
max_cost_for_model,
model_obj,
self.provider_fee,
):
case MaxCostData() as cost:
logger.debug(
@@ -3415,6 +3485,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Handle streaming response for X-Cashu payment, calculating refund if needed.
@@ -3488,7 +3559,7 @@ class BaseUpstreamProvider:
response_data = {"usage": usage_data, "model": model}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
@@ -3559,14 +3630,8 @@ class BaseUpstreamProvider:
if "provider" not in data_json:
self._apply_provider_field(data_json)
changed = True
if (
cost_data
and "usage" in data_json
and data_json["usage"]
):
data_json["usage"]["cost_sats"] = (
cost_data.total_msats // 1000
)
if cost_data and "usage" in data_json and data_json["usage"]:
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
changed = True
if changed:
lines[i] = "data: " + json.dumps(data_json)
@@ -3593,6 +3658,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> Response:
"""Handle non-streaming response for X-Cashu payment, calculating refund if needed.
@@ -3614,7 +3680,9 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
@@ -3743,6 +3811,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse | Response:
"""Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming.
@@ -3786,6 +3855,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
else:
return await self.handle_x_cashu_non_streaming_response(
@@ -3796,6 +3866,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
except Exception as e:
@@ -3981,6 +4052,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4273,6 +4345,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=getattr(request.state, "request_id", None),
model_obj=model_obj,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -4323,6 +4396,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse | Response:
"""Handle Responses API completion response for X-Cashu payment.
@@ -4367,6 +4441,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
else:
return await self.handle_x_cashu_non_streaming_responses_response(
@@ -4377,6 +4452,7 @@ class BaseUpstreamProvider:
max_cost_for_model,
mint,
request_id=request_id,
model_obj=model_obj,
)
except Exception as e:
@@ -4404,6 +4480,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API response for X-Cashu payment.
@@ -4462,7 +4539,7 @@ class BaseUpstreamProvider:
response_data = {"usage": usage_data, "model": model}
try:
cost_data = await self.get_x_cashu_cost(
response_data, max_cost_for_model
response_data, max_cost_for_model, model_obj
)
if cost_data:
if unit == "msat":
@@ -4534,14 +4611,8 @@ class BaseUpstreamProvider:
if "provider" not in data_json:
self._apply_provider_field(data_json)
changed = True
if (
cost_data
and "usage" in data_json
and data_json["usage"]
):
data_json["usage"]["cost_sats"] = (
cost_data.total_msats // 1000
)
if cost_data and "usage" in data_json and data_json["usage"]:
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
changed = True
if changed:
lines[i] = "data: " + json.dumps(data_json)
@@ -4568,6 +4639,7 @@ class BaseUpstreamProvider:
max_cost_for_model: int,
mint: str | None = None,
request_id: str | None = None,
model_obj: Model | None = None,
) -> Response:
"""Handle non-streaming Responses API response for X-Cashu payment."""
logger.debug(
@@ -4578,7 +4650,9 @@ class BaseUpstreamProvider:
try:
response_json = json.loads(content_str)
self._apply_provider_field(response_json)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
cost_data = await self.get_x_cashu_cost(
response_json, max_cost_for_model, model_obj
)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
@@ -5030,7 +5104,9 @@ class BaseUpstreamProvider:
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
self._models_by_id = {
m.forwarded_model_id or m.id: m for m in self._models_cache
}
except Exception as e:
logger.error(
@@ -85,7 +85,9 @@ async def test_balance_never_negative_when_cost_exceeds_reservation(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await _refresh(integration_session, key)
@@ -121,7 +123,9 @@ async def test_balance_floor_at_zero_on_overrun(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await _refresh(integration_session, key)
@@ -162,7 +166,9 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await _refresh(integration_session, key)
@@ -236,19 +242,22 @@ async def test_concurrent_cost_overruns_never_negative(
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation,
)
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=reservation,
)
await asyncio.gather(*[finalize(r) for r in reservations])
# Patch once around the gather: entering the same patch target from
# concurrent tasks un-patches in the wrong order and leaks the mock into
# every later test in the session.
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await asyncio.gather(*(finalize(r) for r in reservations))
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
@@ -298,7 +307,9 @@ async def test_zero_free_balance_overrun_is_safe(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost, None, None
)
await _refresh(integration_session, key)
@@ -371,19 +382,22 @@ async def test_parallel_requests_no_free_inference(
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation,
)
await adjust_payment_for_tokens(
fresh_key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=reservation,
)
await asyncio.gather(*(finalize(r) for r in reservations))
# Patch once around the gather: entering the same patch target from two
# concurrent tasks un-patches in the wrong order and leaks the mock into
# every later test in the session.
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await asyncio.gather(*(finalize(r) for r in reservations))
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
+1 -1
View File
@@ -77,7 +77,7 @@ async def test_child_key_flow(integration_session: AsyncSession) -> None:
try:
adjustment = await adjust_payment_for_tokens(
child_key_db, response_data, integration_session, 500
child_key_db, response_data, integration_session, 500, None, None
)
assert adjustment["total_msats"] == 400
+666
View File
@@ -0,0 +1,666 @@
"""Failover requests are billed and forwarded as the provider that served them.
Covers the whole-system settlement path when two enabled providers expose the
same model under different spellings and prices: the routing winner fails with
a 502, the fallback provider serves, and the response must be billed at the
fallback's configured rate, carry the fallback's model id in the forwarded
request body, and echo the fallback's model id to the client.
"""
import json
from typing import Any, AsyncGenerator
from unittest.mock import patch
import httpx
import pytest
from httpx import AsyncClient
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, ReservationRelease
from routstr.payment.models import Architecture, Model, Pricing
from routstr.proxy import refresh_model_maps
from routstr.upstream.base import BaseUpstreamProvider
CHEAP_BASE_URL = "https://cheap.example.com/v1"
EXPENSIVE_BASE_URL = "https://expensive.example.com/v1"
THIRD_BASE_URL = "https://third.example.com/v1"
def _make_model(
model_id: str,
prompt_sats: float,
completion_sats: float,
max_cost: float = 50.0,
) -> Model:
"""Build a model whose USD and sats pricing rank consistently."""
return Model(
id=model_id,
name=model_id,
created=1,
description="test model",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="gpt",
instruct_type=None,
),
pricing=Pricing(
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
),
sats_pricing=Pricing(
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
),
)
class _StaticProvider(BaseUpstreamProvider):
"""Upstream provider with a fixed model catalog and no remote refresh."""
def __init__(self, base_url: str, api_key: str, fee: float, model: Model) -> None:
super().__init__(base_url, api_key, fee)
self.provider_type = "custom"
self._static_model = model
def get_cached_models(self) -> list[Model]:
return [self._static_model]
async def refresh_models_cache(self) -> None:
pass
async def _install_providers(
providers: list[_StaticProvider],
) -> AsyncGenerator[None, None]:
"""Install providers into the routing maps, restoring the originals after."""
from routstr import proxy
original_upstreams = proxy.get_upstreams()
with patch("routstr.proxy._upstreams", providers):
await refresh_model_maps()
yield
with patch("routstr.proxy._upstreams", original_upstreams):
await refresh_model_maps()
@pytest.fixture
async def dual_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[tuple[_StaticProvider, _StaticProvider], None]:
"""Two same-tail providers under different spellings and prices."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("prova/dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("provb/dual-model", 0.005, 0.010, max_cost=100.0),
)
async for _ in _install_providers([cheap, expensive]):
yield cheap, expensive
def _upstream_response(request: httpx.Request) -> httpx.Response:
"""502 from the cheap (winning) provider; a served completion elsewhere."""
if request.url.host == "cheap.example.com":
return httpx.Response(
502,
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
headers={"content-type": "application/json"},
)
body = {
"id": "chatcmpl-served",
"object": "chat.completion",
"created": 1,
"model": "dual-model",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 500,
"total_tokens": 1500,
},
}
return httpx.Response(
200,
content=json.dumps(body).encode(),
headers={"content-type": "application/json"},
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_serve_billed_at_serving_providers_rate(
authenticated_client: AsyncClient,
dual_provider_maps: tuple[_StaticProvider, _StaticProvider],
integration_session: AsyncSession,
) -> None:
"""A fallback serve is billed at the fallback's price, not the winner's.
The cheap provider ranks first for the shared tail; it 502s and the
expensive provider serves 1000 input + 500 output tokens. At the serving
provider's sats pricing (0.005/0.010 sats per token) that is 10_000 msats;
at the winner's (0.001/0.002) it would be 2_000 msats.
"""
sent_requests: list[httpx.Request] = []
# 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.
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
# cost_calculation binds sats_usd_price at import time, so the price
# patch in the app fixture does not reach it; patch its own binding.
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
payload = response.json()
# Both providers were attempted, cheapest first.
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
# The fallback must be asked for ITS OWN model spelling, not the winner's.
forwarded_body = json.loads(sent_requests[1].content)
assert forwarded_body["model"] == "provb/dual-model"
# The response echo names the model that actually served.
assert payload["model"] == "provb/dual-model"
# Billed at the serving provider's rate: 1000/1000*5000 + 500/1000*10000.
assert payload["cost"]["total_msats"] == 10_000
# The fallback's larger max-cost envelope requires a replacement
# reservation. The failed candidate is released, the serving candidate is
# charged, and no request-owned reservation remains active.
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert sorted(record.status for record in records) == ["charged", "released"]
@pytest.fixture
async def same_id_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Two providers exposing the IDENTICAL model id at different prices."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_same_id_failover_settles_at_serving_price(
authenticated_client: AsyncClient,
same_id_provider_maps: None,
) -> None:
"""Settlement must not re-derive pricing from the response's model string.
Both providers expose the exact same model id, so the forwarded body is
identical either way the only observable difference is the settled
amount. The response's model string resolves to the alias winner (cheap),
but the expensive provider served, so the bill must be 10_000 msats, not
the winner's 2_000.
"""
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 10_000
@pytest.mark.integration
@pytest.mark.asyncio
async def test_version_suffixed_model_id_routes(
authenticated_client: AsyncClient,
same_id_provider_maps: None,
) -> None:
"""A version-suffixed request (``…-YYYYMMDD``) routes to the base model.
Model resolution stripped the suffix but the provider lookup did not, so
such requests resolved a model yet found no provider and 400'd. With the
unified candidate lookup the strip applies to both.
"""
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model-20260101",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
@pytest.fixture
async def fee_split_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Same-tail providers whose fees differ; the serving one charges 1.5x."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.5,
_make_model("dual-model", 0.005, 0.010),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_usd_cost_serve_carries_serving_providers_fee(
authenticated_client: AsyncClient,
fee_split_provider_maps: None,
) -> None:
"""The USD-cost billing path applies the SERVING provider's fee.
The upstream that serves reports ``usage.cost`` in USD, so billing goes
through the USD-cost path where the provider fee is applied explicitly.
The serving provider's fee is 1.5; the alias winner's is 1.0. At 0.001 USD
reported cost and 0.0005 USD/sat: 0.001 * 1.5 / 0.0005 = 3 sats = 3000
msats (fee 1.0 would give 2000).
"""
sent_requests: list[httpx.Request] = []
def usd_cost_response(request: httpx.Request) -> httpx.Response:
if request.url.host == "cheap.example.com":
return httpx.Response(
502,
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
headers={"content-type": "application/json"},
)
body = {
"id": "chatcmpl-usd",
"object": "chat.completion",
"created": 1,
"model": "dual-model",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.001,
},
}
return httpx.Response(
200,
content=json.dumps(body).encode(),
headers={"content-type": "application/json"},
)
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return usd_cost_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 3_000
@pytest.fixture
async def envelope_split_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Same-id providers where the fallback's max cost dwarfs the key balance."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010, max_cost=20_000.0),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_beyond_balance_envelope_is_rejected(
authenticated_client: AsyncClient,
envelope_split_provider_maps: None,
) -> None:
"""A fallback whose max-cost envelope exceeds the balance is not served.
Admission and reservation are sized to the best-ranked candidate's max
cost. When that candidate fails and the next one's envelope exceeds the
key's balance, serving it could settle far beyond what admission allowed,
so the request must be rejected (as it would be if the pricier candidate
were ranked first) instead of forwarded.
"""
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
# The 20_000-sat envelope exceeds the key's 10_000-sat balance: the
# fallback must be rejected before its upstream is ever contacted.
assert response.status_code == 402
assert [r.url.host for r in sent_requests] == ["cheap.example.com"]
@pytest.fixture
async def three_candidate_child_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Second candidate cannot fit the child limit; third restores and serves."""
first = _StaticProvider(
CHEAP_BASE_URL,
"key-first",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
too_large = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-too-large",
1.0,
_make_model("dual-model", 0.002, 0.003, max_cost=100.0),
)
third = _StaticProvider(
THIRD_BASE_URL,
"key-third",
1.0,
_make_model("dual-model", 0.003, 0.004, max_cost=50.0),
)
async for _ in _install_providers([first, too_large, third]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring(
authenticated_client: AsyncClient,
three_candidate_child_maps: None,
integration_session: AsyncSession,
) -> None:
"""A failed child guard cannot leak its parent update into restoration."""
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
child = await integration_session.get(ApiKey, key_hash)
assert child is not None
parent = ApiKey(hashed_key="failover-parent", balance=10_000_000)
child.parent_key_hash = parent.hashed_key
child.balance_limit = 75_000
integration_session.add(parent)
integration_session.add(child)
await integration_session.commit()
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
# The 100-sat candidate is rejected before forwarding; the third serves.
assert [request.url.host for request in sent_requests] == [
"cheap.example.com",
"third.example.com",
]
await integration_session.refresh(parent)
await integration_session.refresh(child)
assert parent.reserved_balance == 0
assert child.reserved_balance == 0
assert parent.total_spent == response.json()["cost"]["total_msats"]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert len(records) == 2
assert sorted(record.status for record in records) == ["charged", "released"]
assert len({record.reserved_msats for record in records}) == 1
assert all(record.status != "active" for record in records)
@pytest.fixture
async def raised_envelope_provider_maps(
patched_db_engine: None,
) -> AsyncGenerator[None, None]:
"""Same-id providers where the fallback needs a larger, affordable reserve."""
cheap = _StaticProvider(
CHEAP_BASE_URL,
"key-cheap",
1.0,
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
)
expensive = _StaticProvider(
EXPENSIVE_BASE_URL,
"key-expensive",
1.0,
_make_model("dual-model", 0.005, 0.010, max_cost=100.0),
)
async for _ in _install_providers([cheap, expensive]):
yield
@pytest.mark.integration
@pytest.mark.asyncio
async def test_failover_reserves_serving_candidates_envelope(
authenticated_client: AsyncClient,
raised_envelope_provider_maps: None,
integration_session: AsyncSession,
) -> None:
"""An affordable pricier fallback is re-reserved, served, and billed.
The fallback's max cost (100 sats) exceeds the winner's (50 sats) but fits
the key's balance, so the reservation is raised to the serving candidate's
envelope and the request completes, billed at the serving rate with the
unused reserve refunded.
"""
sent_requests: list[httpx.Request] = []
async def fake_transport(
request: httpx.Request, *args: Any, **kwargs: Any
) -> httpx.Response:
sent_requests.append(request)
return _upstream_response(request)
with (
patch(
"httpx.AsyncHTTPTransport.handle_async_request",
side_effect=fake_transport,
),
patch(
"routstr.payment.cost_calculation.sats_usd_price",
return_value=0.0005,
),
):
response = await authenticated_client.post(
"/v1/chat/completions",
json={
"model": "dual-model",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 200
assert [r.url.host for r in sent_requests] == [
"cheap.example.com",
"expensive.example.com",
]
assert response.json()["cost"]["total_msats"] == 10_000
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
records = (
await integration_session.exec(
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
)
).all()
assert len(records) == 2
released = next(record for record in records if record.status == "released")
charged = next(record for record in records if record.status == "charged")
assert charged.reserved_msats > released.reserved_msats
assert all(record.status != "active" for record in records)
@@ -62,7 +62,7 @@ async def test_overrun_charges_after_reservation_swept(
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, integration_session, deducted_max_cost
key, response_data, integration_session, deducted_max_cost, None, None
)
await integration_session.refresh(key)
@@ -146,7 +146,11 @@ async def test_free_response_path_closed_end_to_end(
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
key, response_data, session, deducted_max_cost, snapshot
key,
response_data,
session,
deducted_max_cost,
reservation_snapshot=snapshot,
)
async with create_session() as session:
@@ -241,8 +241,10 @@ async def test_http_402_response_shape_on_insufficient_balance(
mock_upstream.prepare_headers = MagicMock(return_value={})
with (
patch("routstr.proxy.get_model_instance", return_value=mock_model),
patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]),
patch(
"routstr.proxy.get_candidates",
return_value=[(mock_model, mock_upstream)],
),
# Patch where it is used (proxy imports it at module level)
patch(
"routstr.proxy.get_max_cost_for_model",
@@ -120,7 +120,9 @@ async def test_finalise_releases_reservation_and_charges_balance(
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
with patch("routstr.auth.calculate_cost", return_value=cost_data):
await adjust_payment_for_tokens(key, response_data, integration_session, cost)
await adjust_payment_for_tokens(
key, response_data, integration_session, cost, None, None
)
await integration_session.refresh(key)
+5 -5
View File
@@ -142,7 +142,7 @@ def test_create_model_mappings_includes_db_override_for_missing_cached_model(
)
assert "azure/gpt-4o" in model_instances
assert provider_map["azure/gpt-4o"] == [provider]
assert [p for _, p in provider_map["azure/gpt-4o"]] == [provider]
assert "gpt-4o" in unique_models
@@ -186,7 +186,7 @@ def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
disabled_model_keys=set(),
)
providers_for_alias = provider_map["azure/gpt-4o"]
providers_for_alias = [p for _, p in provider_map["azure/gpt-4o"]]
assert provider_a in providers_for_alias
assert provider_b in providers_for_alias
assert len(providers_for_alias) == 2
@@ -226,8 +226,8 @@ def test_create_model_mappings_applies_override_only_to_matching_provider(
disabled_model_keys=set(),
)
assert provider_map["provider-b-only"] == [provider_b]
assert set(provider_map["same-id"]) == {provider_a, provider_b}
assert [p for _, p in provider_map["provider-b-only"]] == [provider_b]
assert {p for _, p in provider_map["same-id"]} == {provider_a, provider_b}
def test_create_model_mappings_disables_only_matching_provider() -> None:
@@ -251,4 +251,4 @@ def test_create_model_mappings_disables_only_matching_provider() -> None:
disabled_model_keys={("same-id", 2)},
)
assert provider_map["same-id"] == [provider_a]
assert [p for _, p in provider_map["same-id"]] == [provider_a]
+2 -2
View File
@@ -146,7 +146,7 @@ def test_get_x_cashu_cost_with_usage() -> None:
"usage": {"prompt_tokens": 100, "completion_tokens": 50},
}
result = p.get_x_cashu_cost(response_data, 100000)
result = p.get_x_cashu_cost(response_data, 100000, None)
# Either returns None (needs more data) or a cost object
assert result is not None
@@ -157,7 +157,7 @@ def test_get_x_cashu_cost_no_usage() -> None:
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
response_data = {"model": "gpt-4"}
result = p.get_x_cashu_cost(response_data, 100000)
result = p.get_x_cashu_cost(response_data, 100000, None)
# Without usage, uses max_cost
assert result is not None
+33 -2
View File
@@ -18,6 +18,7 @@ from fastapi.responses import Response, StreamingResponse
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.auth import ReservationSnapshot # noqa: E402
from routstr.core.db import ApiKey # noqa: E402
from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
@@ -498,14 +499,27 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
yield {"type": "message_stop"}
fake_cost = {"total_msats": 4321, "total_usd": 0.00015}
reservation = ReservationSnapshot(
release_id="messages-stream",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=10_000,
)
captured_cost_call: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
fresh_key: Any,
combined_data: Any,
sess: Any,
max_cost: int,
model_obj: Any = None,
provider_fee: Any = None,
reservation_snapshot: Any = None,
) -> dict:
captured_cost_call["combined_data"] = combined_data
captured_cost_call["max_cost"] = max_cost
captured_cost_call["reservation_snapshot"] = reservation_snapshot
return fake_cost
fake_session = MagicMock()
@@ -539,6 +553,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
session=session,
max_cost_for_model=10_000,
model_obj=model,
reservation_snapshot=reservation,
)
assert isinstance(result, StreamingResponse)
@@ -561,6 +576,7 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
assert combined["usage"]["input_tokens"] == 5
assert combined["usage"]["output_tokens"] == 7
assert combined["model"] == "openai/gpt-4o-mini"
assert captured_cost_call["reservation_snapshot"] is reservation
@pytest.mark.asyncio
@@ -588,12 +604,25 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
yield b'event: message_stop\ndata: {"type":"message_stop"}\n\n'
fake_cost = {"total_msats": 999, "total_usd": 0.0001}
reservation = ReservationSnapshot(
release_id="messages-byte-stream",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=10_000,
)
captured: dict[str, Any] = {}
async def fake_adjust(
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
fresh_key: Any,
combined_data: Any,
sess: Any,
max_cost: int,
model_obj: Any = None,
provider_fee: Any = None,
reservation_snapshot: Any = None,
) -> dict:
captured["combined_data"] = combined_data
captured["reservation_snapshot"] = reservation_snapshot
return fake_cost
fake_session = MagicMock()
@@ -626,6 +655,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
session=session,
max_cost_for_model=10_000,
model_obj=model,
reservation_snapshot=reservation,
)
assert isinstance(result, StreamingResponse)
@@ -649,6 +679,7 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
assert combined["usage"]["input_tokens"] == 3
assert combined["usage"]["output_tokens"] == 4
assert combined["model"] == "openai/gpt-4o-mini"
assert captured["reservation_snapshot"] is reservation
# ---------------------------------------------------------------------------
+152
View File
@@ -0,0 +1,152 @@
"""Settlement bills the model and provider fee that actually served.
Covers ``calculate_cost``'s served-identity parameters: a passed ``model_obj``
is billed directly instead of re-deriving pricing from the response's model
string through the alias map (which yields the best-ranked candidate, not the
serving one), and a passed ``provider_fee`` is applied on the USD-cost path
instead of the best-ranked provider's fee. The string/alias fallbacks remain
for callers without routed identity.
"""
import os
from unittest.mock import patch
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.payment.cost_calculation import CostData, calculate_cost
from routstr.payment.models import Architecture, Model, Pricing
def _make_model(
model_id: str, prompt_sats: float, completion_sats: float
) -> Model:
return Model(
id=model_id,
name=model_id,
created=0,
description="",
context_length=64000,
architecture=Architecture(
modality="text->text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="Other",
instruct_type=None,
),
pricing=Pricing(prompt=prompt_sats, completion=completion_sats),
sats_pricing=Pricing(prompt=prompt_sats, completion=completion_sats),
)
WINNER = _make_model("dual-model", 0.001, 0.002)
SERVED = _make_model("dual-model", 0.005, 0.010)
RESPONSE = {
"model": "dual-model",
"usage": {
"prompt_tokens": 1000,
"completion_tokens": 500,
"total_tokens": 1500,
},
}
@pytest.fixture(autouse=True)
def patch_sats_usd_price() -> None: # type: ignore[misc]
with patch(
"routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4
):
yield
@pytest.mark.asyncio
async def test_served_model_pricing_wins_over_alias_lookup() -> None:
"""With ``model_obj`` given, the alias map is not consulted for pricing."""
with patch(
"routstr.proxy.get_model_instance", return_value=WINNER
) as alias_lookup:
result = await calculate_cost(
dict(RESPONSE), max_cost=100_000, model_obj=SERVED
)
assert isinstance(result, CostData)
# 1000/1000 * 5000 + 500/1000 * 10000 msats at the SERVED model's rates.
assert result.total_msats == 10_000
alias_lookup.assert_not_called()
@pytest.mark.asyncio
async def test_string_fallback_still_prices_without_model_obj() -> None:
"""Callers without routed identity keep the alias-map string lookup."""
with patch("routstr.proxy.get_model_instance", return_value=WINNER):
result = await calculate_cost(dict(RESPONSE), max_cost=100_000)
assert isinstance(result, CostData)
assert result.total_msats == 2_000
@pytest.mark.asyncio
async def test_usd_cost_path_applies_given_provider_fee() -> None:
"""The USD-cost path bills the serving provider's fee when supplied."""
from unittest.mock import Mock
response = dict(RESPONSE)
response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type]
best_ranked = Mock(provider_fee=1.0)
with patch(
"routstr.proxy.get_provider_for_model", return_value=[best_ranked]
):
result = await calculate_cost(
response, max_cost=100_000, model_obj=SERVED, provider_fee=1.5
)
assert isinstance(result, CostData)
# 0.001 USD * fee 1.5 / 0.0005 USD-per-sat = 3 sats = 3000 msats.
assert result.total_msats == 3_000
@pytest.mark.asyncio
async def test_usd_cost_path_falls_back_to_best_ranked_fee() -> None:
"""Without a supplied fee, the alias-map provider lookup still applies."""
from unittest.mock import Mock
response = dict(RESPONSE)
response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type]
best_ranked = Mock(provider_fee=2.0)
with patch(
"routstr.proxy.get_provider_for_model", return_value=[best_ranked]
):
result = await calculate_cost(response, max_cost=100_000)
assert isinstance(result, CostData)
assert result.total_msats == 4_000
@pytest.mark.asyncio
async def test_x_cashu_cost_uses_served_model_not_upstream_echo() -> None:
"""``get_x_cashu_cost`` bills the routed model, not the raw model echo.
X-Cashu handlers do not rewrite the upstream's echoed model string, so
without the routed model the settle would look up whatever wire name the
upstream reported. With ``model_obj`` given, the echo must be irrelevant.
"""
from routstr.upstream import GenericUpstreamProvider
provider = GenericUpstreamProvider("http://upstream.example", "key", 1.0)
response = dict(RESPONSE, model="totally-unknown-wire-name")
with patch(
"routstr.proxy.get_model_instance", return_value=WINNER
) as alias_lookup:
cost = await provider.get_x_cashu_cost(
response, max_cost_for_model=100_000, model_obj=SERVED
)
assert cost is not None
assert cost.total_msats == 10_000
alias_lookup.assert_not_called()
+5 -2
View File
@@ -390,8 +390,11 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
revert_mock = AsyncMock(return_value=True)
with (
patch.object(proxy_module, "get_model_instance", return_value=MagicMock()),
patch.object(proxy_module, "get_provider_for_model", return_value=[upstream]),
patch.object(
proxy_module,
"get_candidates",
return_value=[(MagicMock(), upstream)],
),
patch.object(
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
),
@@ -3,6 +3,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.exc import SQLAlchemyError
import routstr.auth as auth_module
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
@@ -152,6 +154,59 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None:
await engine.dispose()
@pytest.mark.asyncio
async def test_generic_background_settlement_uses_explicit_reservation() -> None:
engine = await _engine()
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test-key", provider_fee=1.0
)
key = ApiKey(hashed_key="generic-key", balance=1_000)
cost = MaxCostData(
base_msats=500,
input_msats=0,
output_msats=0,
total_msats=500,
)
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)
context_token = auth_module._current_reservation.set(None)
try:
with (
patch(
"routstr.upstream.base.create_session",
side_effect=lambda: AsyncSession(engine, expire_on_commit=False),
),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
auth_module.adjust_payment_for_tokens,
),
patch("routstr.auth.calculate_cost", AsyncMock(return_value=cost)),
):
await provider._finalize_generic_streaming_payment(
key.hashed_key,
500,
"audio/speech",
model_obj=None,
provider_fee=provider.provider_fee,
reservation_snapshot=snapshot,
)
finally:
auth_module._current_reservation.reset(context_token)
async with AsyncSession(engine, expire_on_commit=False) as session:
settled_key = await session.get(ApiKey, key.hashed_key)
record = await session.get(ReservationRelease, snapshot.release_id)
assert settled_key is not None
assert (settled_key.balance, settled_key.reserved_balance) == (500, 0)
assert record is not None and record.status == "charged"
await engine.dispose()
@pytest.mark.asyncio
async def test_streaming_release_is_terminal_and_suppresses_background_charge() -> None:
provider = BaseUpstreamProvider(
@@ -224,7 +279,7 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() ->
{"model": "test", "usage": None},
session,
500,
snapshot,
reservation_snapshot=snapshot,
)
await session.refresh(first)
+5 -2
View File
@@ -369,8 +369,11 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
revert_mock = AsyncMock(return_value=True)
with (
patch.object(proxy_module, "get_model_instance", return_value=MagicMock()),
patch.object(proxy_module, "get_provider_for_model", return_value=[upstream]),
patch.object(
proxy_module,
"get_candidates",
return_value=[(MagicMock(), upstream)],
),
patch.object(
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
),