Files
routstr-core/routstr/proxy.py
T

1258 lines
48 KiB
Python

import asyncio
import inspect
import json
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
from sqlmodel import select
from .algorithm import create_model_mappings
from .auth import (
ReservationSnapshot,
pay_for_request,
revert_pay_for_request,
validate_bearer_key,
)
from .core import get_logger
from .core.db import (
ApiKey,
AsyncSession,
ModelRow,
UpstreamProviderRow,
create_session,
get_session,
)
from .core.error_scope import (
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
)
from .core.exceptions import UpstreamError
from .core.not_found import build_not_found_response
from .core.settings import settings
from .payment.helpers import (
calculate_discounted_max_cost,
check_token_balance,
create_error_response,
create_upstream_error_response,
get_max_cost_for_model,
)
from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.cooldown import is_cooling_down, record_failure
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
from .upstream.helpers import init_upstreams
from .upstream.model_paths import (
ModelPathSelector,
decode_model_path,
is_openrouter_base_url,
public_model_id,
public_provider_url,
)
from .upstream.request_correction import correct_request, extract_error_message
logger = get_logger(__name__)
MODEL_PATH_HEADER = "x-routstr-model-path"
proxy_router = APIRouter()
_upstreams: list[BaseUpstreamProvider] = []
_provider_map: dict[
str, list[tuple[Model, BaseUpstreamProvider]]
] = {} # All aliases -> sorted [(candidate Model, its Provider)]
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
async def _finish_read_transaction(session: AsyncSession) -> None:
"""Release a read transaction without assuming a particular session mock."""
commit_result = session.commit()
if inspect.isawaitable(commit_result):
await commit_result
async def initialize_upstreams() -> None:
"""Initialize upstream providers from database during application startup."""
global _upstreams
_upstreams = await init_upstreams()
logger.info(f"Initialized {len(_upstreams)} upstream providers")
await refresh_model_maps()
async def reinitialize_upstreams() -> None:
"""Re-initialize upstream providers from database (called after admin changes)."""
global _upstreams
_upstreams = await init_upstreams()
logger.info(
"Re-initialized upstream providers from admin action",
extra={"provider_count": len(_upstreams)},
)
await refresh_model_maps()
def get_upstreams() -> list[BaseUpstreamProvider]:
"""Get the initialized upstream providers.
Returns:
List of upstream provider instances
"""
return _upstreams
def get_candidates(
model_id: str,
) -> list[tuple[Model, BaseUpstreamProvider]] | None:
"""Get the sorted (model, provider) candidate list for a model ID.
Each provider is paired with its own model for the alias, so routing can
forward and bill the candidate that actually serves. Version suffixes
(e.g. ``-20251222``) are stripped as a retry when the exact ID is
unknown, since upstreams may return a specific version of a base model
we track.
"""
if not model_id:
return None
model_id_lower = model_id.lower()
if candidates := _provider_map.get(model_id_lower):
return candidates
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower:
if candidates := _provider_map.get(base_model_id):
return candidates
return None
def _model_ids_match(requested: str, selected: str) -> bool:
if requested.lower() == selected.lower():
return True
return public_model_id(requested).lower() == public_model_id(selected).lower()
def _candidate_for_selector(
selector: ModelPathSelector,
candidates: list[tuple[Model, BaseUpstreamProvider]],
) -> tuple[Model, BaseUpstreamProvider] | None:
"""Resolve a route selector to one candidate.
``candidates`` is ranked by cost, so the first URL match is the cheapest
provider configured against that URL. A selector still carrying a legacy
``provider-id`` keeps pinning that exact provider instead.
"""
for model_obj, upstream in candidates:
if public_provider_url(upstream.base_url) != selector.base_url:
continue
if selector.provider_id is not None and upstream.db_id != selector.provider_id:
continue
return model_obj, upstream
return None
def get_model_instance(model_id: str) -> Model | None:
"""Get the best-ranked Model instance for a model ID."""
candidates = get_candidates(model_id)
return candidates[0][0] if candidates else None
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
"""Get the sorted UpstreamProvider list for a model ID."""
candidates = get_candidates(model_id)
return [provider for _, provider in candidates] if candidates else None
def get_unique_models() -> list[Model]:
"""Get list of unique models (no duplicates from aliases)."""
return list(_unique_models.values())
def _is_tinfoil_attestation_path(path: str) -> bool:
"""Return True for exact Tinfoil attestation routes, with optional slash."""
return path in {
"attestation",
"attestation/",
"tee/attestation",
"tee/attestation/",
}
def _select_unauthenticated_get_upstreams(
path: str, upstreams: list[BaseUpstreamProvider]
) -> list[BaseUpstreamProvider]:
"""Select upstream candidates for unauthenticated GET bypass paths.
Tinfoil attestation endpoints are provider-specific. Trying every enabled
upstream can return an unrelated provider's 404 before Tinfoil is reached,
so route those paths only to Tinfoil providers.
"""
if _is_tinfoil_attestation_path(path):
return [
upstream
for upstream in upstreams
if getattr(upstream, "provider_type", None) == "tinfoil"
]
return upstreams
async def refresh_model_maps() -> None:
"""Refresh global model and provider maps using the cost-based algorithm."""
from sqlalchemy.orm import selectinload
global _provider_map, _unique_models
async with create_session() as session:
# Fetch all providers with their models in a single logical operation
query = select(UpstreamProviderRow).options(
selectinload(UpstreamProviderRow.models) # type: ignore
)
result = await session.exec(query)
provider_rows = result.all()
overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {}
disabled_model_keys: set[tuple[str, int]] = set()
for provider in provider_rows:
if not provider.enabled:
continue
for model in provider.models:
model_key = (model.id.lower(), model.upstream_provider_id)
if model.enabled:
overrides_by_key[model_key] = (model, provider.provider_fee)
else:
disabled_model_keys.add(model_key)
_, _provider_map, _unique_models = create_model_mappings(
upstreams=_upstreams,
overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys,
)
# Keep model-path discovery in sync with admin mutations: disabling or
# deleting a provider must stop advertising its paths immediately rather
# than after the next timed refresh.
from .upstream.model_paths import prune_model_paths_for_inactive_providers
try:
await prune_model_paths_for_inactive_providers()
except Exception as e: # noqa: BLE001 - discovery sync must not break routing
logger.warning(
"Failed to prune model paths for inactive providers",
extra={"error": str(e), "error_type": type(e).__name__},
)
async def refresh_model_maps_periodically() -> None:
"""Background task to refresh model maps every minute."""
import asyncio
while True:
try:
await asyncio.sleep(60)
await refresh_model_maps()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(
"Error refreshing model maps",
extra={"error": str(e), "error_type": type(e).__name__},
)
# Canonical endpoints this proxy will forward, keyed by the path with any
# leading "v1/" and trailing slash removed, mapped to the methods allowed on
# each. The provider credential is attached during forwarding, so endpoint
# permission has to come from this table rather than from the client-supplied
# path: an upstream's key-management, organization, or billing routes live
# under the same origin and must never be reachable through the proxy.
_ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = {
"chat/completions": frozenset({"POST"}),
"completions": frozenset({"POST"}),
"responses": frozenset({"POST"}),
"messages": frozenset({"POST"}),
"embeddings": frozenset({"POST"}),
# TypeSafe System One decision endpoint: POST {state, model, questions}
# -> {answers, usage}. Non-streaming, JSON in/out; billed from the
# response's usage exactly like embeddings.
"systemone": frozenset({"POST"}),
"models": frozenset({"GET"}),
"attestation": frozenset({"GET"}),
"tee/attestation": frozenset({"GET"}),
}
_ALLOWED_METHODS = frozenset({"GET", "POST"})
def _canonical_api_path(path: str) -> str:
"""Reduce a request path to its allowlist key.
OpenAI-style clients reach the same endpoint with or without the ``v1/``
prefix and with or without a trailing slash, so both spellings collapse to
one key. Callers must screen the path with
:func:`_is_ambiguously_spelled_path` first — this function assumes the path
has no dot segments, empty segments, or encoded separators left to resolve.
"""
core = path[:-1] if path.endswith("/") else path
if core.startswith("v1/"):
core = core[len("v1/") :]
return core
def _parse_extra_allowed_endpoints(raw: str) -> dict[str, frozenset[str]]:
"""Parse operator-configured additions to the endpoint allowlist.
Deployments whose provider exposes an endpoint outside the canonical set
opt in explicitly with ``PROXY_EXTRA_ALLOWED_PATHS``, a comma-separated
list of ``METHOD:path`` pairs (e.g. ``POST:v1/rerank,GET:batches``). Every
entry must name one concrete method and one unambiguous path; wildcards
and bare prefixes are deliberately unsupported, so widening the proxy's
reach is always a per-endpoint decision. Malformed entries are dropped
with a warning rather than silently widening or narrowing the surface.
"""
extra: dict[str, frozenset[str]] = {}
for entry in raw.split(","):
entry = entry.strip()
if not entry:
continue
method, separator, endpoint = entry.partition(":")
method = method.strip().upper()
endpoint = endpoint.strip()
if not separator or method not in _ALLOWED_METHODS or not endpoint:
logger.warning(
"Ignoring malformed PROXY_EXTRA_ALLOWED_PATHS entry",
extra={"entry": entry},
)
continue
if _is_ambiguously_spelled_path(endpoint):
logger.warning(
"Ignoring ambiguously spelled PROXY_EXTRA_ALLOWED_PATHS entry",
extra={"entry": entry},
)
continue
if any(character in endpoint for character in "*?["):
# Refuse glob syntax outright. Kept as a literal endpoint name it
# would never match a real request, so the operator would think
# they had widened the proxy when they had not.
logger.warning(
"Ignoring wildcard PROXY_EXTRA_ALLOWED_PATHS entry; "
"list each endpoint explicitly",
extra={"entry": entry},
)
continue
key = _canonical_api_path(endpoint)
extra[key] = extra.get(key, frozenset()) | {method}
return extra
def _is_ambiguously_spelled_path(path: str) -> bool:
"""Reject paths whose spelling could resolve somewhere the allowlist did not.
``{path:path}`` arrives percent-decoded, so a client that sent ``%2e%2e`` or
``%2f`` shows up here as ``..`` / ``/``. Dot segments, backslashes, duplicate
or leading separators, NUL bytes, and any residual encoded separator are
treated as unsafe: they let a caller walk off the canonical API surface (and
onto a sensitive upstream endpoint) even though the literal prefix check
would pass. Reject rather than trying to rewrite the path.
"""
if not path or path != path.strip() or path.startswith("/"):
return True
if "\x00" in path or "\\" in path:
return True
# A single trailing slash is canonical (e.g. "attestation/"); ignore it,
# then no remaining segment may be empty (covers "//") or a dot segment.
core = path[:-1] if path.endswith("/") else path
if any(segment in ("", ".", "..") for segment in core.split("/")):
return True
lowered = path.lower()
return "%2e" in lowered or "%2f" in lowered or "%5c" in lowered
_EXTRA_ALLOWED_ENDPOINTS = _parse_extra_allowed_endpoints(
settings.proxy_extra_allowed_paths
)
def _allowed_methods_for(endpoint: str) -> frozenset[str]:
"""Return the methods allowed on a canonical endpoint, empty if unknown."""
methods = _ALLOWED_ENDPOINTS.get(endpoint, frozenset())
methods |= _EXTRA_ALLOWED_ENDPOINTS.get(endpoint, frozenset())
return methods
def _forwarding_allowed(path: str, method: str) -> bool:
"""Gate which method/path pairs may reach an upstream at all.
The provider credential is attached during forwarding, so an unknown
endpoint must never be forwarded on the caller's say-so. The path is
reduced to its canonical form and looked up in the endpoint table; there is
no prefix match, so a known prefix no longer carries an unknown endpoint
(``v1/organization/api_keys`` is rejected even though ``v1/`` is familiar).
EHBP requests are gated by the same table. Their body is opaque to the
proxy, which is a reason to constrain the destination more tightly, not to
trust the caller's path: the encrypted contract covers the body, never the
endpoint the credential is spent against.
"""
if method not in _ALLOWED_METHODS:
return False
return method in _allowed_methods_for(_canonical_api_path(path))
# Gateway conditions a retry usually clears. 500 is excluded: as likely to be a
# deterministic rejection that fails identically on the next attempt.
_RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5
def _counts_toward_cooldown(status_code: int) -> bool:
"""Provider faults and timeouts only — not client errors or rate limits."""
return status_code >= 500 or status_code == UPSTREAM_ERROR_STATUS
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session)
) -> Response | StreamingResponse:
"""Run proxy setup in a short request session, never across response streaming."""
try:
return await _proxy(request, path, session)
finally:
# FastAPI yield dependencies normally close after the response body is
# sent. Close explicitly so a long stream cannot retain DB resources.
close_result = session.close()
if inspect.isawaitable(close_result):
await close_result
async def _proxy(
request: Request, path: str, session: AsyncSession
) -> Response | StreamingResponse:
# Screen the path before any routing decision: reject ambiguous spellings,
# then require a known API prefix so nothing unknown is forwarded with the
# provider credential attached.
if _is_ambiguously_spelled_path(path):
return build_not_found_response(request, path)
headers = dict(request.headers)
is_ehbp = "ehbp-encapsulated-key" in headers
if not _forwarding_allowed(path, request.method):
return build_not_found_response(request, path)
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
request_body = await request.body()
# EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key
# header and a binary HPKE-sealed body. The proxy cannot parse the body to
# extract the model id, so the SDK sends it in X-Routstr-Model. Forward the
# raw encrypted body to the upstream's /private/ endpoint and stream the
# encrypted response back untouched — the SDK's SecureClient decrypts it.
if is_ehbp:
request_body_dict = {}
model_id = headers.get("x-routstr-model", "")
if not model_id:
return create_error_response(
"invalid_request",
"EHBP request missing X-Routstr-Model header",
400,
request=request,
)
else:
request_body_dict = parse_request_body_json(request_body, path)
if is_responses_api:
model_id = extract_model_from_responses_request(request_body_dict)
else:
model_id = request_body_dict.get("model", "unknown")
# Exact Tinfoil attestation GET routes don't map to models — forward
# without model/cost/auth lookups. Do not prefix-match here: paths such as
# /attestationjunk must continue through normal authentication.
if request.method == "GET" and _is_tinfoil_attestation_path(path):
if MODEL_PATH_HEADER in headers:
return create_error_response(
"unsupported_request",
"Model paths do not apply to attestation",
400,
request=request,
)
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
if not selected_upstreams:
return create_error_response(
"upstream_error",
"No upstream available for unauthenticated GET path",
502,
request=request,
)
last_error_response = None
for i, upstream in enumerate(selected_upstreams):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if (
response.status_code in [424, 502, 503, 429]
and i < len(selected_upstreams) - 1
):
logger.warning(
"Upstream %s returned %s for unauthenticated GET %s, trying next",
upstream.provider_type,
response.status_code,
path,
)
continue
return response
except UpstreamError as e:
logger.warning(
"Upstream %s failed for unauthenticated GET %s: %s",
upstream.provider_type,
path,
e,
)
if i == len(selected_upstreams) - 1:
last_error_response = create_upstream_error_response(e, request)
continue
return last_error_response or create_error_response(
"upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
)
selector: ModelPathSelector | None = None
if MODEL_PATH_HEADER in headers:
selector = decode_model_path(headers[MODEL_PATH_HEADER])
if (
selector is None
or sum(
name.lower() == MODEL_PATH_HEADER for name, _ in request.headers.items()
)
!= 1
):
return create_error_response(
"invalid_request",
f"Malformed {MODEL_PATH_HEADER} header",
400,
request=request,
)
if not isinstance(model_id, str) or not _model_ids_match(
model_id, selector.model_id
):
return create_error_response(
"invalid_request",
f"{MODEL_PATH_HEADER} selects model '{selector.model_id}' but the "
f"request asks for '{model_id}'",
400,
request=request,
)
if "models" in request_body_dict:
return create_error_response(
"invalid_request",
"Model paths cannot be combined with model fallbacks",
400,
request=request,
)
model_id = selector.model_id
candidates = get_candidates(model_id)
if not candidates:
return create_error_response(
"invalid_model", f"Model '{model_id}' not found", 400, request=request
)
if selector is not None:
pinned = _candidate_for_selector(selector, candidates)
if pinned is None:
target = (
f"provider {selector.provider_id}"
if selector.provider_id is not None
else f"'{selector.base_url}'"
)
return create_error_response(
"invalid_model_path",
f"Model '{selector.model_id}' is not routable through {target}",
404,
request=request,
)
# Explicit routes must never enter cross-provider failover.
candidates = [pinned]
if selector.endpoint_tag:
if (
is_ehbp
or not request_body_dict
or not is_openrouter_base_url(pinned[1].base_url)
or _canonical_api_path(path)
not in {"chat/completions", "completions", "responses"}
):
return create_error_response(
"unsupported_request",
"Endpoint pinning requires an OpenRouter completion or Responses JSON request",
400,
request=request,
)
provider_options = request_body_dict.get("provider", {})
if not isinstance(provider_options, dict):
return create_error_response(
"invalid_request",
"provider must be an object",
400,
request=request,
)
request_body_dict = {
**request_body_dict,
"provider": {
**provider_options,
"order": [selector.endpoint_tag],
"allow_fallbacks": False,
},
}
request_body = json.dumps(request_body_dict).encode()
if is_ehbp:
candidates = [
(model, upstream)
for model, upstream in candidates
if upstream.supports_ehbp
]
if not candidates:
return create_error_response(
"unsupported_request",
f"No EHBP-capable provider found for model '{model_id}'",
400,
request=request,
)
# A provider that just failed this model repeatedly is skipped while some
# other candidate can serve it. An explicit route is never rerouted.
if selector is None:
healthy = [
candidate
for candidate in candidates
if not is_cooling_down(candidate[1].base_url, model_id)
]
if healthy:
candidates = healthy
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried.
model_obj = candidates[0][0]
_max_cost_for_model = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
max_cost_for_model = await calculate_discounted_max_cost(
_max_cost_for_model, request_body_dict, model_obj=model_obj
)
check_token_balance(headers, request_body_dict, max_cost_for_model)
if x_cashu := headers.get("x-cashu", None):
last_error = None
for i, (model_obj, upstream) in enumerate(candidates):
try:
if is_ehbp:
if not upstream.supports_ehbp:
logger.warning(
"Upstream %s does not support EHBP for model=%s",
upstream.provider_type,
model_id,
)
continue
return await forward_ehbp_x_cashu_request(
request=request,
x_cashu_token=x_cashu,
path=path,
max_cost_for_model=max_cost_for_model,
model_obj=model_obj,
upstream=upstream,
)
elif is_responses_api:
return await upstream.handle_x_cashu_responses(
request,
x_cashu,
path,
max_cost_for_model,
model_obj,
request_body=request_body,
)
else:
return await upstream.handle_x_cashu(
request,
x_cashu,
path,
max_cost_for_model,
model_obj,
request_body=request_body,
)
except UpstreamError as e:
logger.warning(
"Upstream %s failed (x-cashu) for model=%s: %s",
upstream.provider_type,
model_id,
e,
extra={
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
},
)
if i == len(candidates) - 1:
last_error = e
continue
if last_error is not None:
return create_upstream_error_response(last_error, request)
return create_error_response(
"upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
)
elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(
headers, path, session, auth, max_cost_for_model, model_id
)
else:
if request.method not in ["GET"]:
raise HTTPException(
status_code=401,
detail={
"error": {"type": "invalid_request_error", "code": "unauthorized"}
},
)
logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None
for i, (_, upstream) in enumerate(candidates):
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
if (
response.status_code in [424, 502, 503, 429]
and i < len(candidates) - 1
):
error_message = ""
try:
if hasattr(response, "body"):
body_bytes = response.body
data = json.loads(body_bytes)
if "error" in data:
error_data = data["error"]
if isinstance(error_data, dict):
error_message = error_data.get("message", "")
elif isinstance(error_data, str):
error_message = error_data
except Exception:
pass
await upstream.on_upstream_error_redirect(
response.status_code, error_message
)
logger.warning(
f"Upstream {upstream.provider_type} returned {response.status_code} (GET), trying next provider",
extra={
"status_code": response.status_code,
"upstream": upstream.provider_type,
},
)
continue
return response
except UpstreamError as e:
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
if i == len(candidates) - 1:
last_error_response = create_upstream_error_response(e, request)
continue
return last_error_response or create_error_response(
"upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
)
reservation_snapshot: ReservationSnapshot | None = None
if is_ehbp or request_body_dict:
reservation_snapshot = await pay_for_request(key, max_cost_for_model, session)
# pay_for_request refreshes the key after committing the reservation.
# End that read transaction before waiting on upstream response headers.
await _finish_read_transaction(session)
# Tracks request params already removed in response to upstream rejections,
# shared across providers so a stripped param stays stripped on failover and
# the reactive retry can never loop unboundedly.
already_stripped: set[str] = set()
for i, (model_obj, upstream) in enumerate(candidates):
if i > 0 and request_body_dict:
# The reservation was sized to the previous candidate's envelope;
# settlement bills the serving candidate, so a pricier fallback
# must be re-reserved at its own max cost before it is tried. A
# candidate whose envelope the key cannot cover is rejected, just
# as it would be had it been ranked first.
candidate_max = await get_max_cost_for_model(
model=model_id, session=session, model_obj=model_obj
)
candidate_max = await calculate_discounted_max_cost(
candidate_max, request_body_dict, model_obj=model_obj
)
if candidate_max > max_cost_for_model:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
try:
reservation_snapshot = await pay_for_request(
key, candidate_max, session
)
except HTTPException:
if i == len(candidates) - 1:
raise
reservation_snapshot = await pay_for_request(
key, max_cost_for_model, session
)
await _finish_read_transaction(session)
continue
await _finish_read_transaction(session)
max_cost_for_model = candidate_max
retries_left = settings.upstream_5xx_retry_attempts
retry_index = 0
headers = upstream.prepare_headers(dict(request.headers))
try:
while True:
try:
if is_ehbp:
if not upstream.supports_ehbp:
logger.warning(
"Upstream %s does not support EHBP for model=%s",
upstream.provider_type,
model_id,
)
raise UpstreamError(
f"Provider {upstream.provider_type} does not support EHBP",
status_code=400,
)
response = await forward_ehbp_request(
request=request,
path=path,
headers=headers,
request_body=request_body,
upstream=upstream,
key=key,
max_cost_for_model=max_cost_for_model,
session=session,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
)
elif is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
reservation_snapshot,
)
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
reservation_snapshot,
)
except UpstreamError as e:
# Only a gateway status the upstream itself answered with:
# re-sending the buffered body cannot double-bill. A 502 this
# proxy invented for a transport error or timeout is not
# retried — that request may already be running upstream.
if (
e.from_upstream_response
and e.status_code in _RETRYABLE_UPSTREAM_5XX
and retries_left > 0
):
retries_left -= 1
retry_index += 1
logger.warning(
"Upstream %s returned %s for model=%s; retrying same "
"upstream (attempt %s, %s retries left)",
upstream.provider_type,
e.status_code,
model_id,
retry_index + 1,
retries_left,
extra={
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
"path": path,
"retries_left": retries_left,
},
)
await asyncio.sleep(
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS * retry_index
)
continue
# Let the outer UpstreamError handler manage failover/revert
raise
except Exception as e:
# Unexpected error (not an upstream failure) — revert and propagate
logger.error(
"Unexpected error in upstream request, reverting payment",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"max_cost_for_model": max_cost_for_model,
},
)
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
raise
# Same-provider recovery must not relax an explicit route.
if response.status_code == 400 and not is_ehbp:
correction = correct_request(
request_body,
extract_error_message(response),
already_stripped,
)
if correction is not None and selector is not None:
corrected_body = json.loads(correction.body)
if any(
corrected_body.get(field) != request_body_dict.get(field)
for field in ("model", "provider")
):
correction = None
if correction is not None:
request_body, bad_param = correction.body, correction.label
already_stripped.add(bad_param)
logger.warning(
"Upstream %s rejected param '%s' for model=%s; "
"stripping and retrying same upstream",
upstream.provider_type,
bad_param,
model_id,
extra={
"provider": upstream.provider_type,
"model": model_id,
"stripped_param": bad_param,
"path": path,
},
)
continue
break
if response.status_code != 200:
if _counts_toward_cooldown(response.status_code):
record_failure(upstream.base_url, model_id)
# 424 is an upstream failure re-reported by error_scope.
# 502/503 are upstream errors, 429 rate limits.
should_retry = response.status_code in [
424,
502,
503,
429,
400,
401,
403,
404,
]
if should_retry and i < len(candidates) - 1:
error_message = ""
try:
if hasattr(response, "body"):
body_bytes = response.body
data = json.loads(body_bytes)
if "error" in data:
error_data = data["error"]
if isinstance(error_data, dict):
error_message = error_data.get("message", "")
elif isinstance(error_data, str):
error_message = error_data
except Exception:
pass
await upstream.on_upstream_error_redirect(
response.status_code, error_message
)
logger.warning(
"Upstream %s returned %s for model=%s, trying next provider",
upstream.provider_type,
response.status_code,
model_id,
extra={
"status_code": response.status_code,
"provider": upstream.provider_type,
"model": model_id,
},
)
continue
# 4xx error (user error), or other non-retryable error, or last provider failed
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
logger.warning(
"Upstream request failed, revert payment "
"(provider=%s model=%s status=%s path=%s)",
upstream.provider_type,
model_id,
response.status_code,
path,
extra={
"status_code": response.status_code,
"path": path,
"provider": upstream.provider_type,
"model": model_id,
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
"max_cost_for_model": max_cost_for_model,
},
)
return response
return response
except asyncio.CancelledError:
logger.warning(
"Client disconnected mid-request, reverting reservation",
extra={
"path": path,
"model": model_id,
"key_hash": key.hashed_key[:8] + "...",
"max_cost_for_model": max_cost_for_model,
},
)
# The cancellation has been caught, so complete exact cleanup in
# this task before the request-scoped session can be torn down.
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
raise
except UpstreamError as e:
if _counts_toward_cooldown(e.status_code):
record_failure(upstream.base_url, model_id)
logger.warning(
"Upstream %s failed for model=%s: %s",
upstream.provider_type,
model_id,
e,
extra={
"provider": upstream.provider_type,
"model": model_id,
"status_code": e.status_code,
"retry": i < len(candidates) - 1,
},
)
# If this was the last provider
if i == len(candidates) - 1:
await revert_pay_for_request(
key, session, max_cost_for_model, reservation_snapshot
)
return create_upstream_error_response(e, request)
# Otherwise loop continues to next provider
continue
# Should not be reached given logic above
return create_error_response(
"upstream_error",
"All upstreams failed",
UPSTREAM_ERROR_STATUS,
request=request,
code=UPSTREAM_UNAVAILABLE,
error_scope=ERROR_SCOPE_UPSTREAM,
)
async def get_bearer_token_key(
headers: dict,
path: str,
session: AsyncSession,
auth: str,
min_cost: int = 0,
model_id: str = "unknown",
) -> ApiKey:
"""Handle bearer token authentication proxy requests."""
parts = auth.split()
bearer_key = parts[1] if len(parts) > 1 and parts[0].lower() == "bearer" else ""
refund_address = headers.get("Refund-LNURL", None)
key_expiry_time = headers.get("Key-Expiry-Time", None)
logger.debug(
"Processing bearer token",
extra={
"path": path,
"has_refund_address": bool(refund_address),
"has_expiry_time": bool(key_expiry_time),
"bearer_key_preview": bearer_key[:20] + "..."
if len(bearer_key) > 20
else bearer_key,
"min_cost": min_cost,
},
)
# Validate key_expiry_time header
if key_expiry_time:
try:
key_expiry_time = int(key_expiry_time) # type: ignore
logger.debug(
"Key expiry time validated",
extra={"expiry_time": key_expiry_time, "path": path},
)
except ValueError:
logger.error(
"Invalid Key-Expiry-Time header",
extra={"key_expiry_time": key_expiry_time, "path": path},
)
raise HTTPException(
status_code=400,
detail="Invalid Key-Expiry-Time: must be a valid Unix timestamp",
)
if not refund_address:
logger.error(
"Missing Refund-LNURL header with Key-Expiry-Time",
extra={"path": path, "expiry_time": key_expiry_time},
)
raise HTTPException(
status_code=400,
detail="Error: Refund-LNURL header required when using Key-Expiry-Time",
)
else:
key_expiry_time = None
try:
key = await validate_bearer_key(
bearer_key,
session,
refund_address,
key_expiry_time, # type: ignore
min_cost=min_cost,
)
logger.info(
"Bearer token validated successfully",
extra={
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"key_balance": key.balance,
},
)
return key
except HTTPException as error:
detail: dict[str, Any] = error.detail if isinstance(error.detail, dict) else {}
raw_error = detail.get("error")
error_info = raw_error if isinstance(raw_error, dict) else {}
logger.warning(
"Bearer token rejected",
extra={
"status_code": error.status_code,
"error_code": error_info.get("code"),
"path": path,
"model_id": model_id,
"required_msat": min_cost,
},
)
raise
except Exception as error:
logger.exception(
"Bearer token validation failed",
extra={
"error_type": type(error).__name__,
"path": path,
"model_id": model_id,
"required_msat": min_cost,
},
)
raise
def extract_model_from_responses_request(request_body_dict: dict[str, Any]) -> str:
if model := request_body_dict.get("model"):
return model
if input_data := request_body_dict.get("input"):
if isinstance(input_data, dict) and (model := input_data.get("model")):
return model
if request_body_dict.get("messages"):
return "unknown"
logger.warning(
"No model found in Responses API request",
extra={"body_keys": list(request_body_dict.keys())},
)
return "unknown"
def parse_request_body_json(request_body: bytes, path: str) -> dict[str, Any]:
request_body_dict = {}
if request_body:
try:
request_body_dict = json.loads(request_body)
if "max_tokens" in request_body_dict:
max_tokens_value = request_body_dict["max_tokens"]
if isinstance(max_tokens_value, int):
pass
else:
raise HTTPException(
status_code=400,
detail={"error": "max_tokens must be an integer"},
)
logger.debug(
"Request body parsed",
extra={
"path": path,
"body_keys": list(request_body_dict.keys()),
"model": request_body_dict.get("model", "not_specified"),
},
)
except json.JSONDecodeError as e:
logger.error(
"Invalid JSON in request body",
extra={
"error": str(e),
"path": path,
"body_preview": request_body[:200].decode(errors="ignore")
if request_body
else "empty",
},
)
raise HTTPException(
status_code=400,
detail={
"error": {"type": "invalid_request_error", "code": "invalid_json"}
},
)
return request_body_dict