From 98aecb08f916db3e1b468c9c56ba423e540f7a39 Mon Sep 17 00:00:00 2001 From: Evan Yang Date: Tue, 10 Feb 2026 01:04:08 +0800 Subject: [PATCH 1/6] fix(admin): filter provider remote models by id --- routstr/core/admin.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index c0ccbe93..0fb9c1f4 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -717,9 +717,7 @@ async def get_provider_models(provider_id: int) -> dict[str, object]: ) db_model_ids = {model.id for model in db_models} - filtered_remote_models = [ - m for m in upstream_models if m.name not in db_model_ids - ] + filtered_remote_models = [m for m in upstream_models if m.id not in db_model_ids] return { "provider": { From 6ebe73f2f77179bc1bb8ad9503650a7f0d32f391 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 10 Feb 2026 07:25:13 +0000 Subject: [PATCH 2/6] Fix Azure Kimi routing and DB override model mapping --- routstr/algorithm.py | 70 ++++++++++++++++++++++ routstr/upstream/azure.py | 113 ++++++++++++++++++++++++++++++++---- routstr/upstream/helpers.py | 3 + 3 files changed, 174 insertions(+), 12 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 7dffcae1..845bca57 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -118,6 +118,13 @@ def create_model_mappings( candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} unique_models: dict[str, "Model"] = {} + seen_model_provider: set[tuple[str, str]] = set() + + providers_by_db_id: dict[int, "BaseUpstreamProvider"] = {} + for upstream in upstreams: + db_id = getattr(upstream, "db_id", None) + if isinstance(db_id, int): + providers_by_db_id[db_id] = upstream # Separate OpenRouter from other providers openrouter: "BaseUpstreamProvider" | None = None @@ -148,6 +155,9 @@ def create_model_mappings( ) -> None: """Process all models from a given provider.""" upstream_prefix = getattr(upstream, "upstream_name", None) + provider_key = getattr(upstream, "provider_type", "") or getattr( + upstream, "base_url", "" + ) for model in upstream.get_cached_models(): if not model.enabled or model.id in disabled_model_ids: @@ -189,6 +199,7 @@ def create_model_mappings( # Try to set each alias for alias in aliases: _add_candidate(alias, model_to_use, upstream) + seen_model_provider.add((model_to_use.id.lower(), provider_key.lower())) # Process non-OpenRouter providers first for upstream in other_upstreams: @@ -198,6 +209,65 @@ def create_model_mappings( if openrouter: process_provider_models(openrouter, is_openrouter=True) + # Include enabled DB overrides even when provider discovery misses models. + # This is important for deployment-based providers like Azure. + for model_id, override_data in overrides_by_id.items(): + if model_id in disabled_model_ids: + continue + try: + override_row, provider_fee = override_data + upstream_provider_id = getattr(override_row, "upstream_provider_id", None) + if not isinstance(upstream_provider_id, int): + continue + + upstream = providers_by_db_id.get(upstream_provider_id) + if upstream is None: + continue + + provider_key = getattr(upstream, "provider_type", "") or getattr( + upstream, "base_url", "" + ) + dedupe_key = (model_id.lower(), provider_key.lower()) + if dedupe_key in seen_model_provider: + continue + + model_to_use = _row_to_model( + override_row, apply_provider_fee=True, provider_fee=provider_fee + ) + if not model_to_use.enabled: + continue + + base_id = get_base_model_id(model_to_use.id) + is_openrouter = ( + getattr(upstream, "base_url", "") == "https://openrouter.ai/api/v1" + ) + if not is_openrouter or base_id not in unique_models: + unique_model = model_to_use.copy( + update={ + "id": base_id, + "upstream_provider_id": upstream.provider_type, + } + ) + unique_models[base_id] = unique_model + + aliases = resolve_model_alias( + model_to_use.id, + model_to_use.canonical_slug, + alias_ids=model_to_use.alias_ids, + ) + upstream_prefix = getattr(upstream, "upstream_name", None) + if upstream_prefix and "/" not in model_to_use.id: + prefixed_id = f"{upstream_prefix}/{model_to_use.id}" + if prefixed_id not in aliases: + aliases.append(prefixed_id) + + for alias in aliases: + _add_candidate(alias, model_to_use, upstream) + seen_model_provider.add(dedupe_key) + except Exception: + # Keep model map creation resilient to malformed overrides. + continue + # Sort candidates and build final maps model_instances: dict[str, "Model"] = {} provider_map: dict[str, list["BaseUpstreamProvider"]] = {} diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index b6240fbd..b9293503 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -1,9 +1,14 @@ from typing import TYPE_CHECKING, Mapping +from fastapi import Request +from fastapi.responses import Response, StreamingResponse + from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..core.db import UpstreamProviderRow + from ..auth import ApiKey + from ..core.db import AsyncSession, UpstreamProviderRow + from ..payment.models import Model class AzureUpstreamProvider(BaseUpstreamProvider): @@ -58,19 +63,103 @@ class AzureUpstreamProvider(BaseUpstreamProvider): "platform_url": cls.platform_url, } + def prepare_headers(self, request_headers: dict) -> dict: + """Prepare headers for Azure OpenAI, adding api-key.""" + headers = super().prepare_headers(request_headers) + if self.api_key: + headers["api-key"] = self.api_key + headers.pop("Authorization", None) + headers.pop("authorization", None) + return headers + def prepare_params( self, path: str, query_params: Mapping[str, str] | None ) -> Mapping[str, str]: - """Prepare query parameters for Azure OpenAI, adding API version. - - Args: - path: Request path - query_params: Original query parameters from the client - - Returns: - Query parameters dict with Azure API version added for chat completions - """ + """Prepare query parameters for Azure OpenAI, adding API version.""" params = dict(query_params or {}) - if path.endswith("chat/completions"): - params["api-version"] = self.api_version + # Ensure we use a valid Azure API version format + # Strip any hidden characters like Byte Order Marks (BOM) or whitespace + version = self.api_version.strip().replace("\ufeff", "") + if version == "v1": + version = "2024-02-15-preview" + params["api-version"] = version return params + + async def forward_request( + self, + request: Request, + path: str, + headers: dict, + request_body: bytes | None, + key: "ApiKey", + max_cost_for_model: int, + session: "AsyncSession", + model_obj: "Model", + ) -> Response | StreamingResponse: + """Forward request to Azure OpenAI.""" + # Fix: If base_url contains /openai/v1, remove it + actual_base_url = self.base_url + if "/openai/v1" in actual_base_url: + actual_base_url = actual_base_url.split("/openai/v1")[0] + + # Use canonical_slug as it often stores the deployment name in Azure setups + # otherwise fallback to transform_model_name + deployment_id = getattr( + model_obj, "canonical_slug", None + ) or self.transform_model_name(model_obj.id) + + # Ensure deployment_id doesn't contain a provider prefix (e.g., 'openai/' or 'azure/') + if "/" in deployment_id: + deployment_id = deployment_id.split("/")[-1] + + # Azure format: openai/deployments/{deployment-id}/chat/completions + clean_path = path.lstrip("/") + if clean_path.startswith("v1/"): + clean_path = clean_path[3:] + azure_path = f"openai/deployments/{deployment_id}/{clean_path}" + + # Temporary backup and restore base_url to use cleaned version + original_base = self.base_url + self.base_url = actual_base_url + + # The query params are handled by super().forward_request via prepare_params + # We don't need to manually append them to full_url for the print if we want to be accurate + params = self.prepare_params(path, {}) + full_url = ( + f"{actual_base_url}/{azure_path}?api-version={params.get('api-version')}" + ) + print(f"\n[DEBUG] Azure Forwarding URL: {full_url}") + print(f"[DEBUG] Deployment ID: {deployment_id}") + + try: + response = await super().forward_request( + request, + azure_path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + + # Check if it's an error response to print details + if hasattr(response, "status_code") and response.status_code != 200: + print(f"[DEBUG] Azure Error Status: {response.status_code}") + if hasattr(response, "body"): + print( + f"[DEBUG] Azure Error Body: {response.body.decode() if isinstance(response.body, bytes) else response.body}" + ) + + return response + except Exception as e: + print(f"[DEBUG] Azure Exception: {str(e)}") + raise + finally: + self.base_url = original_base + + def transform_model_name(self, model_id: str) -> str: + """Extract deployment name from model ID.""" + if "/" in model_id: + return model_id.split("/")[-1] + return model_id diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 95b6ca84..92c3a46d 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -205,6 +205,9 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: provider = _instantiate_provider(provider_row) if provider: + # Keep provider DB id on runtime instance so model mapping can + # bind DB overrides to the correct upstream. + setattr(provider, "db_id", provider_row.id) await provider.refresh_models_cache() logger.debug( f"Initialized {provider_row.provider_type} provider", From ce9834d7ec487df520de087dc900794e56cc529e Mon Sep 17 00:00:00 2001 From: Evan Yang Date: Tue, 10 Feb 2026 19:15:25 +0800 Subject: [PATCH 3/6] fix: harden azure routing and model override mapping --- routstr/algorithm.py | 118 ++++++++++++++++++------------ routstr/upstream/azure.py | 97 ++++++------------------ routstr/upstream/base.py | 36 +++++---- routstr/upstream/ollama.py | 41 ++--------- tests/unit/test_algorithm.py | 93 ++++++++++++++++++++++- tests/unit/test_upstream_azure.py | 82 +++++++++++++++++++++ 6 files changed, 297 insertions(+), 170 deletions(-) create mode 100644 tests/unit/test_upstream_azure.py diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 845bca57..87770a25 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -141,6 +141,16 @@ 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 get_provider_identity(upstream: "BaseUpstreamProvider") -> str: + """Get a stable provider identity used for deduplication.""" + db_id = getattr(upstream, "db_id", None) + if isinstance(db_id, int): + return f"db:{db_id}" + + provider_type = str(getattr(upstream, "provider_type", "") or "").lower() + base_url = str(getattr(upstream, "base_url", "") or "").lower() + return f"{provider_type}|{base_url}" + def _add_candidate( alias: str, model: "Model", provider: "BaseUpstreamProvider" ) -> None: @@ -155,9 +165,7 @@ def create_model_mappings( ) -> None: """Process all models from a given provider.""" upstream_prefix = getattr(upstream, "upstream_name", None) - provider_key = getattr(upstream, "provider_type", "") or getattr( - upstream, "base_url", "" - ) + provider_key = get_provider_identity(upstream) for model in upstream.get_cached_models(): if not model.enabled or model.id in disabled_model_ids: @@ -199,7 +207,7 @@ def create_model_mappings( # Try to set each alias for alias in aliases: _add_candidate(alias, model_to_use, upstream) - seen_model_provider.add((model_to_use.id.lower(), provider_key.lower())) + seen_model_provider.add((model_to_use.id.lower(), provider_key)) # Process non-OpenRouter providers first for upstream in other_upstreams: @@ -214,60 +222,80 @@ def create_model_mappings( for model_id, override_data in overrides_by_id.items(): if model_id in disabled_model_ids: continue + override_row, provider_fee = override_data + upstream_provider_id = getattr(override_row, "upstream_provider_id", None) + if not isinstance(upstream_provider_id, int): + continue + + upstream_for_override = providers_by_db_id.get(upstream_provider_id) + if upstream_for_override is None: + continue + + provider_key = get_provider_identity(upstream_for_override) + dedupe_key = (model_id.lower(), provider_key) + if dedupe_key in seen_model_provider: + continue + try: - override_row, provider_fee = override_data - upstream_provider_id = getattr(override_row, "upstream_provider_id", None) - if not isinstance(upstream_provider_id, int): - continue - - upstream = providers_by_db_id.get(upstream_provider_id) - if upstream is None: - continue - - provider_key = getattr(upstream, "provider_type", "") or getattr( - upstream, "base_url", "" - ) - dedupe_key = (model_id.lower(), provider_key.lower()) - if dedupe_key in seen_model_provider: - continue - model_to_use = _row_to_model( override_row, apply_provider_fee=True, provider_fee=provider_fee ) - if not model_to_use.enabled: - continue - - base_id = get_base_model_id(model_to_use.id) - is_openrouter = ( - getattr(upstream, "base_url", "") == "https://openrouter.ai/api/v1" + except Exception as exc: + logger.warning( + "Skipping invalid model override while building model mappings", + extra={ + "model_id": model_id, + "upstream_provider_id": upstream_provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, ) - if not is_openrouter or base_id not in unique_models: - unique_model = model_to_use.copy( - update={ - "id": base_id, - "upstream_provider_id": upstream.provider_type, - } - ) - unique_models[base_id] = unique_model + continue + if not model_to_use.enabled: + continue + base_id = get_base_model_id(model_to_use.id) + is_openrouter = ( + getattr(upstream_for_override, "base_url", "") + == "https://openrouter.ai/api/v1" + ) + if not is_openrouter or base_id not in unique_models: + unique_model = model_to_use.copy( + update={ + "id": base_id, + "upstream_provider_id": upstream_for_override.provider_type, + } + ) + unique_models[base_id] = unique_model + + try: aliases = resolve_model_alias( model_to_use.id, model_to_use.canonical_slug, alias_ids=model_to_use.alias_ids, ) - upstream_prefix = getattr(upstream, "upstream_name", None) - if upstream_prefix and "/" not in model_to_use.id: - prefixed_id = f"{upstream_prefix}/{model_to_use.id}" - if prefixed_id not in aliases: - aliases.append(prefixed_id) - - for alias in aliases: - _add_candidate(alias, model_to_use, upstream) - seen_model_provider.add(dedupe_key) - except Exception: - # Keep model map creation resilient to malformed overrides. + except Exception as exc: + logger.warning( + "Skipping model aliases for invalid override model", + extra={ + "model_id": model_id, + "upstream_provider_id": upstream_provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) continue + upstream_prefix = getattr(upstream_for_override, "upstream_name", None) + if upstream_prefix and "/" not in model_to_use.id: + prefixed_id = f"{upstream_prefix}/{model_to_use.id}" + if prefixed_id not in aliases: + aliases.append(prefixed_id) + + for alias in aliases: + _add_candidate(alias, model_to_use, upstream_for_override) + seen_model_provider.add(dedupe_key) + # Sort candidates and build final maps model_instances: dict[str, "Model"] = {} provider_map: dict[str, list["BaseUpstreamProvider"]] = {} diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index b9293503..e412c466 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -1,13 +1,9 @@ from typing import TYPE_CHECKING, Mapping -from fastapi import Request -from fastapi.responses import Response, StreamingResponse - from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..auth import ApiKey - from ..core.db import AsyncSession, UpstreamProviderRow + from ..core.db import UpstreamProviderRow from ..payment.models import Model @@ -77,86 +73,35 @@ class AzureUpstreamProvider(BaseUpstreamProvider): ) -> Mapping[str, str]: """Prepare query parameters for Azure OpenAI, adding API version.""" params = dict(query_params or {}) - # Ensure we use a valid Azure API version format - # Strip any hidden characters like Byte Order Marks (BOM) or whitespace - version = self.api_version.strip().replace("\ufeff", "") - if version == "v1": + version = (self.api_version or "").replace("\ufeff", "").strip() + if not version or version.lower() == "v1": version = "2024-02-15-preview" params["api-version"] = version return params - async def forward_request( - self, - request: Request, - path: str, - headers: dict, - request_body: bytes | None, - key: "ApiKey", - max_cost_for_model: int, - session: "AsyncSession", - model_obj: "Model", - ) -> Response | StreamingResponse: - """Forward request to Azure OpenAI.""" - # Fix: If base_url contains /openai/v1, remove it - actual_base_url = self.base_url - if "/openai/v1" in actual_base_url: - actual_base_url = actual_base_url.split("/openai/v1")[0] + def normalize_request_path( + self, path: str, model_obj: "Model | None" = None + ) -> str: + """Build Azure deployment-specific request path.""" + clean_path = super().normalize_request_path(path, model_obj).lstrip("/") + if model_obj is None: + return clean_path - # Use canonical_slug as it often stores the deployment name in Azure setups - # otherwise fallback to transform_model_name deployment_id = getattr( model_obj, "canonical_slug", None ) or self.transform_model_name(model_obj.id) + deployment_id = deployment_id.split("/")[-1] + return f"openai/deployments/{deployment_id}/{clean_path}" - # Ensure deployment_id doesn't contain a provider prefix (e.g., 'openai/' or 'azure/') - if "/" in deployment_id: - deployment_id = deployment_id.split("/")[-1] - - # Azure format: openai/deployments/{deployment-id}/chat/completions - clean_path = path.lstrip("/") - if clean_path.startswith("v1/"): - clean_path = clean_path[3:] - azure_path = f"openai/deployments/{deployment_id}/{clean_path}" - - # Temporary backup and restore base_url to use cleaned version - original_base = self.base_url - self.base_url = actual_base_url - - # The query params are handled by super().forward_request via prepare_params - # We don't need to manually append them to full_url for the print if we want to be accurate - params = self.prepare_params(path, {}) - full_url = ( - f"{actual_base_url}/{azure_path}?api-version={params.get('api-version')}" - ) - print(f"\n[DEBUG] Azure Forwarding URL: {full_url}") - print(f"[DEBUG] Deployment ID: {deployment_id}") - - try: - response = await super().forward_request( - request, - azure_path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, - ) - - # Check if it's an error response to print details - if hasattr(response, "status_code") and response.status_code != 200: - print(f"[DEBUG] Azure Error Status: {response.status_code}") - if hasattr(response, "body"): - print( - f"[DEBUG] Azure Error Body: {response.body.decode() if isinstance(response.body, bytes) else response.body}" - ) - - return response - except Exception as e: - print(f"[DEBUG] Azure Exception: {str(e)}") - raise - finally: - self.base_url = original_base + def get_request_base_url( + self, path: str, model_obj: "Model | None" = None + ) -> str: + """Use endpoint root, stripping accidental /openai/v1 suffix if present.""" + base_url = self.base_url.rstrip("/") + marker = "/openai/v1" + if marker in base_url: + base_url = base_url.split(marker, 1)[0].rstrip("/") + return base_url def transform_model_name(self, model_id: str) -> str: """Extract deployment name from model ID.""" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index aabb12ca..9eb16e4f 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -199,6 +199,23 @@ class BaseUpstreamProvider: """ return model_id + def normalize_request_path(self, path: str, model_obj: Model | None = None) -> str: + """Normalize request path before forwarding to upstream.""" + if path.startswith("v1/"): + return path.replace("v1/", "", 1) + return path + + def get_request_base_url( + self, path: str, model_obj: Model | None = None + ) -> str: + """Get upstream base URL used when building forwarding URL.""" + return self.base_url.rstrip("/") + + def build_request_url(self, path: str, model_obj: Model | None = None) -> str: + """Build full upstream URL from normalized path.""" + clean_path = path.lstrip("/") + return f"{self.get_request_base_url(path, model_obj)}/{clean_path}" + def prepare_responses_request_body( self, body: bytes | None, model_obj: Model ) -> bytes | None: @@ -1037,10 +1054,8 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse from upstream with cost tracking """ - if path.startswith("v1/"): - path = path.replace("v1/", "") - - url = f"{self.base_url}/{path}" + path = self.normalize_request_path(path, model_obj) + url = self.build_request_url(path, model_obj) transformed_body = self.prepare_request_body(request_body, model_obj) @@ -1281,11 +1296,8 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse from upstream with cost tracking """ - # Remove v1/ prefix if present for Responses API - if path.startswith("v1/"): - path = path.replace("v1/", "") - - url = f"{self.base_url}/{path}" + path = self.normalize_request_path(path, model_obj) + url = self.build_request_url(path, model_obj) transformed_body = self.prepare_responses_request_body(request_body, model_obj) @@ -1493,10 +1505,8 @@ class BaseUpstreamProvider: Returns: StreamingResponse from upstream """ - if path.startswith("v1/"): - path = path.replace("v1/", "") - - url = f"{self.base_url}/{path}" + path = self.normalize_request_path(path) + url = self.build_request_url(path) logger.info( "Forwarding GET request to upstream", diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 24703eb2..eff5d5bb 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -3,13 +3,11 @@ from __future__ import annotations from typing import TYPE_CHECKING import httpx -from fastapi import Request -from fastapi.responses import Response, StreamingResponse from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow + from ..core.db import UpstreamProviderRow from ..payment.models import Model from ..core.logging import get_logger @@ -67,38 +65,11 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - async def forward_request( - self, - request: Request, - path: str, - headers: dict, - request_body: bytes | None, - key: ApiKey, - max_cost_for_model: int, - session: AsyncSession, - model_obj: Model, - ) -> Response | StreamingResponse: - """Override to use OpenAI-compatible endpoint for proxy requests.""" - if path.startswith("v1/"): - path = path.replace("v1/", "") - - original_base_url = self.base_url - self.base_url = f"{self.base_url}/v1" - - try: - result = await super().forward_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, - ) - return result - finally: - self.base_url = original_base_url + def get_request_base_url( + self, path: str, model_obj: Model | None = None + ) -> str: + """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" + return f"{self.base_url.rstrip('/')}/v1" async def fetch_models(self) -> list[Model]: """Fetch models from Ollama API using /api/tags endpoint.""" diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index fd7a837f..1a647528 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -1,14 +1,18 @@ """Tests for the model prioritization algorithm.""" import os +from types import SimpleNamespace from unittest.mock import Mock +import pytest + # Set required env vars before importing os.environ["UPSTREAM_BASE_URL"] = "http://test" os.environ["UPSTREAM_API_KEY"] = "test" from routstr.algorithm import ( # noqa: E402 calculate_model_cost_score, + create_model_mappings, get_provider_penalty, ) from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 @@ -45,11 +49,21 @@ def create_test_model( ) -def create_test_provider(name: str, base_url: str = "http://test.com") -> Mock: +def create_test_provider( + name: str, + base_url: str = "http://test.com", + *, + db_id: int | None = None, + models: list[Model] | None = None, + upstream_name: str | None = None, +) -> Mock: """Helper to create a test provider mock.""" provider = Mock() provider.provider_type = name provider.base_url = base_url + provider.db_id = db_id + provider.upstream_name = upstream_name or name + provider.get_cached_models.return_value = models or [] return provider @@ -99,3 +113,80 @@ 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_create_model_mappings_includes_db_override_for_missing_cached_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Model overrides should still map when provider discovery misses the model.""" + provider = create_test_provider( + "azure", + "https://example.openai.azure.com/openai/v1", + db_id=7, + models=[], + ) + override_model = create_test_model("azure/gpt-4o") + override_model.canonical_slug = "azure-deployment" + + def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def] + return override_model + + monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model) + + override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=7, enabled=True) + + model_instances, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_id={"azure/gpt-4o": (override_row, 1.01)}, + disabled_model_ids=set(), + ) + + assert "azure/gpt-4o" in model_instances + assert provider_map["azure/gpt-4o"] == [provider] + assert "gpt-4o" in unique_models + + +def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Different provider instances of same type should both survive dedupe.""" + provider_a_model = create_test_model( + "azure/gpt-4o", prompt_price=0.01, completion_price=0.01 + ) + provider_a = create_test_provider( + "azure", + "https://a.openai.azure.com/openai/v1", + db_id=1, + models=[provider_a_model], + upstream_name="azure-a", + ) + provider_b = create_test_provider( + "azure", + "https://b.openai.azure.com/openai/v1", + db_id=2, + models=[], + upstream_name="azure-b", + ) + + override_model = create_test_model( + "azure/gpt-4o", prompt_price=0.001, completion_price=0.001 + ) + override_model.canonical_slug = "azure-b-deployment" + + def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def] + return override_model + + monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model) + + override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=2, enabled=True) + + _, provider_map, _ = create_model_mappings( + upstreams=[provider_a, provider_b], + overrides_by_id={"azure/gpt-4o": (override_row, 1.01)}, + disabled_model_ids=set(), + ) + + providers_for_alias = provider_map["azure/gpt-4o"] + assert provider_a in providers_for_alias + assert provider_b in providers_for_alias + assert len(providers_for_alias) == 2 diff --git a/tests/unit/test_upstream_azure.py b/tests/unit/test_upstream_azure.py new file mode 100644 index 00000000..86c03890 --- /dev/null +++ b/tests/unit/test_upstream_azure.py @@ -0,0 +1,82 @@ +"""Tests for Azure upstream provider request normalization.""" + +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream.azure import AzureUpstreamProvider + + +def create_test_model(model_id: str, canonical_slug: str | None = None) -> Model: + return Model( + id=model_id, + name=model_id, + created=0, + description="test", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test", + instruct_type=None, + ), + pricing=Pricing( + prompt=0.001, + completion=0.001, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + ), + canonical_slug=canonical_slug, + ) + + +def test_prepare_headers_uses_azure_api_key_header() -> None: + provider = AzureUpstreamProvider( + base_url="https://example.openai.azure.com", + api_key="azure-key", + api_version="2024-02-15-preview", + ) + + headers = provider.prepare_headers({"Authorization": "Bearer user-token"}) + + assert headers["api-key"] == "azure-key" + assert "Authorization" not in headers + assert "authorization" not in headers + + +def test_prepare_params_normalizes_azure_api_version() -> None: + provider = AzureUpstreamProvider( + base_url="https://example.openai.azure.com", + api_key="azure-key", + api_version="\ufeff v1 ", + ) + + params = provider.prepare_params("chat/completions", {}) + + assert params["api-version"] == "2024-02-15-preview" + + +def test_normalize_request_path_includes_deployment_id() -> None: + provider = AzureUpstreamProvider( + base_url="https://example.openai.azure.com/openai/v1", + api_key="azure-key", + api_version="2024-02-15-preview", + ) + model = create_test_model("azure/gpt-4o", canonical_slug="deploy-gpt4o") + + path = provider.normalize_request_path("v1/chat/completions", model) + + assert path == "openai/deployments/deploy-gpt4o/chat/completions" + + +def test_get_request_base_url_strips_openai_v1_suffix() -> None: + provider = AzureUpstreamProvider( + base_url="https://example.openai.azure.com/openai/v1", + api_key="azure-key", + api_version="2024-02-15-preview", + ) + model = create_test_model("azure/gpt-4o") + + base_url = provider.get_request_base_url("chat/completions", model) + + assert base_url == "https://example.openai.azure.com" From 9fb6f54d12a0bdfd2bc30374a75905dc52e76361 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 3 Mar 2026 15:29:07 +0100 Subject: [PATCH 4/6] fix negative reserve balance --- routstr/auth.py | 52 +++-- routstr/proxy.py | 6 +- routstr/upstream/base.py | 28 +-- .../test_reserved_balance_negative.py | 221 +++++++++++++++++- 4 files changed, 253 insertions(+), 54 deletions(-) diff --git a/routstr/auth.py b/routstr/auth.py index 6e03d968..fa880f92 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -588,12 +588,15 @@ async def pay_for_request( async def revert_pay_for_request( key: ApiKey, session: AsyncSession, cost_per_request: int -) -> None: +) -> bool: + """Revert a previously reserved payment. Returns True if revert succeeded, + False if the reservation was already released (prevents negative reserved_balance).""" billing_key = await get_billing_key(key, session) stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.reserved_balance) >= cost_per_request) .values( reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, total_requests=col(ApiKey.total_requests) - 1, @@ -607,6 +610,7 @@ async def revert_pay_for_request( child_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.reserved_balance) >= cost_per_request) .values( total_requests=col(ApiKey.total_requests) - 1, reserved_balance=col(ApiKey.reserved_balance) - cost_per_request, @@ -616,8 +620,8 @@ async def revert_pay_for_request( await session.commit() if result.rowcount == 0: - logger.error( - "Failed to revert payment - insufficient reserved balance", + logger.warning( + "Revert skipped - reservation already released (no-op to prevent negative reserved_balance)", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", @@ -625,19 +629,11 @@ async def revert_pay_for_request( "current_reserved_balance": billing_key.reserved_balance, }, ) - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"failed to revert request payment: {cost_per_request} mSats required. {billing_key.balance} available.", - "type": "payment_error", - "code": "payment_error", - } - }, - ) + return False await session.refresh(billing_key) if billing_key.hashed_key != key.hashed_key: await session.refresh(key) + return True async def adjust_payment_for_tokens( @@ -669,17 +665,19 @@ async def adjust_payment_for_tokens( release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == billing_key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost ) ) - await session.exec(release_stmt) # type: ignore[call-overload] + result = await session.exec(release_stmt) # type: ignore[call-overload] # Also release on child key if it's different if billing_key.hashed_key != key.hashed_key: child_release_stmt = ( update(ApiKey) .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.reserved_balance) >= deducted_max_cost) .values( reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost @@ -688,14 +686,24 @@ async def adjust_payment_for_tokens( await session.exec(child_release_stmt) # type: ignore[call-overload] await session.commit() - logger.warning( - "Released reservation without charging (fallback)", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "deducted_max_cost": deducted_max_cost, - }, - ) + if result.rowcount == 0: # type: ignore[union-attr] + logger.warning( + "Release reservation skipped - already released (no-op to prevent negative reserved_balance)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, + ) + else: + logger.warning( + "Released reservation without charging (fallback)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + }, + ) except Exception as e: logger.error( "Failed to release reservation in fallback", diff --git a/routstr/proxy.py b/routstr/proxy.py index 4d8f666d..09b34549 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -300,9 +300,13 @@ async def proxy( session, model_obj, ) + except UpstreamError: + # Let the outer UpstreamError handler manage retry/revert + raise except Exception as e: + # Unexpected error (not an upstream failure) — revert and propagate logger.error( - "Upstream request failed, ensuring payment is reverted", + "Unexpected error in upstream request, reverting payment", extra={ "error": str(e), "error_type": type(e).__name__, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5efe5b0c..1be99b15 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -13,7 +13,7 @@ from fastapi.responses import Response, StreamingResponse from pydantic import BaseModel from sqlmodel import select -from ..auth import adjust_payment_for_tokens, revert_pay_for_request +from ..auth import adjust_payment_for_tokens from ..core import get_logger from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow, create_session from ..core.exceptions import UpstreamError @@ -1279,8 +1279,7 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - + # Don't revert here — proxy.py owns payment revert to avoid double-revert if isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" elif isinstance(exc, httpx.TimeoutException): @@ -1310,13 +1309,9 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + # Don't revert here — proxy.py owns payment revert to avoid double-revert + raise UpstreamError( + "An unexpected server error occurred", status_code=500 ) async def forward_responses_request( @@ -1501,8 +1496,7 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - + # Don't revert here — proxy.py owns payment revert to avoid double-revert if isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" elif isinstance(exc, httpx.TimeoutException): @@ -1532,13 +1526,9 @@ class BaseUpstreamProvider: }, ) - await revert_pay_for_request(key, session, max_cost_for_model) - - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + # Don't revert here — proxy.py owns payment revert to avoid double-revert + raise UpstreamError( + "An unexpected server error occurred", status_code=500 ) async def forward_get_request( diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py index 21f9b6d0..b283d55b 100644 --- a/tests/integration/test_reserved_balance_negative.py +++ b/tests/integration/test_reserved_balance_negative.py @@ -133,13 +133,16 @@ async def test_reserved_balance_with_successful_requests( @pytest.mark.asyncio -async def test_insufficient_reserved_balance_for_revert( +async def test_revert_with_zero_reserved_balance_is_noop( integration_session: AsyncSession, ) -> None: - """Test revert_pay_for_request behavior with insufficient reserved balance.""" + """Test that revert_pay_for_request is a no-op when reserved_balance is 0. + + Previously this would drive reserved_balance negative. With the floor guard, + it should return False and leave reserved_balance at 0. + """ from routstr.auth import revert_pay_for_request - # Create key with zero reserved balance unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}" test_key = ApiKey( hashed_key=unique_key, @@ -149,17 +152,211 @@ async def test_insufficient_reserved_balance_for_revert( integration_session.add(test_key) await integration_session.commit() - # Try to revert more than available - # Note: Current implementation allows reserved_balance to go negative - await revert_pay_for_request(test_key, integration_session, 100) + # Try to revert more than available — should be a no-op + result = await revert_pay_for_request(test_key, integration_session, 100) - # Refresh to get updated values await integration_session.refresh(test_key) - # Current implementation allows negative reserved balance - assert test_key.reserved_balance == -100, ( - f"Expected reserved_balance to be -100, got: {test_key.reserved_balance}" + assert result is False, "Revert should return False when reservation already released" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should remain 0, got: {test_key.reserved_balance}" ) - assert test_key.total_requests == -1, ( - f"Expected total_requests to be -1, got: {test_key.total_requests}" + assert test_key.total_requests == 0, ( + f"Total requests should remain 0, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_revert_with_sufficient_reserved_balance_succeeds( + integration_session: AsyncSession, +) -> None: + """Test that revert_pay_for_request works correctly when there is enough reserved balance.""" + from routstr.auth import revert_pay_for_request + + unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=5000, + reserved_balance=500, + total_requests=3, + ) + integration_session.add(test_key) + await integration_session.commit() + + result = await revert_pay_for_request(test_key, integration_session, 500) + + await integration_session.refresh(test_key) + + assert result is True, "Revert should return True on success" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should be 0, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 2, ( + f"Total requests should be 2, got: {test_key.total_requests}" + ) + assert test_key.balance == 5000, "Balance should not change on revert" + + +@pytest.mark.asyncio +async def test_revert_partial_reserved_balance_is_noop( + integration_session: AsyncSession, +) -> None: + """Test that reverting more than the current reserved_balance is a no-op.""" + from routstr.auth import revert_pay_for_request + + unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=5000, + reserved_balance=50, + total_requests=1, + ) + integration_session.add(test_key) + await integration_session.commit() + + # Try to revert 500 when only 50 is reserved — should be no-op + result = await revert_pay_for_request(test_key, integration_session, 500) + + await integration_session.refresh(test_key) + + assert result is False, "Revert should fail when cost > reserved_balance" + assert test_key.reserved_balance == 50, ( + f"Reserved balance should stay at 50, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 1, ( + f"Total requests should stay at 1, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_double_revert_prevented( + integration_session: AsyncSession, +) -> None: + """Test that calling revert twice doesn't drive reserved_balance negative. + + This simulates the double-revert scenario where both upstream/base.py + and proxy.py attempt to revert the same reservation. + """ + from routstr.auth import revert_pay_for_request + + unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=10000, + reserved_balance=500, + total_requests=5, + ) + integration_session.add(test_key) + await integration_session.commit() + + # First revert — should succeed + result1 = await revert_pay_for_request(test_key, integration_session, 500) + await integration_session.refresh(test_key) + + assert result1 is True + assert test_key.reserved_balance == 0 + assert test_key.total_requests == 4 + + # Second revert of the same amount — should be no-op + result2 = await revert_pay_for_request(test_key, integration_session, 500) + await integration_session.refresh(test_key) + + assert result2 is False, "Second revert should be a no-op" + assert test_key.reserved_balance == 0, ( + f"Reserved balance should stay 0, got: {test_key.reserved_balance}" + ) + assert test_key.total_requests == 4, ( + f"Total requests should stay 4, got: {test_key.total_requests}" + ) + + +@pytest.mark.asyncio +async def test_sequential_reverts_never_go_negative( + integration_session: AsyncSession, +) -> None: + """Test that multiple reverts don't cause negative reserved_balance. + + Simulates the double-revert scenario where multiple code paths + attempt to revert the same reservation. + """ + from routstr.auth import revert_pay_for_request + + unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}" + test_key = ApiKey( + hashed_key=unique_key, + balance=10000, + reserved_balance=500, + total_requests=5, + ) + integration_session.add(test_key) + await integration_session.commit() + + # Run 5 sequential reverts for the same 500 reservation + results = [] + for _ in range(5): + r = await revert_pay_for_request(test_key, integration_session, 500) + results.append(r) + + await integration_session.refresh(test_key) + + # Exactly one should succeed, rest should be no-ops + success_count = sum(1 for r in results if r is True) + assert success_count == 1, ( + f"Exactly one revert should succeed, got {success_count} successes" + ) + assert test_key.reserved_balance == 0, ( + f"Reserved balance should be 0, got: {test_key.reserved_balance}" + ) + assert test_key.reserved_balance >= 0, ( + f"Reserved balance went negative: {test_key.reserved_balance}" + ) + + +@pytest.mark.asyncio +async def test_child_key_revert_floor_guard( + integration_session: AsyncSession, +) -> None: + """Test that child key reserved_balance also has floor guard on revert.""" + from routstr.auth import revert_pay_for_request + + parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}" + child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}" + + parent_key = ApiKey( + hashed_key=parent_key_hash, + balance=10000, + reserved_balance=500, + total_requests=3, + ) + child_key = ApiKey( + hashed_key=child_key_hash, + balance=0, + reserved_balance=500, + total_requests=3, + parent_key_hash=parent_key_hash, + ) + integration_session.add(parent_key) + integration_session.add(child_key) + await integration_session.commit() + + # First revert succeeds + result1 = await revert_pay_for_request(child_key, integration_session, 500) + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key) + + assert result1 is True + assert parent_key.reserved_balance == 0 + assert child_key.reserved_balance == 0 + + # Second revert is a no-op for both parent and child + result2 = await revert_pay_for_request(child_key, integration_session, 500) + await integration_session.refresh(parent_key) + await integration_session.refresh(child_key) + + assert result2 is False + assert parent_key.reserved_balance == 0, ( + f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}" + ) + assert child_key.reserved_balance == 0, ( + f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}" ) From e002c0b66f6c966a1a1432822c8f91c71bda829b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 6 Mar 2026 17:59:49 +0100 Subject: [PATCH 5/6] clean up --- routstr/upstream/base.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 1be99b15..43566ea4 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3083,7 +3083,7 @@ class BaseUpstreamProvider: UpstreamProviderRow.api_key == self.api_key ) result = await session.exec(stmt) - + # .first() returns the object or None if not found provider = result.first() if not provider or not provider.id: @@ -3105,7 +3105,6 @@ class BaseUpstreamProvider: models.append(found_db_model) models_with_fees = [self._apply_provider_fee_to_model(m) for m in models] - print([mode.id for mode in models_with_fees]) try: sats_to_usd = sats_usd_price() From d0c7cc6bd9cc6bf7bdfdb84afee28b68685780a9 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 6 Mar 2026 23:02:47 +0100 Subject: [PATCH 6/6] fix refund token multiple times --- routstr/balance.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index 84ea01d9..b738a6c3 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -205,8 +205,9 @@ async def refund_wallet_endpoint( key: ApiKey = await validate_bearer_key(bearer_value, session) - if cached := await _refund_cache_get(bearer_value): - return cached + if key.total_balance <= 0: + if cached := await _refund_cache_get(bearer_value): + return cached if key.parent_key_hash: raise HTTPException(