chore: apply ruff format repo-wide

CI only runs ruff check, so format drift accumulated. Committed separately
so the reformat noise stays out of functional commits.
This commit is contained in:
9qeklajc
2026-07-26 13:23:11 +02:00
parent 73a3f12469
commit 0a00527626
21 changed files with 132 additions and 124 deletions
+10 -10
View File
@@ -232,7 +232,10 @@ def create_model_mappings(
aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
if (
model_to_use.forwarded_model_id
and model_to_use.forwarded_model_id not in aliases
):
aliases.append(model_to_use.forwarded_model_id)
# Try to set each alias
@@ -322,7 +325,10 @@ def create_model_mappings(
aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
if (
model_to_use.forwarded_model_id
and model_to_use.forwarded_model_id not in aliases
):
aliases.append(model_to_use.forwarded_model_id)
for alias in aliases:
@@ -342,16 +348,10 @@ def create_model_mappings(
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
):
if model.forwarded_model_id and model.forwarded_model_id.lower() == alias:
return 5
if (
model.id
and model.id.lower() == alias
):
if model.id and model.id.lower() == alias:
return 4
model_base = get_base_model_id(model.id)
+24 -5
View File
@@ -260,7 +260,11 @@ async def _lookup_key_no_create(
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,
mint_url: str,
) -> None:
"""Restore balance after a failed refund mint attempt."""
restore_stmt = (
@@ -275,7 +279,11 @@ 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,
"mint_url": mint_url,
},
)
@@ -460,11 +468,23 @@ 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,
key.refund_mint_url or "",
)
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,
key.refund_mint_url or "",
)
error_msg = str(e)
logger.error(
"refund_wallet_endpoint: mint/send failed",
@@ -685,7 +705,6 @@ async def reset_child_key_spent(
return {"success": True, "message": "Child key balance reset successfully."}
@router.api_route(
"/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"],
+12 -14
View File
@@ -68,7 +68,9 @@ async def require_admin_api(request: Request) -> None:
async with create_session() as session:
result = await session.exec(select(CliToken).where(CliToken.token == token))
cli_token = result.first()
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
if cli_token and (
cli_token.expires_at is None or cli_token.expires_at > now_ts
):
cli_token.last_used_at = now_ts
session.add(cli_token)
await session.commit()
@@ -255,16 +257,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) ->
secret = await get_secret(session)
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(
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
new_password = password_update.new_password.strip()
@@ -980,9 +978,7 @@ async def update_upstream_provider_by_slug(
lookup = _validate_slug(payload.slug)
async with create_session() as session:
result = await session.exec(
select(UpstreamProviderRow).where(
UpstreamProviderRow.slug == lookup
)
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup)
)
provider = result.first()
if not provider:
@@ -1669,7 +1665,11 @@ async def get_transactions_api(
)
total = count_result.one()
stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
stmt = (
base.order_by(col(CashuTransaction.created_at).desc())
.offset(offset)
.limit(limit)
)
results = await session.exec(stmt)
transactions = results.all()
@@ -1679,9 +1679,7 @@ 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(
status: str | None = None,
purpose: str | None = None,
+10 -11
View File
@@ -408,7 +408,9 @@ class LogManager:
def get_error_details(self, hours: int = 24, limit: int = 100) -> dict:
def compute() -> dict:
try:
return self._usage_store.get_error_details(hours_back=hours, limit=limit)
return self._usage_store.get_error_details(
hours_back=hours, limit=limit
)
except Exception as e:
logger.error(
f"Usage analytics index failed, falling back to log scan: {e}"
@@ -628,8 +630,7 @@ class LogManager:
stats["total_tokens"] += input_tokens + output_tokens
failed = (
"upstream request failed" in message
or "revert payment" in message
"upstream request failed" in message or "revert payment" in message
)
if failed:
stats["total_requests"] += 1
@@ -787,7 +788,9 @@ class LogManager:
if bucket_key:
model_mix_buckets[bucket_key][model] += 1
if revenue_msats > 0:
model_mix_revenue_buckets[bucket_key][model] += revenue_msats
model_mix_revenue_buckets[bucket_key][model] += (
revenue_msats
)
model_mix_revenue_totals[model] += revenue_msats
if input_tokens > 0 or output_tokens > 0:
token_total = input_tokens + output_tokens
@@ -801,8 +804,7 @@ class LogManager:
bucket["revenue_msats"] += revenue_msats
failed = (
"upstream request failed" in message
or "revert payment" in message
"upstream request failed" in message or "revert payment" in message
)
if failed:
summary_stats["total_requests"] += 1
@@ -872,9 +874,7 @@ class LogManager:
models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True)
latest_errors = [
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_models_requests = [
@@ -1051,8 +1051,7 @@ class LogManager:
bucket["warnings"] += 1
failed = (
"upstream request failed" in message
or "revert payment" in message
"upstream request failed" in message or "revert payment" in message
)
if failed:
bucket["total_requests"] += 1
+14 -15
View File
@@ -314,9 +314,7 @@ class UsageAnalyticsStore:
if column in existing_columns:
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}")
def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None:
@@ -364,7 +362,11 @@ class UsageAnalyticsStore:
self._drop_index_tables_locked(conn)
self._initialize_schema_locked(conn)
files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log"))
files = (
log_files
if log_files is not None
else sorted(self.logs_dir.glob("app_*.log"))
)
for log_file in files:
try:
self._process_log_file_locked(conn, log_file, force_full_read=True)
@@ -568,8 +570,7 @@ class UsageAnalyticsStore:
model_bucket["revenue_msats"] += revenue_msats
failed = (
"upstream request failed" in message
or "revert payment" in message
"upstream request failed" in message or "revert payment" in message
)
if failed:
bucket["total_requests"] += 1
@@ -592,9 +593,9 @@ class UsageAnalyticsStore:
if isinstance(max_cost, (int, float)) and max_cost > 0:
max_cost_float = float(max_cost)
bucket["refunds_msats"] += max_cost_float
model_updates[(minute_key, model)][
"refunds_msats"
] += max_cost_float
model_updates[(minute_key, model)]["refunds_msats"] += (
max_cost_float
)
return (
end_offset,
@@ -1032,7 +1033,9 @@ class UsageAnalyticsStore:
""",
(cutoff_timestamp,),
).fetchone()
total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0
total_error_count = (
int(total_error_count_row[0]) if total_error_count_row else 0
)
return {
"errors": [
@@ -1204,11 +1207,7 @@ class UsageAnalyticsStore:
total_successful = int(row["total_successful"])
total_revenue_msats = float(row["total_revenue_msats"])
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
bucket_ts = str(row["bucket_ts"])
+6 -2
View File
@@ -215,7 +215,9 @@ def _build_window_payload(
summary = dashboard.get("summary", {})
model_usage_mix = dashboard.get("model_usage_mix", {})
summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {})
summary_payload = _build_summary_payload(
summary if isinstance(summary, 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)
@@ -338,7 +340,9 @@ async def publish_usage_analytics() -> None:
nsec = (settings.nsec or "").strip()
if not nsec:
if not warned_missing_nsec:
logger.info("NSEC is not configured; skipping analytics sharing to Nostr")
logger.info(
"NSEC is not configured; skipping analytics sharing to Nostr"
)
warned_missing_nsec = True
await asyncio.sleep(DISABLED_POLL_SECONDS)
continue
+5 -14
View File
@@ -224,9 +224,7 @@ async def calculate_cost(
"Token counts %s in the upstream response but cannot be "
"priced; the request will appear in dashboards with the "
"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={
"base_cost_msats": max_cost,
"model": response_data.get("model", "unknown"),
@@ -303,9 +301,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
# actually deducts from the balance. For non-BYOK providers (e.g.
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
# 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"):
byok_fee = _coerce_usd(usage_data.get("cost"))
return upstream_cost + byok_fee
@@ -336,8 +332,7 @@ def _get_pricing_rates(
``None`` means configured fixed pricing should be used by the caller.
"""
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
@@ -393,12 +388,8 @@ def _get_pricing_rates(
usd_per_sat = sats_usd_price()
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
cache_read_usd = _coerce_usd(
pricing.get("cache_read_input_token_cost")
)
cache_write_usd = _coerce_usd(
pricing.get("cache_creation_input_token_cost")
)
cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost"))
cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost"))
mscr_1k = (
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
if cache_read_usd > 0
+1 -3
View File
@@ -110,9 +110,7 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None:
if not isinstance(usage_data, dict):
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)
# ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic
+1 -3
View File
@@ -94,9 +94,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
deployment_id = deployment_id.split("/")[-1]
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."""
base_url = self.base_url.rstrip("/")
marker = "/openai/v1"
+18 -14
View File
@@ -191,7 +191,9 @@ def _resolve_ehbp_target_url(
otherwise the header is ignored so callers cannot redirect other providers
or leak upstream API keys.
"""
override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
override_header = (
profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
)
if not override_header:
return target_url
enclave_url = _get_header_case_insensitive(headers, override_header)
@@ -295,9 +297,7 @@ def _build_cost_info(
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.
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
@@ -375,9 +375,7 @@ async def _compute_ehbp_actual_cost(
resolved_upstream_model = (
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:
logger.info(
"EHBP served model differs from requested, using actual "
@@ -517,7 +515,9 @@ async def finalize_ehbp_actual_cost_payment(
billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key
total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
total_cost_msats = max(
0, int(cost_info.get("total_msats", reserved_cost_for_model))
)
now = int(time.time())
safe_reserved = case(
@@ -560,7 +560,9 @@ async def finalize_ehbp_actual_cost_payment(
)
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
if result.rowcount == 0 or (
child_result is not None and child_result.rowcount == 0
):
await session.rollback()
logger.error(
"Failed to finalize EHBP usage-based payment",
@@ -690,7 +692,9 @@ async def finalize_ehbp_max_cost_payment(
else:
child_result = None
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
if result.rowcount == 0 or (
child_result is not None and child_result.rowcount == 0
):
await session.rollback()
logger.error(
"Failed to finalize EHBP max-cost payment",
@@ -1034,7 +1038,9 @@ async def forward_ehbp_x_cashu_request(
target_url = _resolve_ehbp_target_url(
target.url, path, headers, provider_type, profile
)
upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile)
upstream_headers = _prepare_ehbp_upstream_headers(
headers, target.headers, profile
)
request_body = await request.body()
# Merge query params into the target URL
@@ -1082,9 +1088,7 @@ async def forward_ehbp_x_cashu_request(
usage_source = (
"header"
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")
)
+1 -3
View File
@@ -94,9 +94,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
"""
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.
Required because the stored ``base_url`` typically points at the
+1 -3
View File
@@ -371,9 +371,7 @@ async def dispatch_gemini_messages(
aggregates).
"""
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:
body: dict = json.loads(request_body)
+3 -1
View File
@@ -20,7 +20,9 @@ class GroqUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider":
def _build_from_row(
cls, provider_row: "UpstreamProviderRow"
) -> "GroqUpstreamProvider":
return cls(
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
+1 -3
View File
@@ -91,9 +91,7 @@ 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`.
Falls back to `default` when the host doesn't match any known provider.
+3 -9
View File
@@ -108,9 +108,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
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.
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
@@ -201,9 +199,7 @@ async def aggregate_anthropic_events_to_message(
raw_json = partial_json.pop(idx, None)
if raw_json is not None and idx < len(blocks):
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:
blocks[idx]["input"] = raw_json
elif etype == "message_delta":
@@ -445,9 +441,7 @@ async def dispatch_anthropic_messages(
on bad input or upstream failure.
"""
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:
body: dict = json.loads(request_body)
+4 -4
View File
@@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility."""
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."""
return f"{self.base_url.rstrip('/')}/v1"
@@ -185,7 +183,9 @@ 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.forwarded_model_id or m.id: m for m in self._models_cache
}
logger.info(
f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)},
+3 -1
View File
@@ -119,7 +119,9 @@ def classify_rate_limit(
retry_match = _RETRY_RE.search(redacted)
if retry_match is not None:
value = float(retry_match.group(1))
retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value
retry_after = (
value / 1000.0 if retry_match.group(2).lower() == "ms" else value
)
limit_name_match = _LIMIT_NAME_RE.search(redacted)
+1 -3
View File
@@ -84,9 +84,7 @@ def extract_error_message(response: Response) -> str:
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.
Returns ``(new_body, param)`` (a new dict, original untouched) when the
+1 -2
View File
@@ -50,8 +50,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
def normalize_request_path(
self, path: str, model_obj: "Model | None" = None
) -> str:
"""Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.
"""
"""Preserve the ``v1/`` prefix when forwarding to an upstream Routstr."""
return path.lstrip("/")
@classmethod
+3 -1
View File
@@ -21,7 +21,9 @@ class XAIUpstreamProvider(BaseUpstreamProvider):
)
@classmethod
def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider":
def _build_from_row(
cls, provider_row: "UpstreamProviderRow"
) -> "XAIUpstreamProvider":
return cls(
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
+10 -3
View File
@@ -243,7 +243,9 @@ 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()})
proof_summary = {
f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id)
f"{k.mint_url}/{k.unit.name}": sum(
p.amount for p in wallet.proofs if p.id == k.id
)
for k in wallet.keysets.values()
}
# Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet
@@ -598,11 +600,16 @@ async def swap_to_primary_mint(
# advance the counter so the next request derives fresh secrets.
logger.warning(
"swap_to_primary_mint: outputs already signed — recovering orphaned proofs",
extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount},
extra={
"mint_quote_id": mint_quote.quote,
"minted_amount": minted_amount,
},
)
try:
for keyset_id in primary_wallet.keysets:
await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
await primary_wallet.restore_tokens_for_keyset(
keyset_id, to=1, batch=25
)
await primary_wallet.load_proofs(reload=True)
post_recovery_balance = primary_wallet.available_balance.amount
balance_gained = post_recovery_balance - pre_mint_balance