mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-31 15:56:14 +00:00
Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ceca0e8efc | ||
|
|
89109dc209 | ||
|
|
234ea19cad | ||
|
|
cd7de3958c | ||
|
|
8c0ac499ef | ||
|
|
28d91227af | ||
|
|
e43ceb2e43 | ||
|
|
689a07f562 |
@@ -14,6 +14,14 @@ jobs:
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v3
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Resolve git metadata
|
||||
id: gitmeta
|
||||
run: |
|
||||
echo "sha=$(git rev-parse --short=7 HEAD)" >> "$GITHUB_OUTPUT"
|
||||
echo "tag=$(git describe --tags --exact-match HEAD 2>/dev/null || true)" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v2
|
||||
@@ -30,6 +38,9 @@ jobs:
|
||||
with:
|
||||
context: .
|
||||
push: true
|
||||
build-args: |
|
||||
GIT_COMMIT=${{ steps.gitmeta.outputs.sha }}
|
||||
GIT_TAG=${{ steps.gitmeta.outputs.tag }}
|
||||
tags: |
|
||||
ghcr.io/routstr/proxy:latest
|
||||
ghcr.io/routstr/core:latest
|
||||
|
||||
@@ -21,6 +21,10 @@ WORKDIR /app
|
||||
|
||||
COPY . .
|
||||
|
||||
ARG GIT_COMMIT=""
|
||||
ARG GIT_TAG=""
|
||||
ENV GIT_COMMIT=${GIT_COMMIT}
|
||||
ENV GIT_TAG=${GIT_TAG}
|
||||
ENV PORT=8000
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
|
||||
|
||||
@@ -41,6 +41,10 @@ RUN uv sync --no-dev
|
||||
# Copy the built UI from the ui-builder stage
|
||||
COPY --from=ui-builder /app/ui/out ./ui_out
|
||||
|
||||
ARG GIT_COMMIT=""
|
||||
ARG GIT_TAG=""
|
||||
ENV GIT_COMMIT=${GIT_COMMIT}
|
||||
ENV GIT_TAG=${GIT_TAG}
|
||||
ENV PORT=8000
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
|
||||
|
||||
+71
-41
@@ -1,5 +1,4 @@
|
||||
import asyncio
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import AsyncGenerator
|
||||
@@ -9,6 +8,8 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, RedirectResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from starlette.exceptions import HTTPException
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
from starlette.types import Scope
|
||||
|
||||
from ..auth import periodic_key_reset
|
||||
from ..balance import balance_router, deprecated_wallet_router
|
||||
@@ -30,16 +31,12 @@ from .logging import get_logger, setup_logging
|
||||
from .middleware import LoggingMiddleware
|
||||
from .settings import SettingsService
|
||||
from .settings import settings as global_settings
|
||||
from .version import __version__
|
||||
|
||||
# Initialize logging first
|
||||
setup_logging()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
if os.getenv("VERSION_SUFFIX") is not None:
|
||||
__version__ = f"0.4.3-{os.getenv('VERSION_SUFFIX')}"
|
||||
else:
|
||||
__version__ = "0.4.3"
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
@@ -103,8 +100,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
btc_price_task = asyncio.create_task(update_prices_periodically())
|
||||
pricing_task = asyncio.create_task(update_sats_pricing())
|
||||
if global_settings.models_refresh_interval_seconds > 0:
|
||||
# Pass the accessor (not its current value) so the loop sees providers
|
||||
# added/changed via reinitialize_upstreams() instead of staying pinned
|
||||
# to the startup snapshot.
|
||||
models_refresh_task = asyncio.create_task(
|
||||
refresh_upstreams_models_periodically(get_upstreams())
|
||||
refresh_upstreams_models_periodically(get_upstreams)
|
||||
)
|
||||
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
|
||||
payout_task = asyncio.create_task(periodic_payout())
|
||||
@@ -194,6 +194,23 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
)
|
||||
|
||||
|
||||
class _ImmutableStaticFiles(StaticFiles):
|
||||
"""Static files with long Cache-Control for content-hashed Next.js assets.
|
||||
|
||||
Files under `/_next/static/` are emitted with content hashes in their
|
||||
filenames and never mutate, so we serve them with a one-year immutable
|
||||
cache header so browsers and CDNs stop revalidating on every reload.
|
||||
"""
|
||||
|
||||
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
|
||||
response = await super().get_response(path, scope)
|
||||
if response.status_code == 200:
|
||||
response.headers["Cache-Control"] = (
|
||||
"public, max-age=31536000, immutable"
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
app = FastAPI(version=__version__, lifespan=lifespan)
|
||||
|
||||
|
||||
@@ -240,7 +257,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
|
||||
app.mount(
|
||||
"/_next",
|
||||
StaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
|
||||
_ImmutableStaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
|
||||
name="next-static",
|
||||
)
|
||||
|
||||
@@ -248,10 +265,12 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
async def serve_root_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "index.html")
|
||||
|
||||
# Add explicit route for /index.txt to redirect to /
|
||||
# Serve the App Router RSC payload for the home page.
|
||||
@app.get("/index.txt", include_in_schema=False)
|
||||
async def redirect_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/")
|
||||
async def serve_root_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
@app.get("/admin")
|
||||
async def admin_redirect() -> FileResponse:
|
||||
@@ -265,82 +284,96 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
async def serve_login_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "login" / "index.html")
|
||||
|
||||
# Add explicit route for /login/index.txt to redirect to /login
|
||||
@app.get("/login/index.txt", include_in_schema=False)
|
||||
async def redirect_login_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/login")
|
||||
async def serve_login_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "login" / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
@app.get("/model", include_in_schema=False)
|
||||
async def serve_models_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "model" / "index.html")
|
||||
|
||||
# Add explicit route for /model/index.txt to redirect to /model
|
||||
@app.get("/model/index.txt", include_in_schema=False)
|
||||
async def redirect_model_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/model")
|
||||
async def serve_model_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "model" / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
@app.get("/providers", include_in_schema=False)
|
||||
async def serve_providers_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "providers" / "index.html")
|
||||
|
||||
# Add explicit route for /providers/index.txt to redirect to /providers
|
||||
@app.get("/providers/index.txt", include_in_schema=False)
|
||||
async def redirect_providers_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/providers")
|
||||
async def serve_providers_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "providers" / "index.txt",
|
||||
media_type="text/x-component",
|
||||
)
|
||||
|
||||
@app.get("/settings", include_in_schema=False)
|
||||
async def serve_settings_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "settings" / "index.html")
|
||||
|
||||
# Add explicit route for /settings/index.txt to redirect to /settings
|
||||
@app.get("/settings/index.txt", include_in_schema=False)
|
||||
async def redirect_settings_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/settings")
|
||||
async def serve_settings_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "settings" / "index.txt",
|
||||
media_type="text/x-component",
|
||||
)
|
||||
|
||||
@app.get("/transactions", include_in_schema=False)
|
||||
async def serve_transactions_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "transactions" / "index.html")
|
||||
|
||||
# Add explicit route for /transactions/index.txt to redirect to /transactions
|
||||
@app.get("/transactions/index.txt", include_in_schema=False)
|
||||
async def redirect_transactions_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/transactions")
|
||||
async def serve_transactions_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "transactions" / "index.txt",
|
||||
media_type="text/x-component",
|
||||
)
|
||||
|
||||
@app.get("/balances", include_in_schema=False)
|
||||
async def serve_balances_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "balances" / "index.html")
|
||||
|
||||
# Add explicit route for /balances/index.txt to redirect to /balances
|
||||
@app.get("/balances/index.txt", include_in_schema=False)
|
||||
async def redirect_balances_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/balances")
|
||||
async def serve_balances_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "balances" / "index.txt",
|
||||
media_type="text/x-component",
|
||||
)
|
||||
|
||||
@app.get("/logs", include_in_schema=False)
|
||||
async def serve_logs_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "logs" / "index.html")
|
||||
|
||||
# Add explicit route for /logs/index.txt to redirect to /logs
|
||||
@app.get("/logs/index.txt", include_in_schema=False)
|
||||
async def redirect_logs_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/logs")
|
||||
async def serve_logs_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "logs" / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
@app.get("/usage", include_in_schema=False)
|
||||
async def serve_usage_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "usage" / "index.html")
|
||||
|
||||
# Add explicit route for /usage/index.txt to redirect to /usage
|
||||
@app.get("/usage/index.txt", include_in_schema=False)
|
||||
async def redirect_usage_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/usage")
|
||||
async def serve_usage_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "usage" / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
@app.get("/unauthorized", include_in_schema=False)
|
||||
async def serve_unauthorized_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "unauthorized" / "index.html")
|
||||
|
||||
# Add explicit route for /unauthorized/index.txt to redirect to /unauthorized
|
||||
@app.get("/unauthorized/index.txt", include_in_schema=False)
|
||||
async def redirect_unauthorized_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/unauthorized")
|
||||
async def serve_unauthorized_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "unauthorized" / "index.txt",
|
||||
media_type="text/x-component",
|
||||
)
|
||||
|
||||
@app.get("/favicon.ico", include_in_schema=False)
|
||||
async def serve_favicon() -> FileResponse:
|
||||
@@ -353,9 +386,6 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
async def serve_icon() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "icon.ico")
|
||||
|
||||
app.mount(
|
||||
"/static", StaticFiles(directory=UI_DIST_PATH, check_dir=True), name="ui-static"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"UI dist directory not found at {UI_DIST_PATH}, skipping static file serving"
|
||||
|
||||
+69
-63
@@ -14,8 +14,54 @@ logger = get_logger(__name__)
|
||||
request_id_context: ContextVar[str | None] = ContextVar("request_id")
|
||||
|
||||
|
||||
# Methods that are never logged: HEAD requests are health probes from
|
||||
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
|
||||
# chatter, not user-meaningful events.
|
||||
_SKIP_LOG_METHODS: frozenset[str] = frozenset({"HEAD", "OPTIONS"})
|
||||
|
||||
# Path prefixes to skip. Includes Next.js static chunks and the admin
|
||||
# dashboard's internal polling API (/admin/api/*) which the UI hits on a timer
|
||||
# to refresh balances, logs, providers, etc. — high volume, low diagnostic
|
||||
# value. Mutating admin actions are recorded separately in the audit log.
|
||||
_SKIP_LOG_PREFIXES: tuple[str, ...] = (
|
||||
"/_next/",
|
||||
"/admin/api/",
|
||||
)
|
||||
|
||||
# Exact paths to skip. RSC payload prefetches (`*/index.txt`) fire automatically
|
||||
# as the user hovers near `<Link>`s, and `/v1/wallet/info` is polled by the UI.
|
||||
_SKIP_LOG_EXACT: frozenset[str] = frozenset(
|
||||
{
|
||||
"/favicon.ico",
|
||||
"/icon.ico",
|
||||
"/v1/wallet/info",
|
||||
"/index.txt",
|
||||
"/login/index.txt",
|
||||
"/model/index.txt",
|
||||
"/providers/index.txt",
|
||||
"/settings/index.txt",
|
||||
"/transactions/index.txt",
|
||||
"/balances/index.txt",
|
||||
"/logs/index.txt",
|
||||
"/usage/index.txt",
|
||||
"/unauthorized/index.txt",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _should_log(method: str, path: str) -> bool:
|
||||
if method in _SKIP_LOG_METHODS:
|
||||
return False
|
||||
if path in _SKIP_LOG_EXACT:
|
||||
return False
|
||||
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
|
||||
|
||||
|
||||
class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""Middleware to log detailed request and response information."""
|
||||
"""Middleware to log proxy interactions and page navigation.
|
||||
|
||||
Skips logging for static assets and Next.js chunks to avoid noise.
|
||||
"""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
# Generate request ID
|
||||
@@ -25,56 +71,20 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
# Set request ID in context for logging
|
||||
token = request_id_context.set(request_id)
|
||||
|
||||
path = request.url.path
|
||||
should_log = _should_log(request.method, path)
|
||||
|
||||
# Start timing
|
||||
start_time = time.time()
|
||||
|
||||
# Log request details
|
||||
request_body = None
|
||||
if request.method in ["POST", "PUT", "PATCH"]:
|
||||
try:
|
||||
# Only read body for non-streaming requests
|
||||
if hasattr(request, "_body"):
|
||||
request_body = await request.body()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Log incoming request
|
||||
logger.info(
|
||||
"Incoming request",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"query_params": dict(request.query_params),
|
||||
"headers": {
|
||||
k: v
|
||||
for k, v in request.headers.items()
|
||||
if k.lower()
|
||||
not in [
|
||||
"authorization",
|
||||
"x-cashu",
|
||||
"cookie",
|
||||
"cf-connecting-ip",
|
||||
"cf-ipcountry",
|
||||
"x-forwarded-for",
|
||||
"x-real-ip",
|
||||
]
|
||||
},
|
||||
"body_size": len(request_body) if request_body else 0,
|
||||
},
|
||||
)
|
||||
|
||||
# Log at TRACE level for full body (security filter will redact sensitive data)
|
||||
if request_body and hasattr(logger, "exception"):
|
||||
logger.exception(
|
||||
"Request body",
|
||||
if should_log:
|
||||
logger.info(
|
||||
"Incoming request",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"body": request_body.decode("utf-8", errors="ignore")[
|
||||
:1000
|
||||
], # Limit size
|
||||
"path": path,
|
||||
"query_params": dict(request.query_params),
|
||||
},
|
||||
)
|
||||
|
||||
@@ -82,36 +92,32 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
try:
|
||||
response = await call_next(request)
|
||||
|
||||
# Calculate duration
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Log response
|
||||
logger.info(
|
||||
"Request completed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"status_code": response.status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
},
|
||||
)
|
||||
if should_log:
|
||||
duration = time.time() - start_time
|
||||
logger.info(
|
||||
"Request completed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": path,
|
||||
"status_code": response.status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
},
|
||||
)
|
||||
if hasattr(response, "headers"):
|
||||
response.headers["x-routstr-request-id"] = request_id
|
||||
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
# Calculate duration
|
||||
# Always log failures, even for skipped paths, so we don't lose errors.
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Log error
|
||||
logger.error(
|
||||
"Request failed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"path": path,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Application version resolution.
|
||||
|
||||
Priority order:
|
||||
|
||||
1. ``VERSION_SUFFIX`` env var (manual override; preserves prior behaviour).
|
||||
2. Bare base version when HEAD is on the matching release tag (detected via
|
||||
``GIT_TAG`` env or ``git describe --tags --exact-match HEAD``).
|
||||
3. ``GIT_COMMIT`` env var (build-time injection) -> ``<base>+g<sha>``.
|
||||
4. Local ``.git`` lookup (source checkouts) -> ``<base>+g<sha>``.
|
||||
5. Fallback: bare base version.
|
||||
|
||||
The ``+g<sha>`` form is PEP 440 local-version syntax so the result remains a
|
||||
valid package version.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
BASE_VERSION = "0.4.3"
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
_GIT_TIMEOUT_SECONDS = 2.0
|
||||
|
||||
|
||||
def _run_git(*args: str) -> str | None:
|
||||
try:
|
||||
result = subprocess.run( # noqa: S603 - fixed argv, no shell
|
||||
["git", *args],
|
||||
cwd=_REPO_ROOT,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=_GIT_TIMEOUT_SECONDS,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.SubprocessError, OSError):
|
||||
return None
|
||||
if result.returncode != 0:
|
||||
return None
|
||||
return result.stdout.strip() or None
|
||||
|
||||
|
||||
def _git_short_sha() -> str | None:
|
||||
sha = os.getenv("GIT_COMMIT", "").strip()
|
||||
if sha:
|
||||
return sha[:7]
|
||||
return _run_git("rev-parse", "--short=7", "HEAD")
|
||||
|
||||
|
||||
def _on_tagged_release() -> bool:
|
||||
tag = os.getenv("GIT_TAG", "").strip()
|
||||
if tag:
|
||||
return tag.lstrip("v") == BASE_VERSION
|
||||
described = _run_git("describe", "--tags", "--exact-match", "HEAD")
|
||||
if not described:
|
||||
return False
|
||||
return described.lstrip("v") == BASE_VERSION
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_version() -> str:
|
||||
suffix = os.getenv("VERSION_SUFFIX")
|
||||
if suffix is not None:
|
||||
return f"{BASE_VERSION}-{suffix}"
|
||||
|
||||
if _on_tagged_release():
|
||||
return BASE_VERSION
|
||||
|
||||
sha = _git_short_sha()
|
||||
if not sha:
|
||||
return BASE_VERSION
|
||||
|
||||
return f"{BASE_VERSION}+g{sha}"
|
||||
|
||||
|
||||
__version__ = get_version()
|
||||
|
||||
__all__ = ["BASE_VERSION", "__version__", "get_version"]
|
||||
@@ -26,7 +26,7 @@ logger = get_logger(__name__)
|
||||
|
||||
def get_app_version() -> str | None:
|
||||
try:
|
||||
from ..core.main import __version__ as imported_version
|
||||
from ..core.version import __version__ as imported_version
|
||||
|
||||
return imported_version
|
||||
except Exception:
|
||||
|
||||
+177
-107
@@ -604,57 +604,83 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
yield prefix + part
|
||||
|
||||
# Stream finished, process usage if found
|
||||
if usage_chunk_data:
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(key.__class__, key.hashed_key)
|
||||
if fresh_key:
|
||||
try:
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
usage_chunk_data,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
remaining_balance_msats = fresh_key.balance
|
||||
# Merge cost into usage
|
||||
usage_chunk_data["usage"]["cost"] = cost_data.get(
|
||||
"total_usd", 0.0
|
||||
)
|
||||
usage_chunk_data["usage"]["cost_sats"] = (
|
||||
cost_data.get("total_msats", 0) // 1000
|
||||
)
|
||||
usage_chunk_data["usage"]["remaining_balance_msats"] = (
|
||||
remaining_balance_msats
|
||||
)
|
||||
# Keep detailed cost in metadata
|
||||
usage_chunk_data["metadata"] = usage_chunk_data.get(
|
||||
"metadata", {}
|
||||
)
|
||||
usage_chunk_data["metadata"]["routstr"] = {
|
||||
"cost": cost_data
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(key.__class__, key.hashed_key)
|
||||
if fresh_key:
|
||||
cost_data: dict
|
||||
try:
|
||||
adjustment_input = (
|
||||
usage_chunk_data
|
||||
if usage_chunk_data is not None
|
||||
else {
|
||||
"model": last_model_seen or "unknown",
|
||||
"usage": None,
|
||||
}
|
||||
usage_chunk_data["metadata"]["routstr"]["cost"][
|
||||
"sats_cost"
|
||||
] = cost_data.get("total_msats", 0) // 1000
|
||||
usage_chunk_data["metadata"]["routstr"]["cost"][
|
||||
"remaining_balance_msats"
|
||||
] = remaining_balance_msats
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
"Error during usage finalization",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
# Fallback: yield original usage chunk if adjustment fails
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
)
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
"Error during usage finalization",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
# Fall back so we still emit a non-zero sats cost downstream.
|
||||
cost_data = {
|
||||
"base_msats": 0,
|
||||
"input_msats": 0,
|
||||
"output_msats": 0,
|
||||
"total_msats": 0,
|
||||
"total_usd": 0.0,
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
}
|
||||
|
||||
if usage_chunk_data is None:
|
||||
if not hasattr(self, "_current_stream_id"):
|
||||
self._current_stream_id = (
|
||||
f"chatcmpl-{uuid.uuid4()}"
|
||||
)
|
||||
usage_chunk_data = {
|
||||
"id": self._current_stream_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"model": last_model_seen or "unknown",
|
||||
"choices": [],
|
||||
"usage": {
|
||||
"prompt_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
),
|
||||
"completion_tokens": cost_data.get(
|
||||
"output_tokens", 0
|
||||
),
|
||||
"total_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
)
|
||||
+ cost_data.get("output_tokens", 0),
|
||||
},
|
||||
}
|
||||
|
||||
try:
|
||||
self.inject_cost_metadata(
|
||||
usage_chunk_data, cost_data, fresh_key
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to inject cost metadata into streaming chunk",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
|
||||
if done_seen:
|
||||
yield b"data: [DONE]\n\n"
|
||||
@@ -926,65 +952,108 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
yield prefix + part
|
||||
|
||||
# Stream finished, process usage if found
|
||||
if usage_chunk_data:
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(key.__class__, key.hashed_key)
|
||||
if fresh_key:
|
||||
try:
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
usage_chunk_data,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
remaining_balance_msats = fresh_key.balance
|
||||
# Merge cost into usage chunk
|
||||
if (
|
||||
"response" in usage_chunk_data
|
||||
and "usage" in usage_chunk_data["response"]
|
||||
):
|
||||
usage_chunk_data["response"]["usage"]["cost"] = (
|
||||
cost_data.get("total_usd", 0.0)
|
||||
)
|
||||
usage_chunk_data["response"]["usage"][
|
||||
"cost_sats"
|
||||
] = cost_data.get("total_msats", 0) // 1000
|
||||
usage_chunk_data["response"]["usage"][
|
||||
"remaining_balance_msats"
|
||||
] = remaining_balance_msats
|
||||
elif "usage" in usage_chunk_data:
|
||||
usage_chunk_data["usage"]["cost"] = cost_data.get(
|
||||
"total_usd", 0.0
|
||||
)
|
||||
usage_chunk_data["usage"]["cost_sats"] = (
|
||||
cost_data.get("total_msats", 0) // 1000
|
||||
)
|
||||
usage_chunk_data["usage"][
|
||||
"remaining_balance_msats"
|
||||
] = remaining_balance_msats
|
||||
|
||||
# Keep detailed cost in metadata
|
||||
usage_chunk_data["metadata"] = usage_chunk_data.get(
|
||||
"metadata", {}
|
||||
)
|
||||
usage_chunk_data["metadata"]["routstr"] = {
|
||||
"cost": cost_data
|
||||
# Always emit a cost-bearing data chunk
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(key.__class__, key.hashed_key)
|
||||
if fresh_key:
|
||||
cost_data: dict
|
||||
try:
|
||||
adjustment_input = (
|
||||
usage_chunk_data
|
||||
if usage_chunk_data is not None
|
||||
else {
|
||||
"model": last_model_seen or "unknown",
|
||||
"usage": None,
|
||||
}
|
||||
usage_chunk_data["metadata"]["routstr"]["cost"][
|
||||
"sats_cost"
|
||||
] = cost_data.get("total_msats", 0) // 1000
|
||||
usage_chunk_data["metadata"]["routstr"]["cost"][
|
||||
"remaining_balance_msats"
|
||||
] = remaining_balance_msats
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
# Fallback: yield original usage chunk if adjustment fails
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
)
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
"Error during Responses API usage finalization",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
cost_data = {
|
||||
"base_msats": 0,
|
||||
"input_msats": 0,
|
||||
"output_msats": 0,
|
||||
"total_msats": 0,
|
||||
"total_usd": 0.0,
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
}
|
||||
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
if usage_chunk_data is None:
|
||||
usage_chunk_data = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"model": last_model_seen or "unknown",
|
||||
"usage": {
|
||||
"input_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
),
|
||||
"output_tokens": cost_data.get(
|
||||
"output_tokens", 0
|
||||
),
|
||||
"total_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
)
|
||||
+ cost_data.get("output_tokens", 0),
|
||||
},
|
||||
},
|
||||
"usage": {
|
||||
"input_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
),
|
||||
"output_tokens": cost_data.get(
|
||||
"output_tokens", 0
|
||||
),
|
||||
"total_tokens": cost_data.get(
|
||||
"input_tokens", 0
|
||||
)
|
||||
+ cost_data.get("output_tokens", 0),
|
||||
},
|
||||
}
|
||||
|
||||
remaining_balance_msats = fresh_key.balance
|
||||
sats_cost = cost_data.get("total_msats", 0) // 1000
|
||||
|
||||
if (
|
||||
"response" in usage_chunk_data
|
||||
and isinstance(usage_chunk_data["response"], dict)
|
||||
and "usage" in usage_chunk_data["response"]
|
||||
):
|
||||
usage_chunk_data["response"]["usage"]["cost"] = (
|
||||
cost_data.get("total_usd", 0.0)
|
||||
)
|
||||
usage_chunk_data["response"]["usage"][
|
||||
"cost_sats"
|
||||
] = sats_cost
|
||||
usage_chunk_data["response"]["usage"][
|
||||
"remaining_balance_msats"
|
||||
] = remaining_balance_msats
|
||||
|
||||
try:
|
||||
self.inject_cost_metadata(
|
||||
usage_chunk_data, cost_data, fresh_key
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to inject cost metadata into Responses streaming chunk",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
||||
|
||||
if done_seen:
|
||||
yield b"data: [DONE]\n\n"
|
||||
@@ -1308,7 +1377,8 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
usage_finalized = True
|
||||
yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
|
||||
# Emit the full combined_data as the cost
|
||||
yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.settings import Settings
|
||||
@@ -122,12 +122,16 @@ async def get_all_models_with_overrides(
|
||||
|
||||
|
||||
async def refresh_upstreams_models_periodically(
|
||||
upstreams: list[BaseUpstreamProvider],
|
||||
upstreams_provider: (
|
||||
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
|
||||
),
|
||||
) -> None:
|
||||
"""Background task to periodically refresh models cache for all providers.
|
||||
|
||||
Args:
|
||||
upstreams: List of upstream provider instances
|
||||
upstreams_provider: Either a callable returning the live upstream list
|
||||
(preferred — picks up providers added/changed via reinitialize_upstreams),
|
||||
or a static list (legacy, will go stale after reinitialize_upstreams).
|
||||
"""
|
||||
import asyncio
|
||||
import random
|
||||
@@ -139,9 +143,14 @@ async def refresh_upstreams_models_periodically(
|
||||
logger.info("Provider models refresh disabled (interval <= 0)")
|
||||
return
|
||||
|
||||
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
|
||||
if callable(upstreams_provider):
|
||||
return upstreams_provider()
|
||||
return upstreams_provider
|
||||
|
||||
while True:
|
||||
try:
|
||||
for upstream in upstreams:
|
||||
for upstream in _resolve_upstreams():
|
||||
try:
|
||||
await upstream.refresh_models_cache()
|
||||
except Exception as e:
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
"""Regression tests for the periodic upstream models refresh loop."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
|
||||
|
||||
|
||||
class _FakeUpstream:
|
||||
"""Minimal stand-in for BaseUpstreamProvider used by the refresh loop.
|
||||
|
||||
Only ``base_url`` (for error logging) and ``refresh_models_cache`` (the call
|
||||
under test) are exercised; everything else stays unused.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str) -> None:
|
||||
self.base_url = f"http://{name}"
|
||||
self.refresh_models_cache = AsyncMock()
|
||||
|
||||
|
||||
def _make_fake_upstream(name: str) -> BaseUpstreamProvider:
|
||||
# The loop only uses duck-typed attributes — cast keeps the test type-clean
|
||||
# without dragging in BaseUpstreamProvider's full constructor.
|
||||
return cast(BaseUpstreamProvider, _FakeUpstream(name))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_loop_picks_up_providers_added_after_startup(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""If a provider is added after the loop starts (e.g. via reinitialize_upstreams),
|
||||
the next loop iteration must refresh it. Previously the loop captured the upstream
|
||||
list at startup and missed any later additions."""
|
||||
from routstr.core.settings import settings as global_settings
|
||||
from routstr.upstream.helpers import refresh_upstreams_models_periodically
|
||||
|
||||
# Tight interval so the test finishes quickly.
|
||||
monkeypatch.setattr(
|
||||
global_settings, "models_refresh_interval_seconds", 1, raising=False
|
||||
)
|
||||
|
||||
initial_upstream = _make_fake_upstream("initial")
|
||||
live_list: list[BaseUpstreamProvider] = [initial_upstream]
|
||||
|
||||
# Stub out the post-iteration sats-pricing refresh so the loop body has no DB deps.
|
||||
async def _noop_pricing_refresh() -> None: # pragma: no cover - trivial stub
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
"routstr.payment.models._update_sats_pricing_once",
|
||||
_noop_pricing_refresh,
|
||||
)
|
||||
|
||||
task = asyncio.create_task(
|
||||
refresh_upstreams_models_periodically(lambda: live_list)
|
||||
)
|
||||
|
||||
try:
|
||||
# Wait for the first iteration to refresh the initial upstream.
|
||||
for _ in range(40):
|
||||
if initial_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
assert initial_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
|
||||
"loop did not refresh the initial upstream within the timeout"
|
||||
)
|
||||
|
||||
# Simulate reinitialize_upstreams: replace the live list contents with new
|
||||
# provider instances. The loop must observe the swap on its next tick.
|
||||
new_upstream = _make_fake_upstream("added-after-startup")
|
||||
live_list[:] = [new_upstream]
|
||||
|
||||
for _ in range(60):
|
||||
if new_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
assert new_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
|
||||
"loop did not refresh the upstream added after startup — "
|
||||
"regression: list snapshot captured at startup"
|
||||
)
|
||||
finally:
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_loop_disabled_when_interval_non_positive(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from routstr.core.settings import settings as global_settings
|
||||
from routstr.upstream.helpers import refresh_upstreams_models_periodically
|
||||
|
||||
monkeypatch.setattr(
|
||||
global_settings, "models_refresh_interval_seconds", 0, raising=False
|
||||
)
|
||||
|
||||
upstream = _make_fake_upstream("never-refreshed")
|
||||
|
||||
# Loop must return immediately without ever touching the upstream.
|
||||
await asyncio.wait_for(
|
||||
refresh_upstreams_models_periodically(lambda: [upstream]),
|
||||
timeout=1.0,
|
||||
)
|
||||
upstream.refresh_models_cache.assert_not_awaited() # type: ignore[attr-defined]
|
||||
@@ -51,7 +51,15 @@ async def test_stream_with_id_injection() -> None:
|
||||
base.adjust_payment_for_tokens = AsyncMock(
|
||||
return_value={"total_usd": 0.1, "total_msats": 100}
|
||||
)
|
||||
base.create_session = MagicMock()
|
||||
# create_session() is used as an async context manager whose entered
|
||||
# value exposes an awaitable .get(). Build a mock that behaves that
|
||||
# way so the post-stream cost-chunk emission can run.
|
||||
mock_session = MagicMock()
|
||||
mock_session.get = AsyncMock(return_value=key)
|
||||
mock_ctx = MagicMock()
|
||||
mock_ctx.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_ctx.__aexit__ = AsyncMock(return_value=None)
|
||||
base.create_session = MagicMock(return_value=mock_ctx)
|
||||
|
||||
streaming_response = await provider.handle_streaming_chat_completion(
|
||||
response=mock_response,
|
||||
|
||||
Reference in New Issue
Block a user