revert: remove unrelated repository formatting

This commit is contained in:
9qeklajc
2026-07-26 20:12:12 +02:00
parent 06dba681c5
commit c5da73f1e9
22 changed files with 319 additions and 259 deletions
+10 -10
View File
@@ -232,10 +232,7 @@ def create_model_mappings(
aliases.append(prefixed_id) aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias # Register forwarded_model_id as a routable alias
if ( if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
model_to_use.forwarded_model_id
and model_to_use.forwarded_model_id not in aliases
):
aliases.append(model_to_use.forwarded_model_id) aliases.append(model_to_use.forwarded_model_id)
# Try to set each alias # Try to set each alias
@@ -325,10 +322,7 @@ def create_model_mappings(
aliases.append(prefixed_id) aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias # Register forwarded_model_id as a routable alias
if ( if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
model_to_use.forwarded_model_id
and model_to_use.forwarded_model_id not in aliases
):
aliases.append(model_to_use.forwarded_model_id) aliases.append(model_to_use.forwarded_model_id)
for alias in aliases: for alias in aliases:
@@ -348,10 +342,16 @@ def create_model_mappings(
forwarded_model_ids, the one whose forwarded_model_id equals the forwarded_model_ids, the one whose forwarded_model_id equals the
requested alias wins. requested alias wins.
""" """
if model.forwarded_model_id and model.forwarded_model_id.lower() == alias: if (
model.forwarded_model_id
and model.forwarded_model_id.lower() == alias
):
return 5 return 5
if model.id and model.id.lower() == alias: if (
model.id
and model.id.lower() == alias
):
return 4 return 4
model_base = get_base_model_id(model.id) model_base = get_base_model_id(model.id)
+5 -24
View File
@@ -260,11 +260,7 @@ async def _lookup_key_no_create(
async def _restore_balance( async def _restore_balance(
session: AsyncSession, session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
hashed_key: str,
balance: int,
reserved_balance: int,
mint_url: str,
) -> None: ) -> None:
"""Restore balance after a failed refund mint attempt.""" """Restore balance after a failed refund mint attempt."""
restore_stmt = ( restore_stmt = (
@@ -279,11 +275,7 @@ async def _restore_balance(
await session.commit() await session.commit()
logger.info( logger.info(
"refund_wallet_endpoint: balance restored after mint failure", "refund_wallet_endpoint: balance restored after mint failure",
extra={ extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
"hashed_key": hashed_key,
"restored_balance": balance,
"mint_url": mint_url,
},
) )
@@ -468,23 +460,11 @@ async def refund_wallet_endpoint(
except HTTPException: except HTTPException:
# Minting failed — restore the debited balance # Minting failed — restore the debited balance
await _restore_balance( await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
session,
key.hashed_key,
pre_debit_balance,
pre_debit_reserved,
key.refund_mint_url or "",
)
raise raise
except Exception as e: except Exception as e:
# Minting failed — restore the debited balance # Minting failed — restore the debited balance
await _restore_balance( await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
session,
key.hashed_key,
pre_debit_balance,
pre_debit_reserved,
key.refund_mint_url or "",
)
error_msg = str(e) error_msg = str(e)
logger.error( logger.error(
"refund_wallet_endpoint: mint/send failed", "refund_wallet_endpoint: mint/send failed",
@@ -705,6 +685,7 @@ async def reset_child_key_spent(
return {"success": True, "message": "Child key balance reset successfully."} return {"success": True, "message": "Child key balance reset successfully."}
@router.api_route( @router.api_route(
"/{path:path}", "/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"], methods=["GET", "POST", "PUT", "DELETE"],
+14 -12
View File
@@ -68,9 +68,7 @@ async def require_admin_api(request: Request) -> None:
async with create_session() as session: async with create_session() as session:
result = await session.exec(select(CliToken).where(CliToken.token == token)) result = await session.exec(select(CliToken).where(CliToken.token == token))
cli_token = result.first() cli_token = result.first()
if cli_token and ( if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
cli_token.expires_at is None or cli_token.expires_at > now_ts
):
cli_token.last_used_at = now_ts cli_token.last_used_at = now_ts
session.add(cli_token) session.add(cli_token)
await session.commit() await session.commit()
@@ -257,12 +255,16 @@ async def update_password(request: Request, password_update: PasswordUpdate) ->
secret = await get_secret(session) secret = await get_secret(session)
if not secret.admin_password_hash: if not secret.admin_password_hash:
raise HTTPException(status_code=500, detail="Admin password not configured") raise HTTPException(
status_code=500, detail="Admin password not configured"
)
if not vault.verify_password( if not vault.verify_password(
password_update.current_password, secret.admin_password_hash password_update.current_password, secret.admin_password_hash
): ):
raise HTTPException(status_code=401, detail="Current password is incorrect") raise HTTPException(
status_code=401, detail="Current password is incorrect"
)
# Validate new password # Validate new password
new_password = password_update.new_password.strip() new_password = password_update.new_password.strip()
@@ -978,7 +980,9 @@ async def update_upstream_provider_by_slug(
lookup = _validate_slug(payload.slug) lookup = _validate_slug(payload.slug)
async with create_session() as session: async with create_session() as session:
result = await session.exec( result = await session.exec(
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) select(UpstreamProviderRow).where(
UpstreamProviderRow.slug == lookup
)
) )
provider = result.first() provider = result.first()
if not provider: if not provider:
@@ -1665,11 +1669,7 @@ async def get_transactions_api(
) )
total = count_result.one() total = count_result.one()
stmt = ( stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
base.order_by(col(CashuTransaction.created_at).desc())
.offset(offset)
.limit(limit)
)
results = await session.exec(stmt) results = await session.exec(stmt)
transactions = results.all() transactions = results.all()
@@ -1679,7 +1679,9 @@ async def get_transactions_api(
} }
@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) @admin_router.get(
"/api/lightning-invoices", dependencies=[Depends(require_admin_api)]
)
async def get_lightning_invoices_api( async def get_lightning_invoices_api(
status: str | None = None, status: str | None = None,
purpose: str | None = None, purpose: str | None = None,
+11 -10
View File
@@ -408,9 +408,7 @@ class LogManager:
def get_error_details(self, hours: int = 24, limit: int = 100) -> dict: def get_error_details(self, hours: int = 24, limit: int = 100) -> dict:
def compute() -> dict: def compute() -> dict:
try: try:
return self._usage_store.get_error_details( return self._usage_store.get_error_details(hours_back=hours, limit=limit)
hours_back=hours, limit=limit
)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Usage analytics index failed, falling back to log scan: {e}" f"Usage analytics index failed, falling back to log scan: {e}"
@@ -630,7 +628,8 @@ class LogManager:
stats["total_tokens"] += input_tokens + output_tokens stats["total_tokens"] += input_tokens + output_tokens
failed = ( failed = (
"upstream request failed" in message or "revert payment" in message "upstream request failed" in message
or "revert payment" in message
) )
if failed: if failed:
stats["total_requests"] += 1 stats["total_requests"] += 1
@@ -788,9 +787,7 @@ class LogManager:
if bucket_key: if bucket_key:
model_mix_buckets[bucket_key][model] += 1 model_mix_buckets[bucket_key][model] += 1
if revenue_msats > 0: if revenue_msats > 0:
model_mix_revenue_buckets[bucket_key][model] += ( model_mix_revenue_buckets[bucket_key][model] += revenue_msats
revenue_msats
)
model_mix_revenue_totals[model] += revenue_msats model_mix_revenue_totals[model] += revenue_msats
if input_tokens > 0 or output_tokens > 0: if input_tokens > 0 or output_tokens > 0:
token_total = input_tokens + output_tokens token_total = input_tokens + output_tokens
@@ -804,7 +801,8 @@ class LogManager:
bucket["revenue_msats"] += revenue_msats bucket["revenue_msats"] += revenue_msats
failed = ( failed = (
"upstream request failed" in message or "revert payment" in message "upstream request failed" in message
or "revert payment" in message
) )
if failed: if failed:
summary_stats["total_requests"] += 1 summary_stats["total_requests"] += 1
@@ -874,7 +872,9 @@ class LogManager:
models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True) models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True)
latest_errors = [ latest_errors = [
item item
for _, item in sorted(latest_errors_heap, key=lambda x: x[0], reverse=True) for _, item in sorted(
latest_errors_heap, key=lambda x: x[0], reverse=True
)
] ]
top_model_limit = max(1, min(model_limit, 20)) top_model_limit = max(1, min(model_limit, 20))
top_models_requests = [ top_models_requests = [
@@ -1051,7 +1051,8 @@ class LogManager:
bucket["warnings"] += 1 bucket["warnings"] += 1
failed = ( failed = (
"upstream request failed" in message or "revert payment" in message "upstream request failed" in message
or "revert payment" in message
) )
if failed: if failed:
bucket["total_requests"] += 1 bucket["total_requests"] += 1
+15 -14
View File
@@ -314,7 +314,9 @@ class UsageAnalyticsStore:
if column in existing_columns: if column in existing_columns:
return return
conn.execute(f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}") conn.execute(
f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}"
)
logger.info(f"Migrated analytics schema: added {table}.{column}") logger.info(f"Migrated analytics schema: added {table}.{column}")
def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None: def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None:
@@ -362,11 +364,7 @@ class UsageAnalyticsStore:
self._drop_index_tables_locked(conn) self._drop_index_tables_locked(conn)
self._initialize_schema_locked(conn) self._initialize_schema_locked(conn)
files = ( files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log"))
log_files
if log_files is not None
else sorted(self.logs_dir.glob("app_*.log"))
)
for log_file in files: for log_file in files:
try: try:
self._process_log_file_locked(conn, log_file, force_full_read=True) self._process_log_file_locked(conn, log_file, force_full_read=True)
@@ -570,7 +568,8 @@ class UsageAnalyticsStore:
model_bucket["revenue_msats"] += revenue_msats model_bucket["revenue_msats"] += revenue_msats
failed = ( failed = (
"upstream request failed" in message or "revert payment" in message "upstream request failed" in message
or "revert payment" in message
) )
if failed: if failed:
bucket["total_requests"] += 1 bucket["total_requests"] += 1
@@ -593,9 +592,9 @@ class UsageAnalyticsStore:
if isinstance(max_cost, (int, float)) and max_cost > 0: if isinstance(max_cost, (int, float)) and max_cost > 0:
max_cost_float = float(max_cost) max_cost_float = float(max_cost)
bucket["refunds_msats"] += max_cost_float bucket["refunds_msats"] += max_cost_float
model_updates[(minute_key, model)]["refunds_msats"] += ( model_updates[(minute_key, model)][
max_cost_float "refunds_msats"
) ] += max_cost_float
return ( return (
end_offset, end_offset,
@@ -1033,9 +1032,7 @@ class UsageAnalyticsStore:
""", """,
(cutoff_timestamp,), (cutoff_timestamp,),
).fetchone() ).fetchone()
total_error_count = ( total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0
int(total_error_count_row[0]) if total_error_count_row else 0
)
return { return {
"errors": [ "errors": [
@@ -1207,7 +1204,11 @@ class UsageAnalyticsStore:
total_successful = int(row["total_successful"]) total_successful = int(row["total_successful"])
total_revenue_msats = float(row["total_revenue_msats"]) total_revenue_msats = float(row["total_revenue_msats"])
total_tokens = int(row["total_tokens"]) total_tokens = int(row["total_tokens"])
if total_successful <= 0 and total_revenue_msats <= 0 and total_tokens <= 0: if (
total_successful <= 0
and total_revenue_msats <= 0
and total_tokens <= 0
):
continue continue
bucket_ts = str(row["bucket_ts"]) bucket_ts = str(row["bucket_ts"])
+2 -6
View File
@@ -215,9 +215,7 @@ def _build_window_payload(
summary = dashboard.get("summary", {}) summary = dashboard.get("summary", {})
model_usage_mix = dashboard.get("model_usage_mix", {}) model_usage_mix = dashboard.get("model_usage_mix", {})
summary_payload = _build_summary_payload( summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {})
summary if isinstance(summary, dict) else {}
)
usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {} usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {}
top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload) top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload)
@@ -340,9 +338,7 @@ async def publish_usage_analytics() -> None:
nsec = (settings.nsec or "").strip() nsec = (settings.nsec or "").strip()
if not nsec: if not nsec:
if not warned_missing_nsec: if not warned_missing_nsec:
logger.info( logger.info("NSEC is not configured; skipping analytics sharing to Nostr")
"NSEC is not configured; skipping analytics sharing to Nostr"
)
warned_missing_nsec = True warned_missing_nsec = True
await asyncio.sleep(DISABLED_POLL_SECONDS) await asyncio.sleep(DISABLED_POLL_SECONDS)
continue continue
+14 -5
View File
@@ -224,7 +224,9 @@ async def calculate_cost(
"Token counts %s in the upstream response but cannot be " "Token counts %s in the upstream response but cannot be "
"priced; the request will appear in dashboards with the " "priced; the request will appear in dashboards with the "
"raw counts and a fixed max-cost charge.", "raw counts and a fixed max-cost charge.",
"are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", "are present"
if (input_tokens > 0 or output_tokens > 0)
else "are zero",
extra={ extra={
"base_cost_msats": max_cost, "base_cost_msats": max_cost,
"model": response_data.get("model", "unknown"), "model": response_data.get("model", "unknown"),
@@ -301,7 +303,9 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
# actually deducts from the balance. For non-BYOK providers (e.g. # actually deducts from the balance. For non-BYOK providers (e.g.
# OpenRouter) usage.cost already equals upstream_inference_cost, so we # OpenRouter) usage.cost already equals upstream_inference_cost, so we
# fall through to the normal ``cost`` lookup below. # fall through to the normal ``cost`` lookup below.
upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) upstream_cost = _coerce_usd(
cost_details.get("upstream_inference_cost")
)
if upstream_cost > 0 and usage_data.get("is_byok"): if upstream_cost > 0 and usage_data.get("is_byok"):
byok_fee = _coerce_usd(usage_data.get("cost")) byok_fee = _coerce_usd(usage_data.get("cost"))
return upstream_cost + byok_fee return upstream_cost + byok_fee
@@ -332,7 +336,8 @@ def _get_pricing_rates(
``None`` means configured fixed pricing should be used by the caller. ``None`` means configured fixed pricing should be used by the caller.
""" """
if settings.fixed_pricing and ( if settings.fixed_pricing and (
settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens settings.fixed_per_1k_input_tokens
or settings.fixed_per_1k_output_tokens
): ):
return None return None
@@ -388,8 +393,12 @@ def _get_pricing_rates(
usd_per_sat = sats_usd_price() usd_per_sat = sats_usd_price()
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) cache_read_usd = _coerce_usd(
cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) pricing.get("cache_read_input_token_cost")
)
cache_write_usd = _coerce_usd(
pricing.get("cache_creation_input_token_cost")
)
mscr_1k = ( mscr_1k = (
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
if cache_read_usd > 0 if cache_read_usd > 0
+3 -1
View File
@@ -110,7 +110,9 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None:
if not isinstance(usage_data, dict): if not isinstance(usage_data, dict):
return None return None
output_tokens = _first_token_count(usage_data, "completion_tokens", "output_tokens") output_tokens = _first_token_count(
usage_data, "completion_tokens", "output_tokens"
)
cache_read, cache_write = _extract_cache_tokens(usage_data) cache_read, cache_write = _extract_cache_tokens(usage_data)
# ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic # ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic
+3 -1
View File
@@ -94,7 +94,9 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
deployment_id = deployment_id.split("/")[-1] deployment_id = deployment_id.split("/")[-1]
return f"openai/deployments/{deployment_id}/{clean_path}" return f"openai/deployments/{deployment_id}/{clean_path}"
def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: def get_request_base_url(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Use endpoint root, stripping accidental /openai/v1 suffix if present.""" """Use endpoint root, stripping accidental /openai/v1 suffix if present."""
base_url = self.base_url.rstrip("/") base_url = self.base_url.rstrip("/")
marker = "/openai/v1" marker = "/openai/v1"
+14 -18
View File
@@ -191,9 +191,7 @@ def _resolve_ehbp_target_url(
otherwise the header is ignored so callers cannot redirect other providers otherwise the header is ignored so callers cannot redirect other providers
or leak upstream API keys. or leak upstream API keys.
""" """
override_header = ( override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
)
if not override_header: if not override_header:
return target_url return target_url
enclave_url = _get_header_case_insensitive(headers, override_header) enclave_url = _get_header_case_insensitive(headers, override_header)
@@ -297,7 +295,9 @@ def _build_cost_info(
return result return result
def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: def _inject_cost_response_headers(
headers: dict[str, str], cost_info: dict
) -> None:
"""Add per-request cost headers to an EHBP response. """Add per-request cost headers to an EHBP response.
Since EHBP response bodies are opaque encrypted blobs, cost cannot be Since EHBP response bodies are opaque encrypted blobs, cost cannot be
@@ -375,7 +375,9 @@ async def _compute_ehbp_actual_cost(
resolved_upstream_model = ( resolved_upstream_model = (
actual_model_obj.forwarded_model_id or actual_model_obj.id actual_model_obj.forwarded_model_id or actual_model_obj.id
) )
resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) resolved_identity = _normalize_upstream_model_id(
resolved_upstream_model
)
if resolved_identity != expected_identity: if resolved_identity != expected_identity:
logger.info( logger.info(
"EHBP served model differs from requested, using actual " "EHBP served model differs from requested, using actual "
@@ -515,9 +517,7 @@ async def finalize_ehbp_actual_cost_payment(
billing_key = await get_billing_key(key, session) billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key billing_key_hash = billing_key.hashed_key
total_cost_msats = max( total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
0, int(cost_info.get("total_msats", reserved_cost_for_model))
)
now = int(time.time()) now = int(time.time())
safe_reserved = case( safe_reserved = case(
@@ -560,9 +560,7 @@ async def finalize_ehbp_actual_cost_payment(
) )
child_result = await session.exec(child_stmt) # type: ignore[call-overload] child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if result.rowcount == 0 or ( if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
child_result is not None and child_result.rowcount == 0
):
await session.rollback() await session.rollback()
logger.error( logger.error(
"Failed to finalize EHBP usage-based payment", "Failed to finalize EHBP usage-based payment",
@@ -692,9 +690,7 @@ async def finalize_ehbp_max_cost_payment(
else: else:
child_result = None child_result = None
if result.rowcount == 0 or ( if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
child_result is not None and child_result.rowcount == 0
):
await session.rollback() await session.rollback()
logger.error( logger.error(
"Failed to finalize EHBP max-cost payment", "Failed to finalize EHBP max-cost payment",
@@ -1038,9 +1034,7 @@ async def forward_ehbp_x_cashu_request(
target_url = _resolve_ehbp_target_url( target_url = _resolve_ehbp_target_url(
target.url, path, headers, provider_type, profile target.url, path, headers, provider_type, profile
) )
upstream_headers = _prepare_ehbp_upstream_headers( upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile)
headers, target.headers, profile
)
request_body = await request.body() request_body = await request.body()
# Merge query params into the target URL # Merge query params into the target URL
@@ -1088,7 +1082,9 @@ async def forward_ehbp_x_cashu_request(
usage_source = ( usage_source = (
"header" "header"
if usage_header_name if usage_header_name
and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) and any(
k.lower() == usage_header_name.lower() for k, _ in resp.headers
)
else ("trailer" if usage_header else "none") else ("trailer" if usage_header else "none")
) )
+3 -1
View File
@@ -94,7 +94,9 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
""" """
return self.base_url.rstrip("/").removesuffix("/openai") + "/openai" return self.base_url.rstrip("/").removesuffix("/openai") + "/openai"
def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: def get_request_base_url(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Route every proxied request to the OpenAI-compat surface. """Route every proxied request to the OpenAI-compat surface.
Required because the stored ``base_url`` typically points at the Required because the stored ``base_url`` typically points at the
+3 -1
View File
@@ -371,7 +371,9 @@ async def dispatch_gemini_messages(
aggregates). aggregates).
""" """
if not request_body: if not request_body:
raise UpstreamError("Missing request body for /v1/messages", status_code=400) raise UpstreamError(
"Missing request body for /v1/messages", status_code=400
)
try: try:
body: dict = json.loads(request_body) body: dict = json.loads(request_body)
+1 -3
View File
@@ -20,9 +20,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider):
) )
@classmethod @classmethod
def _build_from_row( def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider":
cls, provider_row: "UpstreamProviderRow"
) -> "GroqUpstreamProvider":
return cls( return cls(
api_key=provider_row.api_key, api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee, provider_fee=provider_row.provider_fee,
+3 -1
View File
@@ -91,7 +91,9 @@ OLLAMA_HOST_HINTS: tuple[str, ...] = (
) )
def detect_litellm_prefix(base_url: str | None, default: str = DEFAULT_PREFIX) -> str: def detect_litellm_prefix(
base_url: str | None, default: str = DEFAULT_PREFIX
) -> str:
"""Return the litellm provider prefix (`"<provider>/"`) for `base_url`. """Return the litellm provider prefix (`"<provider>/"`) for `base_url`.
Falls back to `default` when the host doesn't match any known provider. Falls back to `default` when the host doesn't match any known provider.
+9 -3
View File
@@ -108,7 +108,9 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
return events, buffer return events, buffer
def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: def events_from_chunk(
chunk: object, sse_buffer: bytes
) -> tuple[list[dict], bytes]:
"""Normalize a stream chunk into one or more event dicts. """Normalize a stream chunk into one or more event dicts.
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
@@ -199,7 +201,9 @@ async def aggregate_anthropic_events_to_message(
raw_json = partial_json.pop(idx, None) raw_json = partial_json.pop(idx, None)
if raw_json is not None and idx < len(blocks): if raw_json is not None and idx < len(blocks):
try: try:
blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} blocks[idx]["input"] = (
json.loads(raw_json) if raw_json else {}
)
except json.JSONDecodeError: except json.JSONDecodeError:
blocks[idx]["input"] = raw_json blocks[idx]["input"] = raw_json
elif etype == "message_delta": elif etype == "message_delta":
@@ -441,7 +445,9 @@ async def dispatch_anthropic_messages(
on bad input or upstream failure. on bad input or upstream failure.
""" """
if not request_body: if not request_body:
raise UpstreamError("Missing request body for /v1/messages", status_code=400) raise UpstreamError(
"Missing request body for /v1/messages", status_code=400
)
try: try:
body: dict = json.loads(request_body) body: dict = json.loads(request_body)
+195 -127
View File
@@ -1,19 +1,15 @@
"""Model-path discovery service. """Model-path discovery service.
Exposes every upstream provider path a Routstr model is reachable through. Exposes every selectable upstream route a Routstr model is reachable through.
This is discovery/visibility data only — routing still selects the cheapest or This PR remains discovery-only: request-side routing will consume the opaque
best provider separately. selectors in a follow-up.
A *path* is the provider string that may appear in Routstr chat completion A path is a standard percent-encoded query string containing the normalized
responses. The strings emitted here are produced by the provider's own upstream URL and, for an exact OpenRouter endpoint, its machine-readable tag.
``discovery_path_for_subprovider`` / ``discovery_base_paths`` hooks, which Display names never participate in identity::
mirror ``_apply_provider_field`` so discovery and response stamping cannot
drift:
- Direct upstream -> ``<provider_type>`` e.g. ``anthropic`` url=https%3A%2F%2Fapi.anthropic.com%2Fv1
- Generic/custom OpenRouter-compatible upstream -> ``generic:<name>`` url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider=google-vertex%2Fus
- Native OpenRouter routing to a sub-provider -> ``openrouter:<name>``
- Native OpenRouter with no usable sub-provider -> ``unknown``
""" """
from __future__ import annotations from __future__ import annotations
@@ -21,10 +17,12 @@ from __future__ import annotations
import asyncio import asyncio
import random import random
import time import time
from typing import TYPE_CHECKING, Callable from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable
from urllib.parse import urlencode
import httpx import httpx
from sqlalchemy import insert, or_ from sqlalchemy import insert
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from sqlmodel import col, delete, select from sqlmodel import col, delete, select
@@ -51,6 +49,46 @@ _PERSIST_CHUNK_SIZE = 500
ModelKey = tuple[str, int] ModelKey = tuple[str, int]
@dataclass(frozen=True)
class EndpointIdentity:
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
tag: str
provider_name: str | None
@dataclass(frozen=True)
class DiscoveredPath:
"""One model route ready for persistence and API serialization."""
model_id: str
path: str
upstream_url: str
provider_tag: str | None = None
provider_name: str | None = None
@dataclass(frozen=True)
class ProviderPathSnapshot:
"""Refresh result plus model IDs whose prior rows must survive degradation."""
paths: tuple[DiscoveredPath, ...]
preserve_model_ids: frozenset[str] = frozenset()
def normalize_upstream_url(base_url: str) -> str:
"""Normalize route identity without changing URL semantics."""
return base_url.rstrip("/")
def encode_model_path(base_url: str, provider_tag: str | None = None) -> str:
"""Encode a stable opaque selector for future request-side routing."""
components = [("url", normalize_upstream_url(base_url))]
if provider_tag:
components.append(("provider", provider_tag))
return urlencode(components)
def _make_http_client() -> httpx.AsyncClient: def _make_http_client() -> httpx.AsyncClient:
"""Client factory, separated so tests can substitute a mock transport.""" """Client factory, separated so tests can substitute a mock transport."""
return httpx.AsyncClient() return httpx.AsyncClient()
@@ -69,9 +107,15 @@ def is_openrouter_base_url(base_url: str | None) -> bool:
def exposed_model_id(model: object) -> str: def exposed_model_id(model: object) -> str:
"""Client-visible ``/v1/models`` id for a cached model.""" """Return exactly the ID advertised by ``/v1/models``.
A forwarded ID is already a public routable alias and must remain intact,
including any slash. Without one, ``/v1/models`` exposes the base ID.
"""
forwarded = getattr(model, "forwarded_model_id", None) forwarded = getattr(model, "forwarded_model_id", None)
return forwarded or getattr(model, "id") if forwarded:
return forwarded
return public_model_id(getattr(model, "id"))
def public_model_id(model_id: str) -> str: def public_model_id(model_id: str) -> str:
@@ -116,7 +160,7 @@ class _RefreshCycleState:
""" """
def __init__(self) -> None: def __init__(self) -> None:
self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {} self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {}
self.rate_limited = False self.rate_limited = False
@@ -127,8 +171,8 @@ async def _fetch_openrouter_endpoint_subproviders(
author_slug: str, author_slug: str,
semaphore: asyncio.Semaphore, semaphore: asyncio.Semaphore,
cycle: _RefreshCycleState, cycle: _RefreshCycleState,
) -> list[str] | None: ) -> list[EndpointIdentity] | None:
"""Return sub-provider names for one model, or ``None`` when unknown. """Return exact endpoint identities for one model, or ``None`` when unknown.
``None`` (not ``[]``) signals a degraded fetch — network failure, rate ``None`` (not ``[]``) signals a degraded fetch — network failure, rate
limit, non-200, or an unparseable payload — so callers can distinguish limit, non-200, or an unparseable payload — so callers can distinguish
@@ -143,7 +187,7 @@ async def _fetch_openrouter_endpoint_subproviders(
url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints"
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
result: list[str] | None result: list[EndpointIdentity] | None
async with semaphore: async with semaphore:
try: try:
resp = await client.get( resp = await client.get(
@@ -177,14 +221,24 @@ async def _fetch_openrouter_endpoint_subproviders(
endpoints = resp.json().get("data", {}).get("endpoints", []) endpoints = resp.json().get("data", {}).get("endpoints", [])
if not isinstance(endpoints, list): if not isinstance(endpoints, list):
endpoints = [] endpoints = []
names: list[str] = [] identities: dict[str, EndpointIdentity] = {}
for endpoint in endpoints: for endpoint in endpoints:
provider_name = ( if not isinstance(endpoint, dict):
endpoint.get("provider_name") if isinstance(endpoint, dict) else None continue
tag = endpoint.get("tag")
if not isinstance(tag, str) or not tag.strip():
continue
provider_name = endpoint.get("provider_name")
identities.setdefault(
tag,
EndpointIdentity(
tag=tag,
provider_name=provider_name
if isinstance(provider_name, str) and provider_name
else None,
),
) )
if provider_name: result = list(identities.values())
names.append(provider_name)
result = list(dict.fromkeys(names))
except Exception as e: # noqa: BLE001 except Exception as e: # noqa: BLE001
logger.warning( logger.warning(
"OpenRouter endpoint discovery bad payload", "OpenRouter endpoint discovery bad payload",
@@ -287,46 +341,41 @@ async def _collect_provider_paths(
overrides_by_key: dict[ModelKey, ModelRow] | None = None, overrides_by_key: dict[ModelKey, ModelRow] | None = None,
disabled_model_keys: set[ModelKey] | None = None, disabled_model_keys: set[ModelKey] | None = None,
cycle: _RefreshCycleState | None = None, cycle: _RefreshCycleState | None = None,
) -> list[tuple[str, str]] | None: ) -> ProviderPathSnapshot:
"""Collect ``(model_id, path)`` pairs for one provider instance. """Collect selectable routes while marking model-level degraded fetches.
Emits the provider's ``discovery_base_paths`` for normal upstreams. For A failed OpenRouter lookup preserves only that model's prior rows. Other
OpenRouter-compatible providers, additionally emits one path per OpenRouter models in the same provider still refresh, so a partial outage cannot erase
sub-provider endpoint via ``discovery_path_for_subprovider`` so the strings valid discovery data or freeze the entire provider snapshot.
match response stamping exactly.
Returns ``None`` when the provider's path set could not be determined this
cycle (every endpoint fetch degraded); callers must then keep previously
persisted rows instead of wiping them.
""" """
cycle = cycle or _RefreshCycleState() cycle = cycle or _RefreshCycleState()
models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys)
base_paths = upstream.discovery_base_paths() upstream_url = normalize_upstream_url(upstream.base_url)
def _base_path(model: object) -> DiscoveredPath:
return DiscoveredPath(
model_id=exposed_model_id(model),
path=encode_model_path(upstream_url),
upstream_url=upstream_url,
)
if not is_openrouter_base_url(upstream.base_url): if not is_openrouter_base_url(upstream.base_url):
return [ return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models))
(exposed_model_id(model), path) for model in models for path in base_paths
]
if not (upstream.provider_type or "").strip(): if not (upstream.provider_type or "").strip():
return [] return ProviderPathSnapshot(paths=())
any_fetch_succeeded = False
any_fetch_attempted = False
semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY)
async with _make_http_client() as client: async with _make_http_client() as client:
async def _for_model(model: object) -> list[tuple[str, str]]: async def _for_model(
nonlocal any_fetch_succeeded, any_fetch_attempted model: object,
) -> tuple[list[DiscoveredPath], str | None]:
model_id = exposed_model_id(model) model_id = exposed_model_id(model)
# Base paths always apply: responses whose upstream payload lacks a
# provider field are stamped with them (see _apply_provider_field).
pairs = [(model_id, path) for path in base_paths]
author_slug = openrouter_author_slug(model) author_slug = openrouter_author_slug(model)
if not author_slug: if not author_slug:
return pairs return [_base_path(model)], None
any_fetch_attempted = True endpoints = await _fetch_openrouter_endpoint_subproviders(
sub_providers = await _fetch_openrouter_endpoint_subproviders(
client, client,
upstream.base_url, upstream.base_url,
upstream.api_key, upstream.api_key,
@@ -334,67 +383,78 @@ async def _collect_provider_paths(
semaphore, semaphore,
cycle, cycle,
) )
if sub_providers is None: if endpoints is None:
return [] return [], model_id
any_fetch_succeeded = True paths = [_base_path(model)]
paths = [ paths.extend(
upstream.discovery_path_for_subprovider(name) for name in sub_providers DiscoveredPath(
] model_id=model_id,
pairs.extend((model_id, path) for path in paths if path) path=encode_model_path(upstream_url, endpoint.tag),
return list(dict.fromkeys(pairs)) upstream_url=upstream_url,
provider_tag=endpoint.tag,
provider_name=endpoint.provider_name,
)
for endpoint in endpoints
)
return paths, None
results = await asyncio.gather( results = await asyncio.gather(
*(_for_model(m) for m in models), return_exceptions=True *(_for_model(model) for model in models), return_exceptions=True
) )
if any_fetch_attempted and not any_fetch_succeeded: paths: list[DiscoveredPath] = []
# Every endpoint lookup degraded (offline, throttled, bad payloads): preserve_model_ids: set[str] = set()
# the true path set is unknown, not empty. for model, result in zip(models, results):
return None
pairs: list[tuple[str, str]] = []
for result in results:
if isinstance(result, BaseException): if isinstance(result, BaseException):
model_id = exposed_model_id(model)
preserve_model_ids.add(model_id)
logger.warning( logger.warning(
"OpenRouter endpoint discovery task errored", "OpenRouter endpoint discovery task errored",
extra={"provider": upstream.provider_type, "error": str(result)}, extra={"provider": upstream.provider_type, "error": str(result)},
) )
continue continue
pairs.extend(result) model_paths, preserved_model_id = result
paths.extend(model_paths)
if preserved_model_id:
preserve_model_ids.add(preserved_model_id)
return pairs return ProviderPathSnapshot(
paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids)
)
async def _persist_provider_paths( async def _persist_provider_paths(
upstream_provider_id: int, pairs: list[tuple[str, str]] upstream_provider_id: int, snapshot: ProviderPathSnapshot
) -> None: ) -> None:
"""Replace all rows for ``upstream_provider_id`` with ``pairs``. """Replace refreshed rows while retaining model-level degraded snapshots."""
unique_paths = list(
Replacement (not upsert) so stale paths disappear when provider config or {(path.model_id, path.path): path for path in snapshot.paths}.values()
upstream availability changes. Rows are written with chunked bulk INSERTs )
so the transaction holds SQLite's write lock briefly — billing writes share
this database file.
"""
unique_pairs = list(dict.fromkeys(pairs))
now = int(time.time()) now = int(time.time())
async with create_session() as session: async with create_session() as session:
await session.exec( # type: ignore[call-overload] delete_stmt = delete(ModelPathRow).where(
delete(ModelPathRow).where( col(ModelPathRow.upstream_provider_id) == upstream_provider_id
col(ModelPathRow.upstream_provider_id) == upstream_provider_id
)
) )
for start in range(0, len(unique_pairs), _PERSIST_CHUNK_SIZE): if snapshot.preserve_model_ids:
chunk = unique_pairs[start : start + _PERSIST_CHUNK_SIZE] delete_stmt = delete_stmt.where(
col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids))
)
await session.exec(delete_stmt) # type: ignore[call-overload]
for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE):
chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE]
await session.execute( await session.execute(
insert(ModelPathRow), insert(ModelPathRow),
[ [
{ {
"model_id": model_id, "model_id": discovered.model_id,
"path": path, "path": discovered.path,
"upstream_url": discovered.upstream_url,
"provider_tag": discovered.provider_tag,
"provider_name": discovered.provider_name,
"upstream_provider_id": upstream_provider_id, "upstream_provider_id": upstream_provider_id,
"updated_at": now, "updated_at": now,
} }
for model_id, path in chunk for discovered in chunk
], ],
) )
await session.commit() await session.commit()
@@ -454,22 +514,22 @@ async def refresh_model_paths(
if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: if upstream.db_id is None or upstream.db_id not in enabled_provider_ids:
continue continue
try: try:
pairs = await _collect_provider_paths( snapshot = await _collect_provider_paths(
upstream, upstream,
overrides_by_key=overrides_by_key, overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys, disabled_model_keys=disabled_model_keys,
cycle=cycle, cycle=cycle,
) )
if pairs is None: if snapshot.preserve_model_ids:
logger.warning( logger.warning(
"Model paths unknown this cycle; keeping previous rows", "Some model paths are unknown; keeping their previous rows",
extra={ extra={
"provider": upstream.provider_type or upstream.base_url, "provider": upstream.provider_type or upstream.base_url,
"db_id": upstream.db_id, "db_id": upstream.db_id,
"preserved_models": len(snapshot.preserve_model_ids),
}, },
) )
continue await _persist_provider_paths(upstream.db_id, snapshot)
await _persist_provider_paths(upstream.db_id, pairs)
except Exception as e: # noqa: BLE001 - isolate per-provider failures except Exception as e: # noqa: BLE001 - isolate per-provider failures
logger.error( logger.error(
"Failed to refresh model paths for provider", "Failed to refresh model paths for provider",
@@ -482,6 +542,21 @@ async def refresh_model_paths(
) )
async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None:
"""Immediately synchronize discovery after an admin provider/model mutation."""
from ..proxy import get_upstreams
matching = [
upstream
for upstream in get_upstreams()
if upstream.db_id == upstream_provider_id
]
if matching:
await refresh_model_paths(matching)
else:
await prune_model_paths_for_inactive_providers()
def _refresh_interval_seconds() -> int: def _refresh_interval_seconds() -> int:
"""Current interval, re-read every loop so runtime setting changes apply.""" """Current interval, re-read every loop so runtime setting changes apply."""
from ..core.settings import settings from ..core.settings import settings
@@ -535,8 +610,19 @@ async def refresh_model_paths_periodically(
break break
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
provider = None
if row.provider_tag or row.provider_name:
provider = {"name": row.provider_name, "slug": row.provider_tag}
return {
"path": row.path,
"upstream_url": row.upstream_url,
"provider": provider,
}
async def get_all_model_paths() -> dict: async def get_all_model_paths() -> dict:
"""All models with their paths, shaped for ``GET /v1/models/paths``.""" """All models with their exact selectable routes."""
async with create_session() as session: async with create_session() as session:
rows = ( rows = (
await session.exec( await session.exec(
@@ -548,49 +634,37 @@ async def get_all_model_paths() -> dict:
) )
).all() ).all()
grouped: dict[str, list[dict]] = {} grouped: dict[str, list[dict[str, Any]]] = {}
seen_paths: dict[str, set[str]] = {} seen_paths: dict[str, set[str]] = {}
updated_at = 0 updated_at = 0
for row in rows: for row in rows:
updated_at = max(updated_at, row.updated_at) updated_at = max(updated_at, row.updated_at)
model_id = public_model_id(row.model_id) if row.path in seen_paths.setdefault(row.model_id, set()):
if row.path in seen_paths.setdefault(model_id, set()):
continue continue
seen_paths[model_id].add(row.path) seen_paths[row.model_id].add(row.path)
grouped.setdefault(model_id, []).append({"path": row.path}) grouped.setdefault(row.model_id, []).append(_serialize_path(row))
# Deterministic output: models sorted by public id, paths sorted within. data = [
data: list[dict] = [] {
for grouped_model_id in sorted(grouped): "id": grouped_model_id,
model_paths = sorted(grouped[grouped_model_id], key=lambda p: str(p["path"])) "paths": sorted(
data.append({"id": grouped_model_id, "paths": model_paths}) grouped[grouped_model_id], key=lambda item: str(item["path"])
),
}
for grouped_model_id in sorted(grouped)
]
return {"data": data, "updated_at": updated_at or None} return {"data": data, "updated_at": updated_at or None}
async def get_paths_for_model(model_id: str) -> dict: async def get_paths_for_model(model_id: str) -> dict:
"""Paths for a single model, shaped for ``GET /v1/models/paths/model``. """Return paths only for the exact model ID advertised by ``/v1/models``."""
Match by the public, unqualified model id, mirroring the model cache alias
behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve
every row whose stored id has the same base model id. The candidate set is
narrowed in SQL (exact id or ``%/<id>`` suffix) so the route does not
materialize the whole table per request.
"""
# The request may be a full stored id ("z-ai/glm-5v-turbo") or an
# already-stripped public id ("fireworks/models/glm-5"); accept both.
accepted_ids = {model_id, public_model_id(model_id)}
async with create_session() as session: async with create_session() as session:
conditions = []
for candidate in accepted_ids:
conditions.append(col(ModelPathRow.model_id) == candidate)
conditions.append(col(ModelPathRow.model_id).endswith(f"/{candidate}"))
rows = ( rows = (
await session.exec( await session.exec(
select(ModelPathRow) select(ModelPathRow)
.where(or_(*conditions)) .where(col(ModelPathRow.model_id) == model_id)
.order_by( .order_by(
col(ModelPathRow.path), col(ModelPathRow.path),
col(ModelPathRow.upstream_provider_id), col(ModelPathRow.upstream_provider_id),
col(ModelPathRow.model_id),
) )
) )
).all() ).all()
@@ -599,15 +673,9 @@ async def get_paths_for_model(model_id: str) -> dict:
paths: list[dict] = [] paths: list[dict] = []
updated_at = 0 updated_at = 0
for row in rows: for row in rows:
# The SQL suffix match is a prefilter; enforce the exact public-id rule.
if (
row.model_id not in accepted_ids
and public_model_id(row.model_id) not in accepted_ids
):
continue
updated_at = max(updated_at, row.updated_at) updated_at = max(updated_at, row.updated_at)
if row.path in seen: if row.path in seen:
continue continue
seen.add(row.path) seen.add(row.path)
paths.append({"path": row.path}) paths.append(_serialize_path(row))
return {"data": paths, "updated_at": updated_at or None} return {"data": paths, "updated_at": updated_at or None}
+4 -4
View File
@@ -66,7 +66,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility.""" """Strip 'ollama/' prefix for Ollama API compatibility."""
return model_id.removeprefix("ollama/") return model_id.removeprefix("ollama/")
def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: def get_request_base_url(
self, path: str, model_obj: Model | None = None
) -> str:
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
return f"{self.base_url.rstrip('/')}/v1" return f"{self.base_url.rstrip('/')}/v1"
@@ -183,9 +185,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
except Exception: except Exception:
self._models_cache = models_with_fees self._models_cache = models_with_fees
self._models_by_id = { self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
m.forwarded_model_id or m.id: m for m in self._models_cache
}
logger.info( logger.info(
f"Refreshed models cache for {self.base_url}", f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)}, extra={"model_count": len(models)},
+1 -3
View File
@@ -119,9 +119,7 @@ def classify_rate_limit(
retry_match = _RETRY_RE.search(redacted) retry_match = _RETRY_RE.search(redacted)
if retry_match is not None: if retry_match is not None:
value = float(retry_match.group(1)) value = float(retry_match.group(1))
retry_after = ( retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value
value / 1000.0 if retry_match.group(2).lower() == "ms" else value
)
limit_name_match = _LIMIT_NAME_RE.search(redacted) limit_name_match = _LIMIT_NAME_RE.search(redacted)
+3 -1
View File
@@ -84,7 +84,9 @@ def extract_error_message(response: Response) -> str:
return "" return ""
def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: def strip_unsupported_param(
body: dict, error_message: str
) -> tuple[dict, str] | None:
"""Drop a top-level param the upstream named as unsupported/deprecated. """Drop a top-level param the upstream named as unsupported/deprecated.
Returns ``(new_body, param)`` (a new dict, original untouched) when the Returns ``(new_body, param)`` (a new dict, original untouched) when the
+2 -1
View File
@@ -50,7 +50,8 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
def normalize_request_path( def normalize_request_path(
self, path: str, model_obj: "Model | None" = None self, path: str, model_obj: "Model | None" = None
) -> str: ) -> str:
"""Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.""" """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.
"""
return path.lstrip("/") return path.lstrip("/")
@classmethod @classmethod
+1 -3
View File
@@ -21,9 +21,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider):
) )
@classmethod @classmethod
def _build_from_row( def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider":
cls, provider_row: "UpstreamProviderRow"
) -> "XAIUpstreamProvider":
return cls( return cls(
api_key=provider_row.api_key, api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee, provider_fee=provider_row.provider_fee,
+3 -10
View File
@@ -243,9 +243,7 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) all_mint_urls = list({k.mint_url for k in wallet.keysets.values()})
proof_summary = { proof_summary = {
f"{k.mint_url}/{k.unit.name}": sum( f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id)
p.amount for p in wallet.proofs if p.id == k.id
)
for k in wallet.keysets.values() for k in wallet.keysets.values()
} }
# Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet
@@ -600,16 +598,11 @@ async def swap_to_primary_mint(
# advance the counter so the next request derives fresh secrets. # advance the counter so the next request derives fresh secrets.
logger.warning( logger.warning(
"swap_to_primary_mint: outputs already signed — recovering orphaned proofs", "swap_to_primary_mint: outputs already signed — recovering orphaned proofs",
extra={ extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount},
"mint_quote_id": mint_quote.quote,
"minted_amount": minted_amount,
},
) )
try: try:
for keyset_id in primary_wallet.keysets: for keyset_id in primary_wallet.keysets:
await primary_wallet.restore_tokens_for_keyset( await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
keyset_id, to=1, batch=25
)
await primary_wallet.load_proofs(reload=True) await primary_wallet.load_proofs(reload=True)
post_recovery_balance = primary_wallet.available_balance.amount post_recovery_balance = primary_wallet.available_balance.amount
balance_gained = post_recovery_balance - pre_mint_balance balance_gained = post_recovery_balance - pre_mint_balance