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"