mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #629 from jeroenubbink/fix/bill-served-candidate
fix: bill and forward failover requests as the provider that actually served
This commit is contained in:
+11
-7
@@ -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
|
||||
|
||||
|
||||
+13
-2
@@ -4,7 +4,7 @@ import math
|
||||
import random
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import case
|
||||
@@ -26,6 +26,9 @@ 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")
|
||||
|
||||
@@ -771,12 +774,18 @@ async def adjust_payment_for_tokens(
|
||||
response_data: dict,
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
model_obj: "Model | None",
|
||||
provider_fee: float | None,
|
||||
) -> dict:
|
||||
"""
|
||||
Adjusts the payment based on token usage in the response.
|
||||
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``.
|
||||
"""
|
||||
@@ -861,7 +870,9 @@ async def adjust_payment_for_tokens(
|
||||
extra={"error": str(e), "fee_msats": fee_msats},
|
||||
)
|
||||
|
||||
match await calculate_cost(response_data, deducted_max_cost):
|
||||
match await calculate_cost(
|
||||
response_data, deducted_max_cost, model_obj, provider_fee
|
||||
):
|
||||
case MaxCostData() as cost:
|
||||
logger.debug(
|
||||
"Using max cost data (no token adjustment)",
|
||||
|
||||
@@ -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
|
||||
|
||||
+73
-42
@@ -37,10 +37,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)
|
||||
|
||||
|
||||
@@ -72,32 +71,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]:
|
||||
@@ -137,7 +148,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
|
||||
@@ -160,7 +171,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,
|
||||
@@ -283,25 +294,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}'",
|
||||
@@ -309,9 +315,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
|
||||
@@ -326,7 +333,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:
|
||||
@@ -364,7 +371,7 @@ async def proxy(
|
||||
"status_code": e.status_code,
|
||||
},
|
||||
)
|
||||
if i == len(upstreams) - 1:
|
||||
if i == len(candidates) - 1:
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
@@ -391,12 +398,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"):
|
||||
@@ -426,7 +433,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(
|
||||
@@ -441,7 +448,31 @@ 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)
|
||||
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)
|
||||
continue
|
||||
max_cost_for_model = candidate_max
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
|
||||
try:
|
||||
@@ -542,7 +573,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"):
|
||||
@@ -622,12 +653,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)
|
||||
return create_upstream_error_response(e, request)
|
||||
|
||||
|
||||
@@ -799,6 +799,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
background_tasks: BackgroundTasks,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -841,6 +842,8 @@ class BaseUpstreamProvider:
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
@@ -1013,6 +1016,8 @@ class BaseUpstreamProvider:
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
@@ -1109,6 +1114,7 @@ class BaseUpstreamProvider:
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1155,6 +1161,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
await session.refresh(key)
|
||||
@@ -1250,6 +1258,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1293,6 +1302,8 @@ class BaseUpstreamProvider:
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
@@ -1422,6 +1433,8 @@ class BaseUpstreamProvider:
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
@@ -1540,6 +1553,7 @@ class BaseUpstreamProvider:
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1589,6 +1603,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
await session.refresh(key)
|
||||
@@ -1693,11 +1709,15 @@ 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=None,
|
||||
provider_fee=None,
|
||||
)
|
||||
logger.debug(
|
||||
"Finalized generic streaming payment in background",
|
||||
@@ -1722,6 +1742,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
async def stream_with_cost(
|
||||
max_cost_for_model: int,
|
||||
@@ -1787,6 +1808,8 @@ class BaseUpstreamProvider:
|
||||
fallback,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
|
||||
@@ -1938,6 +1961,8 @@ class BaseUpstreamProvider:
|
||||
combined_data,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
self.inject_cost_metadata(
|
||||
@@ -1985,6 +2010,7 @@ class BaseUpstreamProvider:
|
||||
deducted_max_cost: int,
|
||||
path: str,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
try:
|
||||
content = await response.aread()
|
||||
@@ -2009,6 +2035,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
self.inject_cost_metadata(response_json, cost_data, key)
|
||||
@@ -2103,6 +2131,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
response_json = messages_dispatch.coerce_litellm_payload(result)
|
||||
@@ -2114,6 +2143,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
self.inject_cost_metadata(response_json, cost_data, key)
|
||||
|
||||
@@ -2153,6 +2184,7 @@ class BaseUpstreamProvider:
|
||||
requested_model,
|
||||
mint,
|
||||
request_id,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
response_json = messages_dispatch.coerce_litellm_payload(result)
|
||||
@@ -2160,7 +2192,9 @@ 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
|
||||
@@ -2205,6 +2239,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Re-emit a litellm Anthropic-event iterator as live SSE bytes
|
||||
with cost reconciliation appended at end of stream."""
|
||||
@@ -2255,6 +2290,8 @@ class BaseUpstreamProvider:
|
||||
fallback,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
return (
|
||||
@@ -2328,6 +2365,8 @@ class BaseUpstreamProvider:
|
||||
combined_data,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
self.inject_cost_metadata(
|
||||
combined_data, cost_data, fresh_key
|
||||
@@ -2365,6 +2404,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.
|
||||
|
||||
@@ -2456,7 +2496,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(
|
||||
@@ -2673,6 +2713,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -2689,6 +2730,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
path,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -2704,6 +2746,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
path,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -2753,6 +2796,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
background_tasks,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
@@ -2766,6 +2810,7 @@ class BaseUpstreamProvider:
|
||||
session,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -3018,6 +3063,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -3033,6 +3079,7 @@ class BaseUpstreamProvider:
|
||||
session,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -3209,13 +3256,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
|
||||
@@ -3229,6 +3282,8 @@ class BaseUpstreamProvider:
|
||||
match await calculate_cost(
|
||||
response_data,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
):
|
||||
case MaxCostData() as cost:
|
||||
logger.debug(
|
||||
@@ -3373,6 +3428,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.
|
||||
|
||||
@@ -3446,7 +3502,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":
|
||||
@@ -3551,6 +3607,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.
|
||||
|
||||
@@ -3572,7 +3629,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
|
||||
@@ -3701,6 +3760,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.
|
||||
|
||||
@@ -3744,6 +3804,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(
|
||||
@@ -3754,6 +3815,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -3939,6 +4001,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)
|
||||
@@ -4231,6 +4294,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)
|
||||
@@ -4281,6 +4345,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.
|
||||
|
||||
@@ -4325,6 +4390,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(
|
||||
@@ -4335,6 +4401,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -4362,6 +4429,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.
|
||||
|
||||
@@ -4420,7 +4488,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":
|
||||
@@ -4526,6 +4594,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(
|
||||
@@ -4536,7 +4605,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
|
||||
|
||||
@@ -82,7 +82,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)
|
||||
|
||||
@@ -116,7 +118,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)
|
||||
|
||||
@@ -155,7 +159,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)
|
||||
|
||||
@@ -223,15 +229,18 @@ 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
|
||||
)
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await asyncio.gather(*[finalize() for _ in range(n_requests)])
|
||||
# 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() for _ in range(n_requests)])
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
@@ -279,7 +288,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)
|
||||
|
||||
@@ -344,15 +355,18 @@ 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
|
||||
)
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await asyncio.gather(finalize(), finalize())
|
||||
# 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(), finalize())
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -0,0 +1,546 @@
|
||||
"""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 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"
|
||||
|
||||
|
||||
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),
|
||||
)
|
||||
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],
|
||||
) -> 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
|
||||
|
||||
|
||||
@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 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,
|
||||
) -> 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
|
||||
@@ -58,7 +58,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)
|
||||
@@ -129,7 +129,7 @@ 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
|
||||
key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -502,7 +502,12 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
|
||||
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,
|
||||
) -> dict:
|
||||
captured_cost_call["combined_data"] = combined_data
|
||||
captured_cost_call["max_cost"] = max_cost
|
||||
@@ -591,7 +596,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
|
||||
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,
|
||||
) -> dict:
|
||||
captured["combined_data"] = combined_data
|
||||
return fake_cost
|
||||
|
||||
@@ -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()
|
||||
@@ -355,8 +355,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)
|
||||
),
|
||||
|
||||
@@ -362,8 +362,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)
|
||||
),
|
||||
|
||||
Reference in New Issue
Block a user