Compare commits

..
Author SHA1 Message Date
redshift da1065fa50 fix: refund users when model returns zero tokens with non-zero cost
When a model produces empty output (zero tokens reported) but the cost
calculation returns a non-zero USD cost, the previous behavior was to
charge the full amount as output_msats. This is incorrect — if no tokens
were produced, the user should not be charged.

Now the else branch:
- Logs a warning with the usd_cost and model for debugging
- Sets input_msats, output_msats, and cost_in_msats to 0 (full refund)
2026-05-06 19:15:05 +08:00
29 changed files with 257 additions and 2339 deletions
+1 -1
View File
@@ -4,7 +4,7 @@ FROM node:23-alpine AS ui-builder
WORKDIR /app/ui
# Install pnpm
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
# Copy UI source
COPY ui/package.json ui/pnpm-lock.yaml* ./
+3 -11
View File
@@ -78,15 +78,9 @@ Which mints to accept payments from:
Automatic profit withdrawal:
| Setting | Description | Default |
| ------------------------------------ | ---------------------------------------------------------------------------------------------------------- | ------- |
| **Lightning Address** | Your LN address for withdrawals | — |
| **Minimum Payout (sat)** | Min available balance (in sats) before profit is paid out. Applies to both `sat` and `msat` mints (auto-converted). | `210` |
| **Payout Interval (seconds)** | How often the payout loop wakes up and checks balances | `900` |
All payout amounts must be positive. Set the minimums above your wallet's
minimum-invoice constraint (typically 1 sat) and high enough to amortise
routing fees.
| Setting | Description |
| --------------------- | ------------------------------- |
| **Lightning Address** | Your LN address for withdrawals |
### Security
@@ -132,8 +126,6 @@ Use environment variables for:
| `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` |
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` |
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
| `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` |
| `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` |
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
+5 -26
View File
@@ -195,15 +195,14 @@ def create_model_mappings(
# Add to unique models
base_id = get_base_model_id(model_to_use.id)
unique_key = model_to_use.forwarded_model_id or base_id
if not is_openrouter or unique_key not in unique_models:
if not is_openrouter or base_id not in unique_models:
unique_model = model_to_use.copy(
update={
"id": base_id,
"upstream_provider_id": upstream.provider_type,
}
)
unique_models[unique_key] = unique_model
unique_models[base_id] = unique_model
# Get all aliases for this model
aliases = resolve_model_alias(
@@ -273,19 +272,18 @@ def create_model_mappings(
continue
base_id = get_base_model_id(model_to_use.id)
unique_key = model_to_use.forwarded_model_id or base_id
is_openrouter = (
getattr(upstream_for_override, "base_url", "")
== "https://openrouter.ai/api/v1"
)
if not is_openrouter or unique_key not in unique_models:
if not is_openrouter or base_id not in unique_models:
unique_model = model_to_use.copy(
update={
"id": base_id,
"upstream_provider_id": upstream_for_override.provider_type,
}
)
unique_models[unique_key] = unique_model
unique_models[base_id] = unique_model
try:
aliases = resolve_model_alias(
@@ -324,26 +322,7 @@ def create_model_mappings(
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is.
forwarded_model_id is the most specific identifier (set per-provider
instance), so a match there should beat a model_id match. This way,
when multiple providers have the same model_id but different
forwarded_model_ids, the one whose forwarded_model_id equals the
requested alias wins.
"""
if (
model.forwarded_model_id
and model.forwarded_model_id.lower() == alias
):
return 5
if (
model.id
and model.id.lower() == alias
):
return 4
"""Rank how strong the mapping of alias->model is."""
model_base = get_base_model_id(model.id)
if model_base == alias:
return 3
+5 -22
View File
@@ -211,20 +211,8 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
_refund_cache[key] = (expiry, value)
async def _lookup_key_no_create(
bearer_value: str, session: AsyncSession
) -> ApiKey | None:
"""Look up an existing API key without creating one Used by the refund endpoint"""
if bearer_value.startswith("sk-"):
return await session.get(ApiKey, bearer_value[3:])
if bearer_value.startswith("cashu"):
hashed = hashlib.sha256(bearer_value.encode()).hexdigest()
return await session.get(ApiKey, hashed)
return None
async def _restore_balance(
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int
) -> None:
"""Restore balance after a failed refund mint attempt."""
restore_stmt = (
@@ -239,7 +227,7 @@ async def _restore_balance(
await session.commit()
logger.info(
"refund_wallet_endpoint: balance restored after mint failure",
extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
extra={"hashed_key": hashed_key, "restored_balance": balance},
)
@@ -294,12 +282,7 @@ async def refund_wallet_endpoint(
)
bearer_value: str = authorization[7:]
key: ApiKey | None = await _lookup_key_no_create(bearer_value, session)
if key is None:
raise HTTPException(
status_code=401,
detail="Key not found. Deposit first via /v1/wallet/create before requesting a refund.",
)
key: ApiKey = await validate_bearer_key(bearer_value, session)
if key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value):
@@ -390,11 +373,11 @@ async def refund_wallet_endpoint(
except HTTPException:
# Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
raise
except Exception as e:
# Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
error_msg = str(e)
if (
"mint" in error_msg.lower()
+2 -8
View File
@@ -6,7 +6,6 @@ from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, RootModel
from pydantic.v1 import ValidationError as PydanticValidationError
from sqlmodel import select
from ..payment.models import _row_to_model, list_models
@@ -161,13 +160,8 @@ async def update_settings(request: Request, update: SettingsUpdate) -> dict:
if field in settings_data:
del settings_data[field]
try:
async with create_session() as session:
new_settings = await SettingsService.update(settings_data, session)
except PydanticValidationError as e:
# Surface validation issues (e.g. non-positive payout amounts)
# as a clean 400 instead of a 500.
raise HTTPException(status_code=400, detail=e.errors()) from e
async with create_session() as session:
new_settings = await SettingsService.update(settings_data, session)
data = new_settings.dict()
if "upstream_api_key" in data:
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
-9
View File
@@ -41,15 +41,6 @@ class Settings(BaseSettings):
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
# Lightning payout configuration
# Minimum available balance (in satoshis) before profit is paid out over
# Lightning
min_payout_sat: int = Field(default=210, gt=0, env="MIN_PAYOUT_SAT")
# Interval (seconds) between periodic payout attempts. Must be positive.
payout_interval_seconds: int = Field(
default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS"
)
# Pricing
# Default behavior: derive pricing from MODELS
# If fixed_pricing is True -> use fixed_cost_per_request and ignore tokens
+3 -8
View File
@@ -55,14 +55,9 @@ def _on_tagged_release() -> bool:
if tag:
return tag.lstrip("v") == BASE_VERSION
described = _run_git("describe", "--tags", "--exact-match", "HEAD")
if described and described.lstrip("v") == BASE_VERSION:
return True
current_branch = _run_git("rev-parse", "--abbrev-ref", "HEAD")
if current_branch == "main":
branch_tag = _run_git("describe", "--tags", "--abbrev=0", "HEAD")
if branch_tag and branch_tag.lstrip("v") == BASE_VERSION:
return True
return False
if not described:
return False
return described.lstrip("v") == BASE_VERSION
@lru_cache(maxsize=1)
+164 -288
View File
@@ -18,10 +18,6 @@ class CostData(BaseModel):
total_usd: float = 0.0
input_tokens: int = 0
output_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_input_tokens: int = 0
cache_read_msats: int = 0
cache_creation_msats: int = 0
class MaxCostData(CostData):
@@ -33,10 +29,11 @@ class CostDataError(BaseModel):
code: str
async def calculate_cost(
async def calculate_cost( # todo: can be sync
response_data: dict, max_cost: int, session: AsyncSession
) -> CostData | MaxCostData | CostDataError:
"""Calculate the cost of an API request based on token usage.
"""
Calculate the cost of an API request based on token usage.
Args:
response_data: Response data containing usage information
@@ -54,7 +51,6 @@ async def calculate_cost(
},
)
# Check for usage data
if "usage" not in response_data or response_data["usage"] is None:
logger.warning(
"No usage data in response — billing at MaxCostData with zero "
@@ -78,25 +74,77 @@ async def calculate_cost(
total_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
cache_read_msats=0,
cache_creation_msats=0,
)
usage_data = response_data["usage"]
# Extract token counts
input_tokens = _extract_token_pair(usage_data, "prompt_tokens", "input_tokens")
output_tokens = _extract_token_pair(usage_data, "completion_tokens", "output_tokens")
def parse_token_count(value: object) -> int:
if isinstance(value, bool):
return 0
if isinstance(value, int):
return max(0, value)
if isinstance(value, float):
return max(0, int(value))
if isinstance(value, str):
try:
return max(0, int(float(value)))
except ValueError:
return 0
return 0
# Extract cache tokens (handles OpenAI vs Anthropic formats)
cache_read_tokens, cache_creation_tokens, input_tokens = _extract_cache_tokens(
usage_data, input_tokens
input_tokens = parse_token_count(usage_data.get("prompt_tokens", 0))
output_tokens = parse_token_count(usage_data.get("completion_tokens", 0))
input_tokens = (
input_tokens
if input_tokens != 0
else parse_token_count(usage_data.get("input_tokens", 0))
)
output_tokens = (
output_tokens
if output_tokens != 0
else parse_token_count(usage_data.get("output_tokens", 0))
)
input_tokens = (
input_tokens
if input_tokens != 0
else parse_token_count(response_data.get("usage", {}).get("input_tokens", 0))
)
output_tokens = (
output_tokens
if output_tokens != 0
else parse_token_count(response_data.get("usage", {}).get("output_tokens", 0))
)
usd_cost = 0.0
input_usd = 0.0
output_usd = 0.0
if "cost_details" in usage_data:
usd_cost = float(
usage_data["cost_details"].get("upstream_inference_cost", 0) or 0
)
input_usd = float(
usage_data["cost_details"].get("upstream_inference_prompt_cost", 0) or 0
)
output_usd = float(
usage_data["cost_details"].get("upstream_inference_completions_cost", 0)
or 0
)
# Fallback to cost field if upstream_inference_cost is 0
if usd_cost == 0 and "cost" in usage_data:
try:
usd_cost = float(usage_data.get("cost", 0) or 0)
except Exception:
pass
MSATS_PER_1K_INPUT_TOKENS: float = (
float(settings.fixed_per_1k_input_tokens) * 1000.0
)
MSATS_PER_1K_OUTPUT_TOKENS: float = (
float(settings.fixed_per_1k_output_tokens) * 1000.0
)
# Try USD cost first
usd_cost = _resolve_usd_cost(usage_data, response_data)
if usd_cost > 0:
if input_tokens == 0 and output_tokens == 0:
logger.warning(
@@ -115,21 +163,53 @@ async def calculate_cost(
},
)
try:
input_usd = _coerce_usd(
usage_data.get("cost_details", {}).get("input_cost", 0)
sats_per_usd = 1.0 / sats_usd_price()
cost_in_sats = usd_cost * sats_per_usd
cost_in_msats = math.ceil(cost_in_sats * 1000)
input_msats = 0
output_msats = 0
if input_usd > 0 or output_usd > 0:
input_msats = int((input_usd * sats_per_usd) * 1000)
output_msats = int((output_usd * sats_per_usd) * 1000)
else:
total_tokens = input_tokens + output_tokens
if total_tokens > 0:
input_ratio = input_tokens / total_tokens
input_msats = int(cost_in_msats * input_ratio)
output_msats = cost_in_msats - input_msats
else:
# No tokens reported — model produced empty output, issue full refund
logger.warning(
"Zero tokens with non-zero USD cost — issuing full refund",
extra={
"usd_cost": usd_cost,
"model": response_data.get("model", "unknown"),
},
)
input_msats = 0
output_msats = 0
cost_in_msats = 0
logger.info(
"Using cost from usage data/details",
extra={
"usd_cost": usd_cost,
"cost_in_sats": cost_in_sats,
"cost_in_msats": cost_in_msats,
"model": response_data.get("model", "unknown"),
},
)
output_usd = _coerce_usd(
usage_data.get("cost_details", {}).get("output_cost", 0)
)
return _calculate_from_usd_cost(
usd_cost,
input_usd,
output_usd,
input_tokens,
cache_read_tokens,
cache_creation_tokens,
output_tokens,
response_data,
return CostData(
base_msats=0,
input_msats=input_msats,
output_msats=output_msats,
total_msats=cost_in_msats,
total_usd=usd_cost,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
except Exception as e:
logger.warning(
@@ -140,22 +220,57 @@ async def calculate_cost(
"model": response_data.get("model", "unknown"),
},
)
# Fall through to token-based calculation
# Fall back to token-based pricing
try:
pricing_rates = _get_pricing_rates(response_data)
except ValueError as e:
return CostDataError(message=str(e), code="pricing_error")
if not settings.fixed_pricing:
response_model = response_data.get("model", "")
logger.debug(
"Using model-based pricing",
extra={"model": response_model},
)
if pricing_rates is None:
input_rate = float(settings.fixed_per_1k_input_tokens) * 1000.0
output_rate = float(settings.fixed_per_1k_output_tokens) * 1000.0
cache_read_rate = input_rate
cache_creation_rate = input_rate
else:
input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates
from ..proxy import get_model_instance
if not (input_rate and output_rate):
model_obj = get_model_instance(response_model)
if not model_obj:
logger.error(
"Invalid model in response",
extra={"response_model": response_model},
)
return CostDataError(
message=f"Invalid model in response: {response_model}",
code="model_not_found",
)
if not model_obj.sats_pricing:
logger.error(
"Model pricing not defined",
extra={"model": response_model, "model_id": response_model},
)
return CostDataError(
message="Model pricing not defined", code="pricing_not_found"
)
try:
mspp = float(model_obj.sats_pricing.prompt)
mspc = float(model_obj.sats_pricing.completion)
except Exception:
return CostDataError(message="Invalid pricing data", code="pricing_invalid")
MSATS_PER_1K_INPUT_TOKENS = mspp * 1_000_000.0
MSATS_PER_1K_OUTPUT_TOKENS = mspc * 1_000_000.0
logger.info(
"Applied model-specific pricing",
extra={
"model": response_model,
"input_price_msats_per_1k": MSATS_PER_1K_INPUT_TOKENS,
"output_price_msats_per_1k": MSATS_PER_1K_OUTPUT_TOKENS,
},
)
if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS):
logger.warning(
"No token pricing configured — billing at flat MaxCostData. "
"Token counts %s in the upstream response but cannot be "
@@ -178,243 +293,12 @@ async def calculate_cost(
total_msats=max_cost,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=0,
cache_creation_msats=0,
)
return _calculate_from_tokens(
input_tokens,
output_tokens,
cache_read_tokens,
cache_creation_tokens,
input_rate,
output_rate,
cache_read_rate,
cache_creation_rate,
response_data,
)
calc_input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)
# ============================================================================
# Helper Functions (ordered by call sequence in calculate_cost)
# ============================================================================
def parse_token_count(value: object) -> int:
"""Parse a token count from various formats (int, float, str, bool)."""
if isinstance(value, bool):
return 0
if isinstance(value, int):
return max(0, value)
if isinstance(value, float):
return max(0, int(value))
if isinstance(value, str):
try:
return max(0, int(float(value)))
except ValueError:
return 0
return 0
def _coerce_usd(value: object) -> float:
"""Coerce a value to USD float, handling various formats safely."""
if value is None or isinstance(value, bool):
return 0.0
if not isinstance(value, (int, float, str)):
return 0.0
try:
return max(0.0, float(value))
except (TypeError, ValueError):
return 0.0
def _extract_token_pair(
usage_data: dict, standard_field: str, alt_field: str
) -> int:
"""Extract token count trying two field names in order."""
value = parse_token_count(usage_data.get(standard_field, 0))
if value > 0:
return value
return parse_token_count(usage_data.get(alt_field, 0))
def _extract_cache_tokens(usage_data: dict, input_tokens: int) -> tuple[int, int, int]:
"""Extract cache tokens, handling OpenAI vs Anthropic formats.
Returns: (cache_read_tokens, cache_creation_tokens, adjusted_input_tokens)
"""
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
cache_creation = parse_token_count(
usage_data.get("cache_creation_input_tokens", 0)
)
# OpenAI: cache is included in input_tokens, subtract it
prompt_details = usage_data.get("prompt_tokens_details")
if isinstance(prompt_details, dict) and not cache_read:
openai_cached = parse_token_count(prompt_details.get("cached_tokens", 0))
if openai_cached:
cache_read = openai_cached
input_tokens = max(0, input_tokens - cache_read)
return cache_read, cache_creation, input_tokens
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
"""Resolve USD cost with clear priority order.
Priority: cost_details.total_cost → total_cost → cost (in both usage and response).
"""
cost_details = usage_data.get("cost_details")
if isinstance(cost_details, dict):
cost = _coerce_usd(cost_details.get("total_cost"))
if cost > 0:
return cost
for source in [usage_data, response_data]:
if not isinstance(source, dict):
continue
for field in ("total_cost", "cost"):
cost = _coerce_usd(source.get(field))
if cost > 0:
return cost
return 0.0
def _get_pricing_rates(
response_data: dict,
) -> tuple[float, float, float, float] | None:
"""Get model-based pricing rates or None if using fixed pricing.
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate)
"""
if settings.fixed_pricing:
return None
from ..proxy import get_model_instance
response_model = response_data.get("model", "")
model_obj = get_model_instance(response_model)
if not model_obj:
logger.error("Invalid model in response", extra={"response_model": response_model})
raise ValueError(f"Invalid model: {response_model}")
if not model_obj.sats_pricing:
logger.error(
"Model pricing not defined",
extra={"model": response_model, "model_id": response_model},
)
raise ValueError("Model pricing not defined")
try:
mspp = float(model_obj.sats_pricing.prompt)
mspc = float(model_obj.sats_pricing.completion)
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
mspp_1k = mspp * 1_000_000.0
mspc_1k = mspc * 1_000_000.0
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
logger.info(
"Applied model-specific pricing",
extra={
"model": response_model,
"input_price_msats_per_1k": mspp_1k,
"output_price_msats_per_1k": mspc_1k,
"cache_read_price_msats_per_1k": mscr_1k,
"cache_write_price_msats_per_1k": mscw_1k,
},
)
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
except Exception as e:
logger.error("Invalid pricing data", extra={"error": str(e)})
raise ValueError("Invalid pricing data") from e
def _calculate_from_usd_cost(
usd_cost: float,
input_usd: float,
output_usd: float,
input_tokens: int,
cache_read_tokens: int,
cache_creation_tokens: int,
output_tokens: int,
response_data: dict,
) -> CostData:
"""Calculate cost from USD figures, deriving input/output split from tokens."""
sats_per_usd = 1.0 / sats_usd_price()
cost_in_sats = usd_cost * sats_per_usd
cost_in_msats = math.ceil(cost_in_sats * 1000)
if input_usd > 0 or output_usd > 0:
input_msats = int((input_usd * sats_per_usd) * 1000)
output_msats = int((output_usd * sats_per_usd) * 1000)
else:
effective_input_tokens = (
input_tokens + cache_read_tokens + cache_creation_tokens
)
total_tokens = effective_input_tokens + output_tokens
input_msats = (
int(cost_in_msats * effective_input_tokens / total_tokens)
if total_tokens > 0
else 0
)
output_msats = cost_in_msats - input_msats
logger.info(
"Using cost from usage data/details",
extra={
"usd_cost": usd_cost,
"cost_in_sats": cost_in_sats,
"cost_in_msats": cost_in_msats,
"model": response_data.get("model", "unknown"),
},
)
return CostData(
base_msats=0,
input_msats=input_msats,
output_msats=output_msats,
total_msats=cost_in_msats,
total_usd=usd_cost,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=0,
cache_creation_msats=0,
)
def _calculate_from_tokens(
input_tokens: int,
output_tokens: int,
cache_read_tokens: int,
cache_creation_tokens: int,
input_rate: float,
output_rate: float,
cache_read_rate: float,
cache_creation_rate: float,
response_data: dict,
) -> CostData:
"""Calculate cost from token counts using pricing rates."""
calc_input_msats = round(input_tokens / 1000 * input_rate, 3)
calc_output_msats = round(output_tokens / 1000 * output_rate, 3)
calc_cache_read_msats = round(cache_read_tokens / 1000 * cache_read_rate, 3)
calc_cache_write_msats = round(
cache_creation_tokens / 1000 * cache_creation_rate, 3
)
token_based_cost = math.ceil(
calc_input_msats
+ calc_output_msats
+ calc_cache_read_msats
+ calc_cache_write_msats
)
calc_output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
token_based_cost = math.ceil(calc_input_msats + calc_output_msats)
total_usd = (token_based_cost / 1000.0) * sats_usd_price()
logger.info(
@@ -422,12 +306,8 @@ def _calculate_from_tokens(
extra={
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": cache_read_tokens,
"cache_creation_input_tokens": cache_creation_tokens,
"input_cost_msats": calc_input_msats,
"output_cost_msats": calc_output_msats,
"cache_read_cost_msats": calc_cache_read_msats,
"cache_creation_cost_msats": calc_cache_write_msats,
"total_cost_msats": token_based_cost,
"total_usd": total_usd,
"model": response_data.get("model", "unknown"),
@@ -442,8 +322,4 @@ def _calculate_from_tokens(
total_usd=total_usd,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=int(calc_cache_read_msats),
cache_creation_msats=int(calc_cache_write_msats),
)
+1 -1
View File
@@ -363,7 +363,7 @@ async def _update_sats_pricing_once() -> None:
for m in upstream.get_cached_models()
]
upstream._models_cache = updated_models
upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models}
upstream._models_by_id = {m.id: m for m in updated_models}
updated_count += len(updated_models)
if updated_count > 0:
+1 -43
View File
@@ -1,9 +1,8 @@
import json
from pathlib import Path
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, Response, StreamingResponse
from fastapi.responses import Response, StreamingResponse
from sqlmodel import select
from .algorithm import create_model_mappings
@@ -151,51 +150,10 @@ async def refresh_model_maps_periodically() -> None:
)
_API_PATH_PREFIXES = ("v1/", "responses")
_NOT_FOUND_HTML_FILE = Path(__file__).parent.parent / "ui_out" / "404.html"
def _read_not_found_html() -> str | None:
try:
return _NOT_FOUND_HTML_FILE.read_text(encoding="utf-8")
except OSError:
return None
_NOT_FOUND_HTML: str | None = _read_not_found_html()
def _build_not_found_response(request: Request, path: str) -> Response:
"""Return a 404 for unknown paths.
"""
accept = request.headers.get("accept", "").lower()
prefers_json = "application/json" in accept and "text/html" not in accept
request_id = getattr(request.state, "request_id", "unknown")
if not prefers_json and _NOT_FOUND_HTML is not None:
return HTMLResponse(content=_NOT_FOUND_HTML, status_code=404)
return JSONResponse(
status_code=404,
content={
"error": {
"message": f"Path '/{path}' not found",
"type": "not_found",
"code": 404,
},
"request_id": request_id,
},
)
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session)
) -> Response | StreamingResponse:
if not path.startswith(_API_PATH_PREFIXES):
return _build_not_found_response(request, path)
headers = dict(request.headers)
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
+21 -290
View File
@@ -42,23 +42,11 @@ from ..payment.models import (
from ..payment.price import sats_usd_price
from ..wallet import recieve_token, send_token
from . import messages_dispatch
from .count_tokens import count_tokens_locally
from .litellm_routing import detect_litellm_prefix
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.
"""
if not content_type:
return False
main = content_type.split(";", 1)[0].strip().lower()
if main in ("application/json", "text/json"):
return True
return main.startswith("application/") and main.endswith("+json")
class TopupData(BaseModel):
"""Universal top-up data schema for Lightning Network invoices."""
@@ -152,51 +140,6 @@ class BaseUpstreamProvider:
"can_show_balance": False,
}
@staticmethod
def _fold_cache_into_input_tokens(usage: object) -> None:
"""Fold cache token counts into ``input_tokens`` / ``prompt_tokens``.
Cost calculation has already used the per-bucket counts to bill the
request correctly; what the client sees in the visible token total
should be a single rolled-up prompt count *including* the cache
portion. The standalone ``cache_read_input_tokens`` /
``cache_creation_input_tokens`` fields are left in place for clients
that want the breakdown.
For Anthropic-shaped responses (``input_tokens`` present), the cache
fields are forced to ``0`` when the upstream omitted them, so the
client always sees a consistent shape.
"""
if not isinstance(usage, dict):
return
# Normalise missing cache fields to 0 on Anthropic-shaped usage so
# downstream consumers can rely on them being present.
if "input_tokens" in usage:
usage.setdefault("cache_read_input_tokens", 0)
usage.setdefault("cache_creation_input_tokens", 0)
try:
cache_read = int(usage.get("cache_read_input_tokens") or 0)
cache_creation = int(usage.get("cache_creation_input_tokens") or 0)
except (TypeError, ValueError):
return
extra = cache_read + cache_creation
if extra <= 0:
return
if "input_tokens" in usage:
try:
usage["input_tokens"] = int(usage.get("input_tokens") or 0) + extra
except (TypeError, ValueError):
pass
if "prompt_tokens" in usage:
try:
usage["prompt_tokens"] = (
int(usage.get("prompt_tokens") or 0) + extra
)
except (TypeError, ValueError):
pass
def inject_cost_metadata(
self,
response_json: dict,
@@ -220,7 +163,6 @@ class BaseUpstreamProvider:
response_json["usage"]["cost"] = total_usd
response_json["usage"]["cost_sats"] = sats_cost
response_json["usage"]["remaining_balance_msats"] = key.balance
self._fold_cache_into_input_tokens(response_json["usage"])
# Inject into Anthropic nested usage block if present
if (
@@ -229,7 +171,6 @@ class BaseUpstreamProvider:
and "usage" in response_json["message"]
):
response_json["message"]["usage"]["sats_cost"] = sats_cost
self._fold_cache_into_input_tokens(response_json["message"]["usage"])
# Unified Routstr metadata
response_json["metadata"] = response_json.get("metadata", {})
@@ -535,8 +476,7 @@ class BaseUpstreamProvider:
async def forward_upstream_error_response(
self, request: Request, path: str, upstream_response: httpx.Response
) -> Response:
"""Log upstream errors and forward the response in a JSON envelope.
"""
"""Log upstream errors and forward the upstream response unchanged."""
status_code = upstream_response.status_code
headers = dict(upstream_response.headers)
content_type = headers.get("content-type") or headers.get("Content-Type", "")
@@ -558,10 +498,9 @@ class BaseUpstreamProvider:
message, upstream_code = self._extract_upstream_error_message(body_bytes)
body_preview = body_bytes.decode("utf-8", errors="ignore").strip()[:500]
is_json_body = _is_json_content_type(content_type)
logger.warning(
"Forwarding upstream error response",
"Forwarding upstream error response as-is",
extra={
"path": path,
"provider": self.provider_type,
@@ -573,7 +512,6 @@ class BaseUpstreamProvider:
"body_preview": body_preview,
"body_read_error": body_read_error,
"method": request.method,
"json_normalized": not is_json_body,
},
)
@@ -601,40 +539,17 @@ class BaseUpstreamProvider:
):
headers.pop(header_name, None)
if is_json_body:
if not content_type:
headers.pop("content-type", None)
headers.pop("Content-Type", None)
media_type = content_type or None
return Response(
content=body_bytes,
status_code=status_code,
headers=headers,
media_type=media_type,
)
if not content_type:
headers.pop("content-type", None)
headers.pop("Content-Type", None)
# Non-JSON upstream error (HTML, plain text, empty, ...). Wrap it in
# the standard JSON envelope so callers don't need a second parser.
for header_name in ("content-type", "Content-Type"):
headers.pop(header_name, None)
envelope = {
"error": {
"message": message or "Upstream returned a non-JSON error response",
"type": "upstream_error",
"code": upstream_code or status_code,
"upstream_status": status_code,
"upstream_content_type": content_type or None,
"upstream_body_preview": body_preview or None,
},
"request_id": getattr(request.state, "request_id", None),
}
media_type = content_type or None
return Response(
content=json.dumps(envelope).encode(),
content=body_bytes,
status_code=status_code,
headers=headers,
media_type="application/json",
media_type=media_type,
)
async def handle_streaming_chat_completion(
@@ -911,7 +826,6 @@ class BaseUpstreamProvider:
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
self._fold_cache_into_input_tokens(response_json["usage"])
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
@@ -1286,7 +1200,6 @@ class BaseUpstreamProvider:
response_json["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
self._fold_cache_into_input_tokens(response_json["usage"])
# Keep detailed cost
response_json["metadata"] = response_json.get("metadata", {})
@@ -1414,42 +1327,6 @@ class BaseUpstreamProvider:
last_model_seen: str | None = None
input_tokens: int = 0
output_tokens: int = 0
cache_read_input_tokens: int = 0
cache_creation_input_tokens: int = 0
total_cost: float = 0.0
input_cost: float = 0.0
output_cost: float = 0.0
def _coerce_usd(value: object) -> float:
if value is None or isinstance(value, bool):
return 0.0
if not isinstance(value, (int, float, str)):
return 0.0
try:
return max(0.0, float(value))
except (TypeError, ValueError):
return 0.0
def _absorb_usd(usage_or_root: dict) -> None:
nonlocal total_cost, input_cost, output_cost
cd = usage_or_root.get("cost_details")
if isinstance(cd, dict):
total_cost = max(
total_cost,
_coerce_usd(cd.get("total_cost")),
)
input_cost = max(
input_cost,
_coerce_usd(cd.get("input_cost")),
)
output_cost = max(
output_cost,
_coerce_usd(cd.get("output_cost")),
)
for field in ("total_cost", "cost"):
total_cost = max(
total_cost, _coerce_usd(usage_or_root.get(field))
)
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
@@ -1509,63 +1386,12 @@ class BaseUpstreamProvider:
output_tokens += usage.get(
"output_tokens", 0
)
# Anthropic's `message_start.usage`
# carries the cumulative cache
# snapshot for the prompt — pick
# the max() so subsequent
# `message_delta.usage` events
# (which only restate the same
# numbers) don't double-count.
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usd(usage)
if usage := data.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
"output_tokens", 0
)
cache_read_input_tokens = max(
cache_read_input_tokens,
int(
usage.get(
"cache_read_input_tokens", 0
)
or 0
),
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
int(
usage.get(
"cache_creation_input_tokens",
0,
)
or 0
),
)
_absorb_usd(usage)
# Some upstreams attach cost fields at
# the event root rather than nested
# under `usage`.
_absorb_usd(data)
except json.JSONDecodeError:
pass
modified_lines.append(line)
@@ -1580,23 +1406,9 @@ class BaseUpstreamProvider:
usage_data = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
}
messages_dispatch.embed_usd_costs(
usage_data,
total_cost,
input_cost,
output_cost,
)
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
if input_tokens > 0 or output_tokens > 0:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if fresh_key:
@@ -1833,7 +1645,6 @@ class BaseUpstreamProvider:
response_json["usage"], dict
):
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
self._fold_cache_into_input_tokens(response_json["usage"])
response_headers: dict[str, str] = {}
if cost_data:
@@ -1882,11 +1693,6 @@ class BaseUpstreamProvider:
last_model_seen: str | None = None
input_tokens = 0
output_tokens = 0
cache_read_input_tokens = 0
cache_creation_input_tokens = 0
total_cost = 0.0
input_cost = 0.0
output_cost = 0.0
async def finalize_without_usage() -> bytes | None:
nonlocal usage_finalized
@@ -1945,51 +1751,21 @@ class BaseUpstreamProvider:
# double-count.
input_tokens = max(input_tokens, annotated.input_tokens)
output_tokens = max(output_tokens, annotated.output_tokens)
cache_read_input_tokens = max(
cache_read_input_tokens,
annotated.cache_read_input_tokens,
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
annotated.cache_creation_input_tokens,
)
total_cost = max(total_cost, annotated.total_cost)
input_cost = max(input_cost, annotated.input_cost)
output_cost = max(output_cost, annotated.output_cost)
yield annotated.sse_bytes
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
if input_tokens > 0 or output_tokens > 0:
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
if fresh_key:
try:
rebuilt_usage: dict = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": (
cache_read_input_tokens
),
"cache_creation_input_tokens": (
cache_creation_input_tokens
),
}
messages_dispatch.embed_usd_costs(
rebuilt_usage,
total_cost,
input_cost,
output_cost,
)
combined_data: dict = {
"model": last_model_seen or "unknown",
"usage": rebuilt_usage,
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
},
}
cost_data = await adjust_payment_for_tokens(
fresh_key,
@@ -2052,11 +1828,6 @@ class BaseUpstreamProvider:
last_model_seen: str | None = None
input_tokens = 0
output_tokens = 0
cache_read_input_tokens = 0
cache_creation_input_tokens = 0
total_cost = 0.0
input_cost = 0.0
output_cost = 0.0
async for annotated in messages_dispatch.stream_annotated_events(
iterator, requested_model
@@ -2066,16 +1837,6 @@ class BaseUpstreamProvider:
# See _stream_litellm_messages for why this is max() not +=.
input_tokens = max(input_tokens, annotated.input_tokens)
output_tokens = max(output_tokens, annotated.output_tokens)
cache_read_input_tokens = max(
cache_read_input_tokens, annotated.cache_read_input_tokens
)
cache_creation_input_tokens = max(
cache_creation_input_tokens,
annotated.cache_creation_input_tokens,
)
total_cost = max(total_cost, annotated.total_cost)
input_cost = max(input_cost, annotated.input_cost)
output_cost = max(output_cost, annotated.output_cost)
buffered.append(annotated.sse_bytes)
response_headers: dict[str, str] = {
@@ -2083,13 +1844,7 @@ class BaseUpstreamProvider:
"Connection": "keep-alive",
}
if (
input_tokens == 0
and output_tokens == 0
and cache_read_input_tokens == 0
and cache_creation_input_tokens == 0
and total_cost == 0
):
if input_tokens == 0 and output_tokens == 0:
logger.warning(
"x-cashu /v1/messages stream finished with no usage data "
"— refund cannot be computed and the client effectively "
@@ -2103,25 +1858,13 @@ class BaseUpstreamProvider:
},
)
if (
input_tokens > 0
or output_tokens > 0
or cache_read_input_tokens > 0
or cache_creation_input_tokens > 0
or total_cost > 0
):
rebuilt_usage: dict = {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
}
messages_dispatch.embed_usd_costs(
rebuilt_usage, total_cost, input_cost, output_cost
)
if input_tokens > 0 or output_tokens > 0:
response_data: dict = {
"model": last_model_seen or "unknown",
"usage": rebuilt_usage,
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
},
}
try:
cost_data = await self.get_x_cashu_cost(
@@ -2197,12 +1940,6 @@ class BaseUpstreamProvider:
"""
path = self.normalize_request_path(path, model_obj)
if (
path.endswith("messages/count_tokens")
and not self.supports_anthropic_messages
):
return count_tokens_locally(request_body, model_obj)
if (
path.endswith("messages")
and not path.endswith("count_tokens")
@@ -3401,12 +3138,6 @@ class BaseUpstreamProvider:
request_body = await request.body()
if (
path.endswith("messages/count_tokens")
and not self.supports_anthropic_messages
):
return count_tokens_locally(request_body, model_obj)
if (
path.endswith("messages")
and not path.endswith("count_tokens")
@@ -4574,7 +4305,7 @@ 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.id: m for m in self._models_cache}
except Exception as e:
logger.error(
-108
View File
@@ -1,108 +0,0 @@
"""Local handling of Anthropic ``/v1/messages/count_tokens`` for upstreams
that do not natively expose the endpoint.
Most non-Anthropic upstreams (OpenAI-compat, Gemini OpenAI-compat,
OpenRouter chat-completions, generic providers) return 400/404 when asked
to ``POST /messages/count_tokens``. Claude Code and other Anthropic SDK
clients call this endpoint before each turn to size context windows and
trigger compaction, so a failure breaks the whole chat.
We answer locally. ``litellm.token_counter`` understands the Anthropic
message shape and the per-model tokenizers, so we prefer it. If it raises
(unknown model, encoding lookup failure, ...), we fall back to the
project's own ``estimate_tokens`` heuristic, which is always defined and
never raises.
"""
from __future__ import annotations
import json
from typing import Any
import litellm
from fastapi.responses import Response
from ..core import get_logger
from ..payment.helpers import estimate_tokens
from ..payment.models import Model
logger = get_logger(__name__)
def _parse_request_body(request_body: bytes | None) -> dict[str, Any]:
if not request_body:
return {}
try:
parsed = json.loads(request_body)
except (ValueError, TypeError):
return {}
return parsed if isinstance(parsed, dict) else {}
def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
messages = body.get("messages")
if not isinstance(messages, list):
messages = []
system = body.get("system")
if isinstance(system, str) and system:
messages = [{"role": "system", "content": system}, *messages]
elif isinstance(system, list):
text = "".join(
block.get("text", "")
for block in system
if isinstance(block, dict) and block.get("type") == "text"
)
if text:
messages = [{"role": "system", "content": text}, *messages]
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
return int(
litellm.token_counter(
model=model,
messages=messages,
tools=tools,
)
)
def count_tokens_locally(
request_body: bytes | None,
model_obj: Model | None,
) -> Response:
"""Return an Anthropic-compatible count_tokens response without
touching the upstream. Always returns 200; never raises."""
body = _parse_request_body(request_body)
model_name = ""
if model_obj is not None:
model_name = model_obj.forwarded_model_id or model_obj.id or ""
if not model_name:
body_model = body.get("model")
if isinstance(body_model, str):
model_name = body_model
input_tokens: int
try:
input_tokens = _count_with_litellm(model_name, body)
except Exception as exc:
messages = body.get("messages")
fallback_messages = messages if isinstance(messages, list) else []
input_tokens = estimate_tokens(fallback_messages)
logger.debug(
"litellm token_counter failed; using local estimator",
extra={
"model": model_name,
"error": str(exc),
"error_type": type(exc).__name__,
"estimated_tokens": input_tokens,
},
)
payload = {"input_tokens": max(0, int(input_tokens))}
return Response(
content=json.dumps(payload).encode(),
status_code=200,
media_type="application/json",
)
+6 -109
View File
@@ -248,28 +248,12 @@ class AnnotatedEvent(NamedTuple):
streaming paths in ``BaseUpstreamProvider`` consume ``sse_bytes`` plus
the tallies and only differ in whether they stream live or buffer
first.
``cache_read_input_tokens`` and ``cache_creation_input_tokens`` are
surfaced separately so the cost path can price them against the cache
rate rather than fold them silently into the regular input bucket.
``total_cost`` / ``input_cost`` / ``output_cost`` carry any
USD cost figures the upstream attached to this event (from
``usage.cost``, ``usage.total_cost``, or ``usage.cost_details``) so the
streaming paths can re-embed them in the rebuilt ``usage`` dict and let
``calculate_cost`` convert directly USD→sats instead of falling back to
token-based math.
"""
event: dict
sse_bytes: bytes
input_tokens: int
output_tokens: int
cache_read_input_tokens: int
cache_creation_input_tokens: int
total_cost: float
input_cost: float
output_cost: float
model: str | None
@@ -288,65 +272,21 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
in_tokens = 0
out_tokens = 0
cache_read_tokens = 0
cache_create_tokens = 0
total_cost = 0.0
input_cost = 0.0
output_cost = 0.0
model: str | None = None
def _coerce_float(value: object) -> float:
if value is None or isinstance(value, bool):
return 0.0
if not isinstance(value, (int, float, str)):
return 0.0
try:
return max(0.0, float(value))
except (TypeError, ValueError):
return 0.0
def _accumulate(usage: dict) -> None:
nonlocal in_tokens, out_tokens, cache_read_tokens, cache_create_tokens
nonlocal total_cost, input_cost, output_cost
in_tokens += int(usage.get("input_tokens") or 0)
out_tokens += int(usage.get("output_tokens") or 0)
cache_read_tokens += int(usage.get("cache_read_input_tokens") or 0)
cache_create_tokens += int(usage.get("cache_creation_input_tokens") or 0)
total_cost += _coerce_float(usage.get("total_cost"))
input_cost += _coerce_float(usage.get("input_cost"))
output_cost += _coerce_float(usage.get("output_cost"))
msg_for_meta = event.get("message")
if isinstance(msg_for_meta, dict):
if msg_for_meta.get("model"):
model = str(msg_for_meta["model"])
usage = msg_for_meta.get("usage")
if isinstance(usage, dict):
_accumulate(usage)
in_tokens += int(usage.get("input_tokens") or 0)
out_tokens += int(usage.get("output_tokens") or 0)
if isinstance(event.get("usage"), dict):
_accumulate(event["usage"])
# Some upstreams (notably OpenRouter-style proxies) attach cost fields
# directly at the event root rather than inside ``usage``.
for field in ("total_cost", "cost"):
total_cost = max(total_cost, _coerce_float(event.get(field)))
input_cost = max(input_cost, _coerce_float(event.get("input_cost")))
output_cost = max(output_cost, _coerce_float(event.get("output_cost")))
root_cost_details = event.get("cost_details")
if isinstance(root_cost_details, dict):
total_cost = max(
total_cost,
_coerce_float(root_cost_details.get("total_cost")),
)
input_cost = max(
input_cost,
_coerce_float(root_cost_details.get("input_cost")),
)
output_cost = max(
output_cost,
_coerce_float(root_cost_details.get("output_cost")),
)
usage = event["usage"]
in_tokens += int(usage.get("input_tokens") or 0)
out_tokens += int(usage.get("output_tokens") or 0)
event_type = str(event.get("type") or "")
payload = json.dumps(event)
@@ -355,18 +295,7 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
else:
sse_bytes = f"data: {payload}\n\n".encode()
return AnnotatedEvent(
event,
sse_bytes,
in_tokens,
out_tokens,
cache_read_tokens,
cache_create_tokens,
total_cost,
input_cost,
output_cost,
model,
)
return AnnotatedEvent(event, sse_bytes, in_tokens, out_tokens, model)
async def stream_annotated_events(
@@ -386,38 +315,6 @@ async def stream_annotated_events(
yield annotate_event(event, requested_model)
def embed_usd_costs(
usage: dict,
total_cost: float,
input_cost: float,
output_cost: float,
) -> None:
"""Mutate ``usage`` so ``calculate_cost`` will pick up the USD totals.
Mirrors the upstream shape: when any USD figure is present, attach
``cost`` (used by the simple-fallback branch in ``calculate_cost``) and
a ``cost_details`` block (used by the preferred branch — also gives the
input/output USD split when we have one).
"""
if total_cost <= 0 and input_cost <= 0 and output_cost <= 0:
return
cost_details: dict[str, float] = {}
effective_total = total_cost
if effective_total <= 0 and (input_cost > 0 or output_cost > 0):
effective_total = input_cost + output_cost
if effective_total > 0:
cost_details["total_cost"] = effective_total
usage["cost"] = effective_total
if input_cost > 0:
cost_details["input_cost"] = input_cost
if output_cost > 0:
cost_details["output_cost"] = output_cost
if cost_details:
usage["cost_details"] = cost_details
def compute_refund(amount: int, unit: str, cost_msats: int) -> int:
if unit == "msat":
return amount - cost_msats
+1 -1
View File
@@ -185,7 +185,7 @@ class OllamaUpstreamProvider(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.id: m for m in self._models_cache}
logger.info(
f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)},
-11
View File
@@ -18,10 +18,6 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
provider_type = "routstr"
default_base_url = None
platform_url = None
# Upstream Routstr nodes serve `/v1/messages` natively, so forward the
# request as-is instead of round-tripping through litellm's
# Anthropic→OpenAI translator.
supports_anthropic_messages = True
def __init__(
self,
@@ -47,13 +43,6 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
)
self.settings = provider_settings or {}
def normalize_request_path(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.
"""
return path.lstrip("/")
@classmethod
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
+5 -12
View File
@@ -48,8 +48,6 @@ async def recieve_token(
if token_obj.mint not in settings.cashu_mints:
return await swap_to_primary_mint(token_obj, wallet)
await wallet.load_mint(keyset_id=token_obj.keysets[0])
wallet.verify_proofs_dleq(token_obj.proofs)
await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
@@ -504,11 +502,11 @@ async def fetch_all_balances(
async def periodic_payout() -> None:
if not settings.receive_ln_address:
logger.warning("RECEIVE_LN_ADDRESS is not set, periodic payout disabled")
return
while True:
await asyncio.sleep(settings.payout_interval_seconds)
print(settings.payout_interval_seconds)
if not settings.receive_ln_address:
continue
await asyncio.sleep(60 * 15)
try:
async with db.create_session() as session:
for mint_url in settings.cashu_mints:
@@ -526,12 +524,7 @@ async def periodic_payout() -> None:
user_balance = user_balance // 1000
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
# Threshold is configured in sats; convert for msat wallets.
min_amount = (
settings.min_payout_sat
if unit == "sat"
else settings.min_payout_sat * 1000
)
min_amount = 210 if unit == "sat" else 210000
if available_balance > min_amount:
amount_received = await raw_send_to_lnurl(
wallet,
+8 -8
View File
@@ -372,17 +372,17 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
topup_amount_sat = 500
topup_token = await testmint_wallet.mint_tokens(topup_amount_sat)
key_looked_up = asyncio.Event()
validate_called = asyncio.Event()
allow_refund_to_continue = asyncio.Event()
original_lookup = balance_module._lookup_key_no_create
original_validate_bearer_key = balance_module.validate_bearer_key
delayed_once = False
async def delayed_lookup_key_no_create(*args: Any, **kwargs: Any) -> ApiKey | None:
async def delayed_validate_bearer_key(*args: Any, **kwargs: Any) -> ApiKey:
nonlocal delayed_once
key = await original_lookup(*args, **kwargs)
if not delayed_once and key is not None:
key = await original_validate_bearer_key(*args, **kwargs)
if not delayed_once:
delayed_once = True
key_looked_up.set()
validate_called.set()
await allow_refund_to_continue.wait()
return key
@@ -390,7 +390,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
return await authenticated_client.post("/v1/wallet/refund")
async def issue_topup() -> Any:
await key_looked_up.wait()
await validate_called.wait()
try:
return await authenticated_client.post(
"/v1/wallet/topup", params={"cashu_token": topup_token}
@@ -399,7 +399,7 @@ async def test_refund_rejects_concurrent_topup_on_same_key(
allow_refund_to_continue.set()
with patch(
"routstr.balance._lookup_key_no_create", new=delayed_lookup_key_no_create
"routstr.balance.validate_bearer_key", new=delayed_validate_bearer_key
):
refund_response, topup_response = await asyncio.gather(
issue_refund(), issue_topup()
+5 -51
View File
@@ -168,12 +168,12 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No
refund_token = "cashuArefund_apikey_token"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store,
@@ -203,12 +203,12 @@ async def test_apikey_refund_logs_token() -> None:
refund_token = "cashuAlogged_token"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
@@ -232,12 +232,12 @@ async def test_apikey_refund_log_includes_path() -> None:
refund_token = "cashuApath_token"
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
@@ -269,7 +269,6 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
key = _make_api_key(balance=5000, refund_currency="sat")
session = MagicMock()
session.get = AsyncMock(return_value=key)
# Debit returns rowcount=0 → balance changed concurrently
session.exec = AsyncMock(return_value=_update_result(0))
session.commit = AsyncMock()
@@ -277,6 +276,7 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", mock_send_token),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
@@ -333,11 +333,11 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
# First exec call = debit (succeeds), second = restore
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
@@ -355,49 +355,3 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
assert exc_info.value.status_code == 503
# Verify two exec calls: debit + restore
assert session.exec.await_count == 2
# ---------------------------------------------------------------------------
# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_refund_fresh_cashu_bearer_returns_401() -> None:
"""Fresh Cashu token not in DB must get 401, never create a new ApiKey."""
from fastapi import HTTPException
session = MagicMock()
session.get = AsyncMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer cashuAfresh_never_deposited_token",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 401
session.get.assert_awaited_once()
# No add/commit → no key was persisted
session.add.assert_not_called()
session.commit.assert_not_called()
@pytest.mark.asyncio
async def test_refund_unknown_sk_bearer_returns_401() -> None:
"""Unknown sk- key not in DB must get 401."""
from fastapi import HTTPException
session = MagicMock()
session.get = AsyncMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-unknownhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 401
session.get.assert_awaited_once()
-345
View File
@@ -1,345 +0,0 @@
"""Tests for cache token handling in cost calculation.
Covers OpenAI vs Anthropic caching formats, edge cases, and billing accuracy.
"""
import os
from unittest.mock import AsyncMock, patch
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
from routstr.core.settings import settings
from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost
@pytest.fixture
def mock_session() -> AsyncMock:
"""Mock AsyncSession for cost calculation tests."""
return AsyncMock()
@pytest.fixture(autouse=True)
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
"""Mock settings and price to use fixed pricing."""
monkeypatch.setattr(settings, "fixed_pricing", True)
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001)
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001)
@pytest.fixture(autouse=True)
def patch_sats_usd_price() -> None: # type: ignore[misc]
"""Patch sats_usd_price to avoid initialization issues."""
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
yield
# ============================================================================
# Test 1: OpenAI Cache Format
# ============================================================================
@pytest.mark.asyncio
async def test_openai_cache_subtraction(mock_session: AsyncMock) -> None:
"""OpenAI includes cached_tokens in prompt_tokens, subtract them."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 2000, # ← Includes 1000 cached
"completion_tokens": 100,
"prompt_tokens_details": {
"cached_tokens": 1000 # ← Extracted separately
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 1000 # 2000 - 1000
assert result.cache_read_input_tokens == 1000
assert result.output_tokens == 100
# ============================================================================
# Test 2: Anthropic Cache Format
# ============================================================================
@pytest.mark.asyncio
async def test_anthropic_cache_additive(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Anthropic cache tokens are separate (additive) from input_tokens."""
response = {
"model": "claude-3-5-sonnet",
"usage": {
"input_tokens": 500, # ← Regular input only
"output_tokens": 100,
"cache_creation_input_tokens": 1500, # ← Additive, not included above
"cache_read_input_tokens": 0,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 500
assert result.cache_creation_input_tokens == 1500
assert result.cache_read_input_tokens == 0
assert result.output_tokens == 100
# ============================================================================
# Test 3: Invalid Cache (Edge Case)
# ============================================================================
@pytest.mark.asyncio
async def test_cache_read_exceeds_prompt_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle buggy upstream reporting cached > prompt_tokens."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"prompt_tokens_details": {
"cached_tokens": 150 # ← Invalid! Greater than prompt
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
# Should not go negative
assert isinstance(result, CostData)
assert result.input_tokens == 0 # max(0, 100 - 150)
assert result.cache_read_input_tokens == 150
assert result.output_tokens == 50
# ============================================================================
# Test 4: Malformed Token Values
# ============================================================================
@pytest.mark.asyncio
async def test_malformed_cache_tokens_coerce_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle non-numeric cache token values."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"cache_read_input_tokens": "-50", # ← String, negative
"prompt_tokens_details": {
"cached_tokens": "invalid" # ← Non-numeric string
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
# Both should coerce to 0
assert isinstance(result, CostData)
assert result.cache_read_input_tokens == 0
assert result.input_tokens == 100 # No subtraction if cache_read = 0
# ============================================================================
# Test 5: Anthropic Cache Not Subtracted
# ============================================================================
@pytest.mark.asyncio
async def test_anthropic_cache_not_subtracted(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Anthropic cache fields should NOT be subtracted from input_tokens."""
response = {
"model": "claude-3-5-sonnet",
"usage": {
"input_tokens": 500,
"completion_tokens": 100,
"cache_read_input_tokens": 200, # ← Additive, don't subtract
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
# Anthropic: input_tokens stays as-is
assert isinstance(result, CostData)
assert result.input_tokens == 500 # NOT 300
assert result.cache_read_input_tokens == 200
# ============================================================================
# Test 6: Only Cache Read, No Regular Input
# ============================================================================
@pytest.mark.asyncio
async def test_only_cache_read_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle response with only cache read tokens."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 0,
"completion_tokens": 50,
"prompt_tokens_details": {
"cached_tokens": 1000
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 0 # max(0, 0 - 1000)
assert result.cache_read_input_tokens == 1000
assert result.output_tokens == 50
# ============================================================================
# Test 7: Only Cache Creation
# ============================================================================
@pytest.mark.asyncio
async def test_only_cache_creation_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle response with only cache creation tokens (Anthropic)."""
response = {
"model": "claude-3-5-sonnet",
"usage": {
"input_tokens": 500,
"output_tokens": 100,
"cache_creation_input_tokens": 2000,
"cache_read_input_tokens": 0,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 500
assert result.cache_creation_input_tokens == 2000
assert result.cache_read_input_tokens == 0
assert result.output_tokens == 100
# ============================================================================
# Test 8: Both Cache Read and Creation
# ============================================================================
@pytest.mark.asyncio
async def test_both_cache_read_and_creation(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle response with both cache read and creation."""
response = {
"model": "claude-3-5-sonnet",
"usage": {
"input_tokens": 300,
"output_tokens": 100,
"cache_creation_input_tokens": 2000,
"cache_read_input_tokens": 500,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 300
assert result.cache_creation_input_tokens == 2000
assert result.cache_read_input_tokens == 500
assert result.output_tokens == 100
# ============================================================================
# Test 9: Token Field Fallback
# ============================================================================
@pytest.mark.asyncio
async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Verify fallback order for token extraction."""
# When prompt_tokens is not present, fall back to input_tokens
response = {
"model": "gpt-4",
"usage": {
"input_tokens": 250,
"completion_tokens": 50,
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 250
assert result.output_tokens == 50
# ============================================================================
# Test 10: Float Token Values
# ============================================================================
@pytest.mark.asyncio
async def test_float_token_values_coerced_to_int(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle float token values by converting to int."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100.7, # Float
"completion_tokens": 50.3, # Float
"cache_read_input_tokens": 25.9, # Float
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.input_tokens == 100 # Floored
assert result.output_tokens == 50 # Floored
assert result.cache_read_input_tokens == 25 # Floored
# ============================================================================
# Test 11: Boolean Cache Tokens
# ============================================================================
@pytest.mark.asyncio
async def test_boolean_cache_tokens_coerced_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle boolean cache token values by coercing to zero."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"cache_read_input_tokens": True, # Boolean
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.cache_read_input_tokens == 0 # Boolean coerced to 0
assert result.input_tokens == 100 # No subtraction
# ============================================================================
# Test 12: Zero Cache Tokens
# ============================================================================
@pytest.mark.asyncio
async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""Handle explicit zero cache tokens."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"prompt_tokens_details": {
"cached_tokens": 0
}
}
}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, CostData)
assert result.cache_read_input_tokens == 0
assert result.input_tokens == 100
# ============================================================================
# Test 13: Missing Usage Block
# ============================================================================
@pytest.mark.asyncio
async def test_missing_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""When usage is missing, return MaxCostData with zero tokens."""
response = {"model": "gpt-4", "choices": [{"message": {"content": "test"}}]}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, MaxCostData)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 0
assert result.output_tokens == 0
# ============================================================================
# Test 14: Null Usage Block
# ============================================================================
@pytest.mark.asyncio
async def test_null_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
"""When usage is null, return MaxCostData with zero tokens."""
response = {"model": "gpt-4", "usage": None}
result = await calculate_cost(response, max_cost=100000, session=mock_session)
assert isinstance(result, MaxCostData)
assert result.input_tokens == 0
assert result.cache_read_input_tokens == 0
-177
View File
@@ -1,177 +0,0 @@
"""Unit tests for the local count_tokens shim.
The shim runs whenever an upstream that does not support Anthropic's
``/v1/messages`` endpoint is asked for a token count. It must always
return a 200 JSON ``{"input_tokens": N}`` response and must never raise.
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import patch
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream import count_tokens as count_tokens_module
from routstr.upstream.count_tokens import count_tokens_locally
def _make_model(model_id: str = "anthropic/claude-3-5-sonnet") -> Model:
pricing = Pricing(prompt=0.000003, completion=0.000015)
architecture = Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="cl100k_base",
instruct_type=None,
)
return Model(
id=model_id,
name=model_id,
created=0,
description="",
context_length=200_000,
architecture=architecture,
pricing=pricing,
)
def _body(payload: dict[str, Any]) -> bytes:
return json.dumps(payload).encode()
def _read_payload(response: Any) -> dict[str, Any]:
body = response.body if isinstance(response.body, bytes) else bytes(response.body)
return json.loads(body.decode())
def test_returns_input_tokens_for_simple_messages() -> None:
model = _make_model()
request_body = _body(
{
"model": model.id,
"messages": [{"role": "user", "content": "hello world"}],
}
)
response = count_tokens_locally(request_body, model)
assert response.status_code == 200
assert response.media_type == "application/json"
payload = _read_payload(response)
assert "input_tokens" in payload
assert isinstance(payload["input_tokens"], int)
assert payload["input_tokens"] >= 0
def test_falls_back_to_estimator_when_litellm_raises() -> None:
model = _make_model()
request_body = _body(
{
"model": model.id,
"messages": [{"role": "user", "content": "this is a longer message"}],
}
)
with patch.object(
count_tokens_module,
"_count_with_litellm",
side_effect=RuntimeError("boom"),
):
response = count_tokens_locally(request_body, model)
assert response.status_code == 200
payload = _read_payload(response)
assert payload["input_tokens"] >= 1
def test_handles_missing_model_object() -> None:
request_body = _body(
{
"model": "anthropic/claude-3-5-sonnet",
"messages": [{"role": "user", "content": "hi"}],
}
)
response = count_tokens_locally(request_body, None)
assert response.status_code == 200
payload = _read_payload(response)
assert payload["input_tokens"] >= 0
def test_handles_empty_request_body() -> None:
response = count_tokens_locally(b"", _make_model())
assert response.status_code == 200
payload = _read_payload(response)
assert payload["input_tokens"] >= 0
def test_handles_malformed_json() -> None:
response = count_tokens_locally(b"not-json", _make_model())
assert response.status_code == 200
payload = _read_payload(response)
assert payload["input_tokens"] >= 0
def test_includes_system_prompt_in_count() -> None:
model = _make_model()
short = _body(
{
"model": model.id,
"messages": [{"role": "user", "content": "hi"}],
}
)
with_system = _body(
{
"model": model.id,
"system": "You are a helpful assistant with a long preamble " * 10,
"messages": [{"role": "user", "content": "hi"}],
}
)
short_count = _read_payload(count_tokens_locally(short, model))["input_tokens"]
long_count = _read_payload(count_tokens_locally(with_system, model))["input_tokens"]
assert long_count > short_count
def test_supports_anthropic_system_block_list() -> None:
model = _make_model()
request_body = _body(
{
"model": model.id,
"system": [{"type": "text", "text": "be terse" * 50}],
"messages": [{"role": "user", "content": "ok"}],
}
)
response = count_tokens_locally(request_body, model)
payload = _read_payload(response)
assert payload["input_tokens"] > 0
def test_uses_forwarded_model_id_when_present() -> None:
model = _make_model("anthropic/claude-3-5-sonnet")
model.forwarded_model_id = "claude-3-5-sonnet-20241022"
request_body = _body(
{
"model": "ignored",
"messages": [{"role": "user", "content": "hi"}],
}
)
captured: dict[str, Any] = {}
def _capture(model_name: str, body: dict[str, Any]) -> int:
captured["model"] = model_name
return 7
with patch.object(count_tokens_module, "_count_with_litellm", side_effect=_capture):
response = count_tokens_locally(request_body, model)
assert captured["model"] == "claude-3-5-sonnet-20241022"
assert _read_payload(response)["input_tokens"] == 7
@@ -1,333 +0,0 @@
"""Tests for cost accumulation in streaming message dispatch.
Verifies that costs are correctly summed across multiple streaming events,
not taking only the maximum.
"""
import os
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.upstream.messages_dispatch import annotate_event
# ============================================================================
# Test 1: Cost Accumulation via AnnotatedEvent
# ============================================================================
@pytest.mark.unit
def test_annotate_event_extracts_costs() -> None:
"""Each event should report its own costs."""
event = {
"type": "content_block_start",
"usage": {"total_cost": 0.010, "input_cost": 0.008, "output_cost": 0.002}
}
result = annotate_event(event, None)
assert result.total_cost == 0.010
assert result.input_cost == 0.008
assert result.output_cost == 0.002
# ============================================================================
# Test 2: Multiple Events with Incremental Costs
# ============================================================================
@pytest.mark.unit
def test_multiple_events_sum_costs() -> None:
"""When processing multiple events, costs should accumulate (not max)."""
# Event 1: initial costs
event1 = {
"type": "message_start",
"usage": {"input_tokens": 100, "total_cost": 0.010}
}
result1 = annotate_event(event1, None)
# Event 2: additional costs
event2 = {
"type": "content_block_start",
"usage": {"output_tokens": 50, "total_cost": 0.005}
}
result2 = annotate_event(event2, None)
# Event 3: more costs
event3 = {
"type": "message_delta",
"usage": {"output_tokens": 25, "total_cost": 0.008}
}
result3 = annotate_event(event3, None)
# Each event should report its own cost (before accumulation)
assert result1.total_cost == 0.010
assert result2.total_cost == 0.005
assert result3.total_cost == 0.008
# When summed for billing: 0.010 + 0.005 + 0.008 = 0.023
# This verifies the cost accumulation fix (using += instead of max())
total_from_events = result1.total_cost + result2.total_cost + result3.total_cost
assert total_from_events == 0.023
# ============================================================================
# Test 3: Token Accumulation Consistency
# ============================================================================
@pytest.mark.unit
def test_tokens_and_costs_extracted_independently() -> None:
"""Tokens and costs should be extracted independently per event."""
event = {
"type": "content_block_delta",
"usage": {
"input_tokens": 100,
"output_tokens": 50,
"cache_read_input_tokens": 30,
"total_cost": 0.007,
"input_cost": 0.005,
"output_cost": 0.002
}
}
result = annotate_event(event, None)
# All values should be extracted
assert result.input_tokens == 100
assert result.output_tokens == 50
assert result.cache_read_input_tokens == 30
assert result.total_cost == 0.007
assert result.input_cost == 0.005
assert result.output_cost == 0.002
# ============================================================================
# Test 4: Cost in Message vs Root
# ============================================================================
@pytest.mark.unit
def test_cost_extracted_from_message_usage() -> None:
"""Costs in message.usage should be extracted correctly."""
event = {
"type": "message_start",
"message": {
"usage": {
"input_tokens": 100,
"total_cost": 0.010,
"input_cost": 0.008,
"output_cost": 0.002
}
}
}
result = annotate_event(event, None)
assert result.input_tokens == 100
assert result.total_cost == 0.010
assert result.input_cost == 0.008
assert result.output_cost == 0.002
# ============================================================================
# Test 5: Cost at Event Root Level
# ============================================================================
@pytest.mark.unit
def test_cost_extracted_from_event_root() -> None:
"""Costs at event root should be extracted (OpenRouter style)."""
event = {
"type": "content_block_delta",
"usage": {"output_tokens": 25},
"total_cost": 0.005, # ← At root level
"cost": 0.005
}
result = annotate_event(event, None)
assert result.output_tokens == 25
assert result.total_cost == 0.005
# ============================================================================
# Test 6: Cost Details in Event Root
# ============================================================================
@pytest.mark.unit
def test_cost_details_extracted_from_event_root() -> None:
"""cost_details at event root should be extracted correctly."""
event = {
"type": "message_delta",
"cost_details": {
"total_cost": 0.015,
"input_cost": 0.010,
"output_cost": 0.005
}
}
result = annotate_event(event, None)
assert result.total_cost == 0.015
assert result.input_cost == 0.010
assert result.output_cost == 0.005
# ============================================================================
# Test 7: No Duplicated Dict Lookups
# ============================================================================
@pytest.mark.unit
def test_annotate_event_no_duplicate_lookups() -> None:
"""Verify that dict lookups are not duplicated (fix for copy-paste error)."""
# This is tested implicitly through proper extraction
event = {
"type": "message_delta",
"cost_details": {
"total_cost": 0.020,
"input_cost": 0.015,
"output_cost": 0.005
}
}
result = annotate_event(event, None)
# Should extract each field exactly once, correctly
assert result.total_cost == 0.020
assert result.input_cost == 0.015
assert result.output_cost == 0.005
# ============================================================================
# Test 8: Model Name Extraction
# ============================================================================
@pytest.mark.unit
def test_model_extracted_from_event() -> None:
"""Model name should be extracted from event."""
event = {
"type": "message_start",
"message": {"model": "claude-3-5-sonnet"}
}
result = annotate_event(event, None)
assert result.model == "claude-3-5-sonnet"
# ============================================================================
# Test 9: Model Name Override
# ============================================================================
@pytest.mark.unit
def test_model_name_override() -> None:
"""Requested model should override actual model in event."""
event = {
"type": "message_start",
"message": {"model": "actual-model"}
}
# Override with requested_model
result = annotate_event(event, requested_model="alias-model")
# Event should be modified to use requested_model
assert event["message"]["model"] == "alias-model" # type: ignore[index]
assert result.model == "alias-model"
# ============================================================================
# Test 10: SSE Encoding
# ============================================================================
@pytest.mark.unit
def test_sse_bytes_encoded() -> None:
"""Event should be encoded as SSE bytes."""
event = {
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""}
}
result = annotate_event(event, None)
assert result.sse_bytes is not None
assert isinstance(result.sse_bytes, bytes)
assert b"event: content_block_start" in result.sse_bytes or b"data:" in result.sse_bytes
# ============================================================================
# Test 11: Missing Cost Fields Default to Zero
# ============================================================================
@pytest.mark.unit
def test_missing_cost_fields_default_to_zero() -> None:
"""When cost fields are missing, should default to 0.0."""
event = {
"type": "content_block_delta",
"usage": {"output_tokens": 25}
# No cost fields
}
result = annotate_event(event, None)
assert result.total_cost == 0.0
assert result.input_cost == 0.0
assert result.output_cost == 0.0
# ============================================================================
# Test 12: Cache Tokens in Events
# ============================================================================
@pytest.mark.unit
def test_cache_tokens_extracted_from_event() -> None:
"""Cache tokens should be extracted from event usage."""
event = {
"type": "message_delta",
"usage": {
"input_tokens": 100,
"output_tokens": 50,
"cache_read_input_tokens": 200,
"cache_creation_input_tokens": 0
}
}
result = annotate_event(event, None)
assert result.input_tokens == 100
assert result.output_tokens == 50
assert result.cache_read_input_tokens == 200
assert result.cache_creation_input_tokens == 0
# ============================================================================
# Test 13: Malformed Cost Values
# ============================================================================
@pytest.mark.unit
def test_malformed_cost_values_coerced() -> None:
"""Malformed cost values should be coerced to 0.0."""
event = {
"type": "message_delta",
"usage": {
"input_tokens": 100,
"output_tokens": 50,
"total_cost": "invalid", # ← Non-numeric
}
}
result = annotate_event(event, None)
# Invalid cost should default to 0.0
assert result.total_cost == 0.0
# ============================================================================
# Test 14: Negative Cost Values
# ============================================================================
@pytest.mark.unit
def test_negative_cost_values_clamped() -> None:
"""Negative cost values should be clamped to 0.0."""
event = {
"type": "message_delta",
"usage": {
"input_tokens": 100,
"output_tokens": 50,
"total_cost": -0.05 # ← Negative
}
}
result = annotate_event(event, None)
# Negative cost should be clamped to 0.0
assert result.total_cost == 0.0
# ============================================================================
# Test 15: Streaming Event Type Preserved
# ============================================================================
@pytest.mark.unit
def test_event_type_preserved_in_sse() -> None:
"""Event type should be preserved in SSE encoding."""
event_types = ["message_start", "content_block_start", "content_block_delta", "message_delta"]
for event_type in event_types:
event = {"type": event_type}
result = annotate_event(event, None)
assert result.event["type"] == event_type # type: ignore[index]
# SSE should include the event type line if present
if event_type:
assert f"event: {event_type}".encode() in result.sse_bytes
+25 -37
View File
@@ -908,7 +908,7 @@ async def test_forward_request_skips_litellm_when_provider_supports_messages() -
@pytest.mark.asyncio
async def test_forward_request_handles_count_tokens_locally() -> None:
async def test_forward_request_skips_litellm_for_count_tokens() -> None:
provider = _make_provider()
key = _make_key()
model = _make_model()
@@ -923,25 +923,19 @@ async def test_forward_request_handles_count_tokens_locally() -> None:
with patch.object(
provider,
"prepare_request_body",
side_effect=AssertionError("upstream should not be called"),
side_effect=RuntimeError("stop here"),
):
response = await provider.forward_request(
request=request,
path="messages/count_tokens",
headers={},
request_body=_anthropic_request_body(),
key=key,
max_cost_for_model=10_000,
session=session,
model_obj=model,
)
assert response.status_code == 200
body = response.body if isinstance(response.body, bytes) else bytes(response.body)
payload = json.loads(body.decode())
assert "input_tokens" in payload
assert isinstance(payload["input_tokens"], int)
assert payload["input_tokens"] >= 0
with pytest.raises(RuntimeError, match="stop here"):
await provider.forward_request(
request=request,
path="messages/count_tokens",
headers={},
request_body=_anthropic_request_body(),
key=key,
max_cost_for_model=10_000,
session=session,
model_obj=model,
)
# ---------------------------------------------------------------------------
@@ -1016,7 +1010,7 @@ async def test_forward_x_cashu_request_skips_litellm_when_native_messages() -> N
@pytest.mark.asyncio
async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None:
async def test_forward_x_cashu_request_skips_litellm_for_count_tokens() -> None:
provider = _make_provider()
model = _make_model()
request = _make_request()
@@ -1030,24 +1024,18 @@ async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None:
with patch.object(
provider,
"prepare_request_body",
side_effect=AssertionError("upstream should not be called"),
side_effect=RuntimeError("stop here"),
):
response = await provider.forward_x_cashu_request(
request=request,
path="v1/messages/count_tokens",
headers={},
amount=5_000,
unit="sat",
max_cost_for_model=10_000,
model_obj=model,
mint="https://mint",
payment_token_hash="h",
)
assert response.status_code == 200
body = response.body if isinstance(response.body, bytes) else bytes(response.body)
payload = json.loads(body.decode())
assert "input_tokens" in payload
with pytest.raises(RuntimeError, match="stop here"):
await provider.forward_x_cashu_request(
request=request,
path="v1/messages/count_tokens",
headers={},
amount=5_000,
unit="sat",
max_cost_for_model=10_000,
model_obj=model,
)
# ---------------------------------------------------------------------------
-64
View File
@@ -1,64 +0,0 @@
"""Tests for the built-in 404 handler in routstr.proxy."""
from __future__ import annotations
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from routstr import proxy
from routstr.proxy import proxy_router
def _make_app() -> FastAPI:
app = FastAPI()
app.include_router(proxy_router)
return app
@pytest.mark.skipif(
proxy._NOT_FOUND_HTML is None,
reason="UI bundle (ui_out/404.html) not present in this environment",
)
def test_unknown_path_returns_html_404_for_browser() -> None:
client = TestClient(_make_app())
response = client.get("/some/random/page", headers={"accept": "text/html"})
assert response.status_code == 404
assert response.headers["content-type"].startswith("text/html")
assert "404" in response.text
def test_unknown_path_returns_json_404_for_api_client() -> None:
client = TestClient(_make_app())
response = client.get(
"/some/random/page", headers={"accept": "application/json"}
)
assert response.status_code == 404
assert response.headers["content-type"].startswith("application/json")
payload = response.json()
assert payload["error"]["type"] == "not_found"
assert payload["error"]["code"] == 404
assert "/some/random/page" in payload["error"]["message"]
def test_root_path_returns_404_for_proxy_router() -> None:
client = TestClient(_make_app())
response = client.get("/", headers={"accept": "application/json"})
assert response.status_code == 404
def test_v1_path_is_not_intercepted_by_404_handler() -> None:
"""Paths starting with v1/ must reach the proxy logic, not the 404 handler."""
client = TestClient(_make_app(), raise_server_exceptions=False)
response = client.get("/v1/anything")
if response.status_code == 404:
# Any 404 here must come from inner proxy logic, not our HTML page.
assert "<!DOCTYPE html>" not in response.text
def test_json_returned_when_ui_html_missing(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(proxy, "_NOT_FOUND_HTML", None)
client = TestClient(_make_app())
response = client.get("/some/random/page", headers={"accept": "text/html"})
assert response.status_code == 404
assert response.headers["content-type"].startswith("application/json")
+1 -50
View File
@@ -1,12 +1,11 @@
import os
import pytest
from pydantic.v1 import ValidationError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import text
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.settings import Settings, SettingsService
from routstr.core.settings import SettingsService
@pytest.mark.asyncio
@@ -47,54 +46,6 @@ async def test_settings_db_precedence_over_env() -> None:
assert again.enable_analytics_sharing is False
def test_payout_settings_have_sensible_defaults() -> None:
s = Settings()
assert s.min_payout_sat == 210
assert s.payout_interval_seconds == 900
@pytest.mark.parametrize(
"field,bad_value",
[
("min_payout_sat", 0),
("min_payout_sat", -1),
("payout_interval_seconds", 0),
("payout_interval_seconds", -10),
],
)
def test_payout_settings_reject_invalid_values(field: str, bad_value: int) -> None:
kwargs: dict[str, object] = {field: bad_value}
with pytest.raises(ValidationError):
Settings(**kwargs) # type: ignore[arg-type]
def test_payout_settings_accept_custom_positive_values() -> None:
s = Settings(min_payout_sat=500, payout_interval_seconds=60)
assert s.min_payout_sat == 500
assert s.payout_interval_seconds == 60
@pytest.mark.asyncio
async def test_payout_settings_persist_via_settings_service() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
await SettingsService.initialize(session)
updated = await SettingsService.update(
{"min_payout_sat": 1000, "payout_interval_seconds": 300}, session
)
assert updated.min_payout_sat == 1000
assert updated.payout_interval_seconds == 300
@pytest.mark.asyncio
async def test_payout_settings_update_rejects_invalid() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
await SettingsService.initialize(session)
with pytest.raises(ValidationError):
await SettingsService.update({"min_payout_sat": 0}, session)
@pytest.mark.asyncio
async def test_settings_initialize_discards_unknown_keys() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
-150
View File
@@ -1,150 +0,0 @@
"""Tests for ``BaseUpstreamProvider.forward_upstream_error_response``.
Upstream services (e.g. an Express server that doesn't expose ``/messages``)
sometimes return a non-JSON error body. The proxy must surface those errors
in a consistent JSON envelope so clients don't have to parse HTML.
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import Mock
import httpx
import pytest
from routstr.upstream.base import BaseUpstreamProvider, _is_json_content_type
def _make_request(request_id: str = "req-123") -> Mock:
request = Mock(spec=["method", "state"])
request.method = "POST"
request.state = Mock()
request.state.request_id = request_id
return request
def _make_upstream_response(
*,
body: bytes,
status_code: int = 404,
content_type: str | None = "text/html",
extra_headers: dict[str, str] | None = None,
) -> httpx.Response:
headers: dict[str, str] = {}
if content_type is not None:
headers["content-type"] = content_type
if extra_headers:
headers.update(extra_headers)
return httpx.Response(status_code=status_code, headers=headers, content=body)
@pytest.fixture
def provider() -> BaseUpstreamProvider:
return BaseUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="k", provider_fee=1.0
)
@pytest.mark.parametrize(
"content_type,expected",
[
("application/json", True),
("application/json; charset=utf-8", True),
("text/json", True),
("application/problem+json", True),
("application/vnd.api+json", True),
("text/html", False),
("text/html; charset=utf-8", False),
("text/plain", False),
("", False),
(None, False),
],
)
def test_is_json_content_type(content_type: str | None, expected: bool) -> None:
assert _is_json_content_type(content_type) is expected
@pytest.mark.asyncio
async def test_html_error_is_normalized_to_json_envelope(
provider: BaseUpstreamProvider,
) -> None:
html_body = (
b"<!DOCTYPE html><html><head><title>Error</title></head>"
b"<body><pre>Cannot POST /messages</pre></body></html>"
)
upstream = _make_upstream_response(body=html_body, status_code=404)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 404
assert response.media_type == "application/json"
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["upstream_status"] == 404
assert payload["error"]["upstream_content_type"] == "text/html"
assert "Cannot POST /messages" in payload["error"]["upstream_body_preview"]
assert payload["request_id"] == "req-123"
# The upstream's text/html content-type must not survive — Response()
# sets the JSON content-type for us via media_type.
assert response.headers["content-type"].startswith("application/json")
@pytest.mark.asyncio
async def test_plain_text_error_is_normalized(
provider: BaseUpstreamProvider,
) -> None:
upstream = _make_upstream_response(
body=b"Service Unavailable", status_code=503, content_type="text/plain"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 503
assert response.media_type == "application/json"
payload = json.loads(bytes(response.body))
assert payload["error"]["message"] == "Service Unavailable"
@pytest.mark.asyncio
async def test_empty_body_with_non_json_content_type_normalizes(
provider: BaseUpstreamProvider,
) -> None:
upstream = _make_upstream_response(
body=b"", status_code=502, content_type="text/html"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 502
assert response.media_type == "application/json"
payload = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["upstream_body_preview"] is None
@pytest.mark.asyncio
async def test_json_error_body_is_passed_through_unchanged(
provider: BaseUpstreamProvider,
) -> None:
json_body = json.dumps(
{"error": {"message": "Invalid model", "type": "invalid_request_error"}}
).encode()
upstream = _make_upstream_response(
body=json_body, status_code=400, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 400
assert bytes(response.body) == json_body
assert response.media_type == "application/json"
-36
View File
@@ -74,39 +74,3 @@ async def test_get_balance_returns_none_on_connect_timeout(
balance = await provider.get_balance()
assert balance is None
def test_normalize_request_path_keeps_v1_prefix() -> None:
"""Routstr upstream stores ``base_url`` without ``/v1``; the prefix
must stay on the path so ``build_request_url`` produces ``/v1/<endpoint>``
instead of ``/<endpoint>`` (which the upstream Routstr 404s with HTML)."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
assert provider.normalize_request_path("v1/messages") == "v1/messages"
assert provider.normalize_request_path("/v1/messages") == "v1/messages"
assert (
provider.normalize_request_path("v1/chat/completions")
== "v1/chat/completions"
)
def test_build_request_url_for_v1_messages() -> None:
"""Forwarding ``/v1/messages`` must hit the upstream's ``/v1/messages``."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
normalized = provider.normalize_request_path("v1/messages")
assert (
provider.build_request_url(normalized)
== "https://privateprovider.xyz/v1/messages"
)
def test_supports_anthropic_messages_natively() -> None:
"""Routstr nodes serve ``/v1/messages`` directly, so the proxy must
forward as-is instead of round-tripping through litellm."""
assert RoutstrUpstreamProvider.supports_anthropic_messages is True
-130
View File
@@ -32,18 +32,9 @@ interface SettingsData {
onion_url?: string;
cashu_mints?: string[];
relays?: string[];
receive_ln_address?: string;
min_payout_sat?: number;
payout_interval_seconds?: number;
[key: string]: unknown;
}
const PAYOUT_KEYS = [
'receive_ln_address',
'min_payout_sat',
'payout_interval_seconds',
] as const;
const HANDLED_KEYS = [
'name',
'description',
@@ -57,7 +48,6 @@ const HANDLED_KEYS = [
'admin_password',
'id',
'updated_at',
...PAYOUT_KEYS,
];
const IGNORED_KEYS = [
@@ -379,7 +369,6 @@ export function AdminSettings() {
const cashuMintsChanged = hasFieldChanged('cashu_mints');
const relaysChanged = hasFieldChanged('relays');
const analyticsSharingChanged = hasFieldChanged('enable_analytics_sharing');
const payoutChanged = PAYOUT_KEYS.some(hasFieldChanged);
const advancedKeys = Object.keys(settings).filter(
(key) => !HANDLED_KEYS.includes(key) && !IGNORED_KEYS.includes(key)
);
@@ -412,44 +401,8 @@ export function AdminSettings() {
setNewRelay('');
};
const resetAnalyticsSharing = () => resetFields(['enable_analytics_sharing']);
const resetPayout = () => resetFields([...PAYOUT_KEYS]);
const resetAdvanced = () => resetFields(advancedKeys);
const payoutFields: ReadonlyArray<{
key: (typeof PAYOUT_KEYS)[number];
label: string;
placeholder: string;
type: 'text' | 'number';
helpText: string;
min?: number;
}> = [
{
key: 'receive_ln_address',
label: 'Lightning Receive Address',
placeholder: 'you@walletofsatoshi.com or LNURL',
type: 'text',
helpText:
'Lightning address (or LNURL) profits are paid out to. Leave empty to disable periodic payouts.',
},
{
key: 'min_payout_sat',
label: 'Minimum Payout (sat)',
placeholder: '210',
type: 'number',
min: 1,
helpText:
'Wallet payouts only fire when at least this many satoshis are available. Must be > 0.',
},
{
key: 'payout_interval_seconds',
label: 'Payout Interval (seconds)',
placeholder: '900',
type: 'number',
min: 1,
helpText: 'How often the payout loop wakes up to check balances.',
},
];
if (loading) {
return (
<div className='space-y-4'>
@@ -667,89 +620,6 @@ export function AdminSettings() {
) : null}
</Card>
{/* Lightning Payout Settings */}
<Card>
<CardHeader>
<CardTitle>Lightning Payout Settings</CardTitle>
<CardDescription>
Tune how node profit is paid out over Lightning. Amounts must be
positive and above your wallet&apos;s minimum-invoice constraints.
</CardDescription>
</CardHeader>
<CardContent className='space-y-4'>
{payoutFields.map((field) => {
const value = settings[field.key];
if (field.type === 'number') {
return (
<div key={field.key} className='space-y-2'>
<Label htmlFor={field.key}>{field.label}</Label>
<Input
id={field.key}
type='number'
min={field.min}
value={
typeof value === 'number'
? value
: value === undefined || value === null
? ''
: Number(value)
}
placeholder={field.placeholder}
onChange={(e) => {
const raw = e.target.value;
if (raw === '') {
handleInputChange(field.key, undefined);
} else {
const parsed = Number(raw);
handleInputChange(
field.key,
Number.isFinite(parsed) ? parsed : undefined
);
}
}}
/>
<p className='text-muted-foreground text-xs'>
{field.helpText}
</p>
</div>
);
}
return (
<div key={field.key} className='space-y-2'>
<Label htmlFor={field.key}>{field.label}</Label>
<Input
id={field.key}
value={(value as string) || ''}
placeholder={field.placeholder}
onChange={(e) =>
handleInputChange(field.key, e.target.value)
}
/>
<p className='text-muted-foreground text-xs'>
{field.helpText}
</p>
</div>
);
})}
</CardContent>
{payoutChanged ? (
<CardFooter className='justify-start'>
<div className='flex w-full flex-col gap-2 sm:w-auto sm:flex-row sm:items-center'>
<Button
variant='outline'
onClick={resetPayout}
disabled={loading || saving}
>
Cancel
</Button>
<Button onClick={handleSave} disabled={loading || saving}>
{saving ? 'Saving...' : 'Save'}
</Button>
</div>
</CardFooter>
) : null}
</Card>
{/* Relays */}
<Card>
<CardHeader>
-3
View File
@@ -6,9 +6,6 @@ const nextConfig: NextConfig = {
images: {
unoptimized: true,
},
turbopack: {
root: __dirname,
},
};
export default nextConfig;
-7
View File
@@ -1,6 +1,5 @@
{
"name": "routstr-service",
"packageManager": "pnpm@10.15.0",
"version": "0.1.0",
"private": true,
"scripts": {
@@ -89,11 +88,5 @@
"prettier-plugin-tailwindcss": "^0.7.2",
"tailwindcss": "^4.2.0",
"typescript": "^5.9.3"
},
"pnpm": {
"onlyBuiltDependencies": [
"sharp",
"unrs-resolver"
]
}
}