From e7f4c98475f63dd10254d9ae62bd79c681bea8a6 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Wed, 7 Jan 2026 17:39:05 +0100 Subject: [PATCH] ranked provider fallback on upstream errors --- routstr/algorithm.py | 148 +++++++++++-------------------- routstr/core/exceptions.py | 9 ++ routstr/proxy.py | 173 +++++++++++++++++++++++++------------ routstr/upstream/base.py | 31 +++++-- 4 files changed, 203 insertions(+), 158 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 537f07f5..7dffcae1 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -84,93 +84,26 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: return penalty -def should_prefer_model( - candidate_model: "Model", - candidate_provider: "BaseUpstreamProvider", - current_model: "Model", - current_provider: "BaseUpstreamProvider", - alias: str, -) -> bool: - """Determine if candidate model should replace current model for an alias. - - This is the core decision function for model prioritization. It considers: - 1. Alias matching quality (exact match vs. canonical slug match) - 2. Model cost (lower is better) - 3. Provider penalties (e.g., slight preference against OpenRouter) - - Args: - candidate_model: The new model being considered - candidate_provider: Provider offering the candidate model - current_model: The currently selected model for this alias - current_provider: Provider offering the current model - alias: The model alias being mapped - - Returns: - True if candidate should replace current, False otherwise - """ - - def get_base_model_id(model_id: str) -> str: - """Get base model ID by removing provider prefix.""" - return model_id.split("/", 1)[1] if "/" in model_id else model_id - - def alias_priority(model: "Model") -> int: - """Rank how strong the mapping of alias->model is. - - Highest priority when alias exactly equals the model ID without provider prefix. - Next when alias equals canonical slug without prefix. Otherwise lowest. - """ - model_base = get_base_model_id(model.id) - if model_base == alias: - return 3 - if model.canonical_slug: - canonical_base = get_base_model_id(model.canonical_slug) - if canonical_base == alias: - return 2 - return 1 - - candidate_alias_priority = alias_priority(candidate_model) - current_alias_priority = alias_priority(current_model) - - # If candidate has better alias match, prefer it regardless of cost - if candidate_alias_priority > current_alias_priority: - return True - - # If current has better alias match, keep it regardless of cost - if current_alias_priority > candidate_alias_priority: - return False - - # Same alias priority - compare costs - candidate_cost = calculate_model_cost_score(candidate_model) - current_cost = calculate_model_cost_score(current_model) - - # Apply provider penalties - candidate_adjusted = candidate_cost * get_provider_penalty(candidate_provider) - current_adjusted = current_cost * get_provider_penalty(current_provider) - - # Prefer lower adjusted cost - should_replace = candidate_adjusted < current_adjusted - - return should_replace - - def create_model_mappings( upstreams: list["BaseUpstreamProvider"], overrides_by_id: dict[str, tuple], disabled_model_ids: set[str], -) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]: +) -> tuple[ + dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"] +]: """Create optimal model mappings based on cost and provider preferences. This is the main entry point for the algorithm. It processes all upstream providers and creates three mappings based on cost optimization: 1. model_instances: alias -> Model (all model aliases mapped to their Model objects) - 2. provider_map: alias -> UpstreamProvider (which provider to use for each alias) + 2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias) 3. unique_models: base_id -> Model (unique models without provider prefixes) The algorithm: - Processes non-OpenRouter providers first (they're typically cheaper) - Then processes OpenRouter models (they can still win if cheaper) - - For each model alias, uses should_prefer_model() to select the best provider + - For each model alias, collects all candidates and sorts them by priority and cost. Args: upstreams: List of all upstream provider instances @@ -183,8 +116,7 @@ def create_model_mappings( from .payment.models import _row_to_model from .upstream.helpers import resolve_model_alias - model_instances: dict[str, "Model"] = {} - provider_map: dict[str, "BaseUpstreamProvider"] = {} + candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} unique_models: dict[str, "Model"] = {} # Separate OpenRouter from other providers @@ -202,24 +134,14 @@ def create_model_mappings( """Get base model ID by removing provider prefix.""" return model_id.split("/", 1)[1] if "/" in model_id else model_id - def _maybe_set_alias( + def _add_candidate( alias: str, model: "Model", provider: "BaseUpstreamProvider" ) -> None: - """Set alias to model/provider if not set or if new model is preferred.""" + """Add candidate model/provider for an alias.""" alias_lower = alias.lower() - existing_model = model_instances.get(alias_lower) - if not existing_model: - # No existing mapping, set it - model_instances[alias_lower] = model - provider_map[alias_lower] = provider - else: - # Check if candidate should replace existing - existing_provider = provider_map[alias_lower] - if should_prefer_model( - model, provider, existing_model, existing_provider, alias - ): - model_instances[alias_lower] = model - provider_map[alias_lower] = provider + if alias_lower not in candidates: + candidates[alias_lower] = [] + candidates[alias_lower].append((model, provider)) def process_provider_models( upstream: "BaseUpstreamProvider", is_openrouter: bool = False @@ -266,21 +188,55 @@ def create_model_mappings( # Try to set each alias for alias in aliases: - _maybe_set_alias(alias, model_to_use, upstream) + _add_candidate(alias, model_to_use, upstream) - # Process non-OpenRouter providers first (they're typically cheaper) + # Process non-OpenRouter providers first for upstream in other_upstreams: process_provider_models(upstream, is_openrouter=False) - # Process OpenRouter last - models only win if they're cheaper or better matched + # Process OpenRouter last if openrouter: process_provider_models(openrouter, is_openrouter=True) - # Log provider distribution + # Sort candidates and build final maps + model_instances: dict[str, "Model"] = {} + provider_map: dict[str, list["BaseUpstreamProvider"]] = {} + + def alias_priority(model: "Model", alias: str) -> int: + """Rank how strong the mapping of alias->model is.""" + model_base = get_base_model_id(model.id) + if model_base == alias: + return 3 + if model.canonical_slug: + canonical_base = get_base_model_id(model.canonical_slug) + if canonical_base == alias: + return 2 + return 1 + + for alias, items in candidates.items(): + # Sort key: (priority DESC, cost ASC) + # Using negative cost for DESC sort overall to keep high priority first + def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]: + model, provider = item + priority = alias_priority(model, alias) + cost = calculate_model_cost_score(model) + penalty = get_provider_penalty(provider) + adjusted_cost = cost * penalty + return (priority, -adjusted_cost) + + items.sort(key=sort_key, reverse=True) + + best_model, best_provider = items[0] + model_instances[alias] = best_model + provider_map[alias] = [p for _, p in items] + + # Log provider distribution (using top provider for stats) provider_counts: dict[str, int] = {} - for provider in provider_map.values(): - provider_name = getattr(provider, "upstream_name", "unknown") - provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 + for providers in provider_map.values(): + if providers: + provider = providers[0] + provider_name = getattr(provider, "upstream_name", "unknown") + provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 logger.debug( f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)", diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index e74d3e05..64111047 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -6,6 +6,15 @@ from .logging import get_logger logger = get_logger(__name__) +class UpstreamError(Exception): + """Exception raised when an upstream provider fails.""" + + def __init__(self, message: str, status_code: int = 502): + self.message = message + self.status_code = status_code + super().__init__(message) + + async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse: """Handle HTTP exceptions and include request ID in response.""" request_id = getattr(request.state, "request_id", "unknown") diff --git a/routstr/proxy.py b/routstr/proxy.py index 1d5aaa96..b77e3430 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -16,7 +16,11 @@ from .core.db import ( create_session, get_session, ) +<<<<<<< Current (Your changes) from .core.settings import settings +======= +from .core.exceptions import UpstreamError +>>>>>>> Incoming (Background Agent changes) from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, @@ -33,7 +37,7 @@ proxy_router = APIRouter() _upstreams: list[BaseUpstreamProvider] = [] _model_instances: dict[str, Model] = {} # All aliases -> Model -_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider +_provider_map: dict[str, list[BaseUpstreamProvider]] = {} # All aliases -> List[Provider] _unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates) @@ -70,8 +74,8 @@ def get_model_instance(model_id: str) -> Model | None: return _model_instances.get(model_id.lower()) -def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: - """Get UpstreamProvider for model ID from global cache.""" +def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None: + """Get UpstreamProvider list for model ID from global cache.""" return _provider_map.get(model_id.lower()) @@ -172,8 +176,8 @@ async def proxy( "invalid_model", f"Model '{model_id}' not found", 400, request=request ) - upstream = get_provider_for_model(model_id) - if not upstream: + upstreams = get_provider_for_model(model_id) + if not upstreams: return create_error_response( "invalid_model", f"No provider found for model '{model_id}'", @@ -181,6 +185,9 @@ async def proxy( request=request, ) + # Use first provider for initial checks/cost calculation + primary_upstream = upstreams[0] + _max_cost_for_model = await get_max_cost_for_model( model=model_id, session=session, model_obj=model_obj ) @@ -190,14 +197,31 @@ async def proxy( check_token_balance(headers, request_body_dict, max_cost_for_model) if x_cashu := headers.get("x-cashu", None): - if is_responses_api: - return await upstream.handle_x_cashu_responses( - request, x_cashu, path, max_cost_for_model, model_obj - ) - else: - return await upstream.handle_x_cashu( - request, x_cashu, path, max_cost_for_model, model_obj - ) + last_error = None + for i, upstream in enumerate(upstreams): + try: + if is_responses_api: + return await upstream.handle_x_cashu_responses( + request, x_cashu, path, max_cost_for_model, model_obj + ) + else: + return await upstream.handle_x_cashu( + request, x_cashu, path, max_cost_for_model, model_obj + ) + except UpstreamError as e: + logger.warning( + f"Upstream {upstream.provider_type} failed (x-cashu): {e}" + ) + if i == len(upstreams) - 1: + last_error = e + continue + + return create_error_response( + "upstream_error", + str(last_error) if last_error else "All upstreams failed", + 502, + request=request, + ) elif auth := headers.get("authorization", None): key = await get_bearer_token_key(headers, path, session, auth) @@ -212,56 +236,93 @@ async def proxy( ) logger.debug("Processing unauthenticated GET request", extra={"path": path}) - headers = upstream.prepare_headers(dict(request.headers)) - return await upstream.forward_get_request(request, path, headers) + + last_error_response = None + for i, upstream in enumerate(upstreams): + try: + headers = upstream.prepare_headers(dict(request.headers)) + return await upstream.forward_get_request(request, path, headers) + except UpstreamError as e: + logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}") + if i == len(upstreams) - 1: + last_error_response = create_error_response( + "upstream_error", str(e), 502, request=request + ) + continue + return last_error_response or create_error_response( + "upstream_error", "All upstreams failed", 502, request=request + ) if request_body_dict: await pay_for_request(key, max_cost_for_model, session) - headers = upstream.prepare_headers(dict(request.headers)) + for i, upstream in enumerate(upstreams): + headers = upstream.prepare_headers(dict(request.headers)) - if is_responses_api: - response = await upstream.forward_responses_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, - ) - else: - response = await upstream.forward_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, - ) + try: + if is_responses_api: + response = await upstream.forward_responses_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + else: + response = await upstream.forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) - if response.status_code != 200: - await revert_pay_for_request(key, session, max_cost_for_model) - logger.warning( - "Upstream request failed, revert payment", - extra={ - "status_code": response.status_code, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "key_balance": key.balance, - "max_cost_for_model": max_cost_for_model, - "upstream_headers": response.headers - if hasattr(response, "headers") - else None, - }, - ) - # Return the mapped error response generated earlier rather than masking with 502 - return response + if response.status_code != 200: + # 4xx error (user error), or other non-retryable error + await revert_pay_for_request(key, session, max_cost_for_model) + logger.warning( + "Upstream request failed, revert payment", + extra={ + "status_code": response.status_code, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "key_balance": key.balance, + "max_cost_for_model": max_cost_for_model, + "upstream_headers": response.headers + if hasattr(response, "headers") + else None, + }, + ) + return response - return response + return response + + except UpstreamError as e: + logger.warning( + f"Upstream {upstream.provider_type} failed: {e}", + extra={"retry": i < len(upstreams) - 1}, + ) + + # If this was the last provider + if i == len(upstreams) - 1: + await revert_pay_for_request(key, session, max_cost_for_model) + return create_error_response( + "upstream_error", str(e), 502, request=request + ) + + # Otherwise loop continues to next provider + continue + + # Should not be reached given logic above + return create_error_response( + "upstream_error", "All upstreams failed", 502, request=request + ) async def get_bearer_token_key( diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 1c42da8c..48349489 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -15,6 +15,7 @@ from pydantic import BaseModel from ..auth import adjust_payment_for_tokens from ..core import get_logger from ..core.db import ApiKey, AsyncSession, create_session +from ..core.exceptions import UpstreamError if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -1148,6 +1149,14 @@ class BaseUpstreamProvider: ) if response.status_code != 200: + if response.status_code >= 500: + await response.aclose() + await client.aclose() + raise UpstreamError( + f"Upstream returned status {response.status_code}", + status_code=response.status_code, + ) + try: mapped_error = await self.map_upstream_error_response( request, path, response @@ -1238,6 +1247,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1265,9 +1277,7 @@ class BaseUpstreamProvider: else: error_message = f"Error connecting to upstream service: {error_type}" - return create_error_response( - "upstream_error", error_message, 502, request=request - ) + raise UpstreamError(error_message, status_code=502) except Exception as exc: await client.aclose() @@ -1380,6 +1390,14 @@ class BaseUpstreamProvider: ) if response.status_code != 200: + if response.status_code >= 500: + await response.aclose() + await client.aclose() + raise UpstreamError( + f"Upstream returned status {response.status_code}", + status_code=response.status_code, + ) + try: mapped_error = await self.map_upstream_error_response( request, path, response @@ -1447,6 +1465,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1474,9 +1495,7 @@ class BaseUpstreamProvider: else: error_message = f"Error connecting to upstream service: {error_type}" - return create_error_response( - "upstream_error", error_message, 502, request=request - ) + raise UpstreamError(error_message, status_code=502) except Exception as exc: await client.aclose()