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 e998287b..40c79637 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -16,6 +16,7 @@ from .core.db import ( create_session, get_session, ) +from .core.exceptions import UpstreamError from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, @@ -31,7 +32,9 @@ 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) @@ -68,8 +71,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()) @@ -154,8 +157,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}'", @@ -163,6 +166,10 @@ async def proxy( request=request, ) + # todo figure out cost calculation since fallback provider is usually not the same price + # 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 ) @@ -172,14 +179,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) @@ -194,56 +218,152 @@ 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)) + response = await upstream.forward_get_request(request, path, headers) + + if response.status_code in [502, 429] and i < len(upstreams) - 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(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: + # Check if we should retry (502 Upstream Error or 429 Rate Limit) + should_retry = response.status_code in [502, 429, 400, 401, 403, 404] + if should_retry and i < len(upstreams) - 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 - return response + await upstream.on_upstream_error_redirect( + response.status_code, error_message + ) + + logger.warning( + f"Upstream {upstream.provider_type} returned {response.status_code}, trying next provider", + extra={ + "status_code": response.status_code, + "upstream": upstream.provider_type, + }, + ) + 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) + 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 + + 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..9dc7ae7b 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 @@ -340,6 +341,20 @@ class BaseUpstreamProvider: message = preview[:500] return message, upstream_code + async def on_upstream_error_redirect( + self, status_code: int, error_message: str + ) -> None: + """Hook called when the proxy redirects to another provider due to an error. + + Subclasses can implement this to perform actions like disabling the provider + if it's out of balance. + + Args: + status_code: The HTTP status code returned by the upstream + error_message: The error message extracted from the upstream response + """ + pass + async def map_upstream_error_response( self, request: Request, path: str, upstream_response: httpx.Response ) -> Response: @@ -1148,6 +1163,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 +1261,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1265,9 +1291,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 +1404,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 +1479,9 @@ class BaseUpstreamProvider: background=background_tasks, ) + except UpstreamError: + raise + except httpx.RequestError as exc: await client.aclose() error_type = type(exc).__name__ @@ -1474,9 +1509,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() diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 0b4dfb08..01a5decb 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -196,6 +196,37 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) return [] + async def on_upstream_error_redirect( + self, status_code: int, error_message: str + ) -> None: + if "insufficient balance" in error_message.lower(): + logger.warning( + f"Disabling PPQ.AI provider ({self.base_url}) due to insufficient balance", + extra={"error": error_message}, + ) + from sqlmodel import select + + from ..core.db import UpstreamProviderRow, create_session + + async with create_session() as session: + statement = select(UpstreamProviderRow).where( + UpstreamProviderRow.base_url == self.base_url + and UpstreamProviderRow.api_key == self.api_key + ) + result = await session.exec(statement) + provider = result.first() + + if provider: + provider.enabled = False + session.add(provider) + await session.commit() + + # Trigger re-initialization of providers + # Import here to avoid circular dependency + from ..proxy import reinitialize_upstreams + + await reinitialize_upstreams() + async def create_account(self) -> dict[str, object]: """Create a new PPQ.AI account. diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index d367819f..fd7a837f 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -10,7 +10,6 @@ os.environ["UPSTREAM_API_KEY"] = "test" from routstr.algorithm import ( # noqa: E402 calculate_model_cost_score, get_provider_penalty, - should_prefer_model, ) from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 @@ -100,100 +99,3 @@ def test_get_provider_penalty_openrouter() -> None: provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1") penalty = get_provider_penalty(provider) assert penalty == 1.001 - - -def test_should_prefer_model_cheaper_wins() -> None: - """Test that cheaper model is preferred.""" - cheap_model = create_test_model("cheap", prompt_price=0.001, completion_price=0.002) - expensive_model = create_test_model( - "expensive", prompt_price=0.03, completion_price=0.06 - ) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # Cheaper model should win - assert should_prefer_model( - cheap_model, provider1, expensive_model, provider2, "test-alias" - ) - - # More expensive model should not win - assert not should_prefer_model( - expensive_model, provider2, cheap_model, provider1, "test-alias" - ) - - -def test_should_prefer_model_exact_match_wins() -> None: - """Test that exact alias match beats cheaper price.""" - # Make model IDs match the alias differently - exact_match = create_test_model( - "test-model", prompt_price=0.03, completion_price=0.06 - ) - no_match = create_test_model( - "other-model", prompt_price=0.001, completion_price=0.002 - ) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # Exact match should win even though it's more expensive - assert should_prefer_model( - exact_match, provider1, no_match, provider2, "test-model" - ) - - -def test_should_prefer_model_openrouter_slight_penalty() -> None: - """Test that OpenRouter has slight penalty compared to other providers.""" - model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002) - model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002) - - regular_provider = create_test_provider("regular", "http://provider.com") - openrouter_provider = create_test_provider( - "openrouter", "https://openrouter.ai/api/v1" - ) - - # Regular provider should be preferred over OpenRouter at same cost - assert should_prefer_model( - model1, regular_provider, model2, openrouter_provider, "test-alias" - ) - - # OpenRouter should not replace regular provider at same cost - assert not should_prefer_model( - model2, openrouter_provider, model1, regular_provider, "test-alias" - ) - - -def test_should_prefer_model_openrouter_can_win_if_cheaper() -> None: - """Test that OpenRouter can still win if significantly cheaper.""" - cheap_model = create_test_model( - "cheap", prompt_price=0.0001, completion_price=0.0002 - ) - expensive_model = create_test_model( - "expensive", prompt_price=0.03, completion_price=0.06 - ) - - regular_provider = create_test_provider("regular", "http://provider.com") - openrouter_provider = create_test_provider( - "openrouter", "https://openrouter.ai/api/v1" - ) - - # OpenRouter should win if it's much cheaper (even with penalty) - assert should_prefer_model( - cheap_model, - openrouter_provider, - expensive_model, - regular_provider, - "test-alias", - ) - - -def test_should_prefer_model_same_cost_first_wins() -> None: - """Test that when costs are identical, current model is kept.""" - model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002) - model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002) - - provider1 = create_test_provider("provider1") - provider2 = create_test_provider("provider2") - - # When costs are equal, should not replace - assert not should_prefer_model(model2, provider2, model1, provider1, "test-alias")