fix: harden azure routing and model override mapping

This commit is contained in:
Evan Yang
2026-02-10 19:15:25 +08:00
parent 6ebe73f2f7
commit ce9834d7ec
6 changed files with 297 additions and 170 deletions
+46 -18
View File
@@ -141,6 +141,16 @@ def create_model_mappings(
"""Get base model ID by removing provider prefix.""" """Get base model ID by removing provider prefix."""
return model_id.split("/", 1)[1] if "/" in model_id else model_id 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( def _add_candidate(
alias: str, model: "Model", provider: "BaseUpstreamProvider" alias: str, model: "Model", provider: "BaseUpstreamProvider"
) -> None: ) -> None:
@@ -155,9 +165,7 @@ def create_model_mappings(
) -> None: ) -> None:
"""Process all models from a given provider.""" """Process all models from a given provider."""
upstream_prefix = getattr(upstream, "upstream_name", None) upstream_prefix = getattr(upstream, "upstream_name", None)
provider_key = getattr(upstream, "provider_type", "") or getattr( provider_key = get_provider_identity(upstream)
upstream, "base_url", ""
)
for model in upstream.get_cached_models(): for model in upstream.get_cached_models():
if not model.enabled or model.id in disabled_model_ids: if not model.enabled or model.id in disabled_model_ids:
@@ -199,7 +207,7 @@ def create_model_mappings(
# Try to set each alias # Try to set each alias
for alias in aliases: for alias in aliases:
_add_candidate(alias, model_to_use, upstream) _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 # Process non-OpenRouter providers first
for upstream in other_upstreams: for upstream in other_upstreams:
@@ -214,59 +222,79 @@ def create_model_mappings(
for model_id, override_data in overrides_by_id.items(): for model_id, override_data in overrides_by_id.items():
if model_id in disabled_model_ids: if model_id in disabled_model_ids:
continue continue
try:
override_row, provider_fee = override_data override_row, provider_fee = override_data
upstream_provider_id = getattr(override_row, "upstream_provider_id", None) upstream_provider_id = getattr(override_row, "upstream_provider_id", None)
if not isinstance(upstream_provider_id, int): if not isinstance(upstream_provider_id, int):
continue continue
upstream = providers_by_db_id.get(upstream_provider_id) upstream_for_override = providers_by_db_id.get(upstream_provider_id)
if upstream is None: if upstream_for_override is None:
continue continue
provider_key = getattr(upstream, "provider_type", "") or getattr( provider_key = get_provider_identity(upstream_for_override)
upstream, "base_url", "" dedupe_key = (model_id.lower(), provider_key)
)
dedupe_key = (model_id.lower(), provider_key.lower())
if dedupe_key in seen_model_provider: if dedupe_key in seen_model_provider:
continue continue
try:
model_to_use = _row_to_model( model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee override_row, apply_provider_fee=True, provider_fee=provider_fee
) )
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__,
},
)
continue
if not model_to_use.enabled: if not model_to_use.enabled:
continue continue
base_id = get_base_model_id(model_to_use.id) base_id = get_base_model_id(model_to_use.id)
is_openrouter = ( is_openrouter = (
getattr(upstream, "base_url", "") == "https://openrouter.ai/api/v1" getattr(upstream_for_override, "base_url", "")
== "https://openrouter.ai/api/v1"
) )
if not is_openrouter or base_id not in unique_models: if not is_openrouter or base_id not in unique_models:
unique_model = model_to_use.copy( unique_model = model_to_use.copy(
update={ update={
"id": base_id, "id": base_id,
"upstream_provider_id": upstream.provider_type, "upstream_provider_id": upstream_for_override.provider_type,
} }
) )
unique_models[base_id] = unique_model unique_models[base_id] = unique_model
try:
aliases = resolve_model_alias( aliases = resolve_model_alias(
model_to_use.id, model_to_use.id,
model_to_use.canonical_slug, model_to_use.canonical_slug,
alias_ids=model_to_use.alias_ids, alias_ids=model_to_use.alias_ids,
) )
upstream_prefix = getattr(upstream, "upstream_name", None) 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: if upstream_prefix and "/" not in model_to_use.id:
prefixed_id = f"{upstream_prefix}/{model_to_use.id}" prefixed_id = f"{upstream_prefix}/{model_to_use.id}"
if prefixed_id not in aliases: if prefixed_id not in aliases:
aliases.append(prefixed_id) aliases.append(prefixed_id)
for alias in aliases: for alias in aliases:
_add_candidate(alias, model_to_use, upstream) _add_candidate(alias, model_to_use, upstream_for_override)
seen_model_provider.add(dedupe_key) seen_model_provider.add(dedupe_key)
except Exception:
# Keep model map creation resilient to malformed overrides.
continue
# Sort candidates and build final maps # Sort candidates and build final maps
model_instances: dict[str, "Model"] = {} model_instances: dict[str, "Model"] = {}
+20 -75
View File
@@ -1,13 +1,9 @@
from typing import TYPE_CHECKING, Mapping from typing import TYPE_CHECKING, Mapping
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider
if TYPE_CHECKING: if TYPE_CHECKING:
from ..auth import ApiKey from ..core.db import UpstreamProviderRow
from ..core.db import AsyncSession, UpstreamProviderRow
from ..payment.models import Model from ..payment.models import Model
@@ -77,86 +73,35 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
) -> Mapping[str, str]: ) -> Mapping[str, str]:
"""Prepare query parameters for Azure OpenAI, adding API version.""" """Prepare query parameters for Azure OpenAI, adding API version."""
params = dict(query_params or {}) params = dict(query_params or {})
# Ensure we use a valid Azure API version format version = (self.api_version or "").replace("\ufeff", "").strip()
# Strip any hidden characters like Byte Order Marks (BOM) or whitespace if not version or version.lower() == "v1":
version = self.api_version.strip().replace("\ufeff", "")
if version == "v1":
version = "2024-02-15-preview" version = "2024-02-15-preview"
params["api-version"] = version params["api-version"] = version
return params return params
async def forward_request( def normalize_request_path(
self, self, path: str, model_obj: "Model | None" = None
request: Request, ) -> str:
path: str, """Build Azure deployment-specific request path."""
headers: dict, clean_path = super().normalize_request_path(path, model_obj).lstrip("/")
request_body: bytes | None, if model_obj is None:
key: "ApiKey", return clean_path
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( deployment_id = getattr(
model_obj, "canonical_slug", None model_obj, "canonical_slug", None
) or self.transform_model_name(model_obj.id) ) 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] deployment_id = deployment_id.split("/")[-1]
return f"openai/deployments/{deployment_id}/{clean_path}"
# Azure format: openai/deployments/{deployment-id}/chat/completions def get_request_base_url(
clean_path = path.lstrip("/") self, path: str, model_obj: "Model | None" = None
if clean_path.startswith("v1/"): ) -> str:
clean_path = clean_path[3:] """Use endpoint root, stripping accidental /openai/v1 suffix if present."""
azure_path = f"openai/deployments/{deployment_id}/{clean_path}" base_url = self.base_url.rstrip("/")
marker = "/openai/v1"
# Temporary backup and restore base_url to use cleaned version if marker in base_url:
original_base = self.base_url base_url = base_url.split(marker, 1)[0].rstrip("/")
self.base_url = actual_base_url return 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: def transform_model_name(self, model_id: str) -> str:
"""Extract deployment name from model ID.""" """Extract deployment name from model ID."""
+23 -13
View File
@@ -199,6 +199,23 @@ class BaseUpstreamProvider:
""" """
return model_id 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( def prepare_responses_request_body(
self, body: bytes | None, model_obj: Model self, body: bytes | None, model_obj: Model
) -> bytes | None: ) -> bytes | None:
@@ -1037,10 +1054,8 @@ class BaseUpstreamProvider:
Returns: Returns:
Response or StreamingResponse from upstream with cost tracking Response or StreamingResponse from upstream with cost tracking
""" """
if path.startswith("v1/"): path = self.normalize_request_path(path, model_obj)
path = path.replace("v1/", "") url = self.build_request_url(path, model_obj)
url = f"{self.base_url}/{path}"
transformed_body = self.prepare_request_body(request_body, model_obj) transformed_body = self.prepare_request_body(request_body, model_obj)
@@ -1281,11 +1296,8 @@ class BaseUpstreamProvider:
Returns: Returns:
Response or StreamingResponse from upstream with cost tracking Response or StreamingResponse from upstream with cost tracking
""" """
# Remove v1/ prefix if present for Responses API path = self.normalize_request_path(path, model_obj)
if path.startswith("v1/"): url = self.build_request_url(path, model_obj)
path = path.replace("v1/", "")
url = f"{self.base_url}/{path}"
transformed_body = self.prepare_responses_request_body(request_body, model_obj) transformed_body = self.prepare_responses_request_body(request_body, model_obj)
@@ -1493,10 +1505,8 @@ class BaseUpstreamProvider:
Returns: Returns:
StreamingResponse from upstream StreamingResponse from upstream
""" """
if path.startswith("v1/"): path = self.normalize_request_path(path)
path = path.replace("v1/", "") url = self.build_request_url(path)
url = f"{self.base_url}/{path}"
logger.info( logger.info(
"Forwarding GET request to upstream", "Forwarding GET request to upstream",
+6 -35
View File
@@ -3,13 +3,11 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import httpx import httpx
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider
if TYPE_CHECKING: if TYPE_CHECKING:
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow from ..core.db import UpstreamProviderRow
from ..payment.models import Model from ..payment.models import Model
from ..core.logging import get_logger from ..core.logging import get_logger
@@ -67,38 +65,11 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility.""" """Strip 'ollama/' prefix for Ollama API compatibility."""
return model_id.removeprefix("ollama/") return model_id.removeprefix("ollama/")
async def forward_request( def get_request_base_url(
self, self, path: str, model_obj: Model | None = None
request: Request, ) -> str:
path: str, """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
headers: dict, return f"{self.base_url.rstrip('/')}/v1"
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
async def fetch_models(self) -> list[Model]: async def fetch_models(self) -> list[Model]:
"""Fetch models from Ollama API using /api/tags endpoint.""" """Fetch models from Ollama API using /api/tags endpoint."""
+92 -1
View File
@@ -1,14 +1,18 @@
"""Tests for the model prioritization algorithm.""" """Tests for the model prioritization algorithm."""
import os import os
from types import SimpleNamespace
from unittest.mock import Mock from unittest.mock import Mock
import pytest
# Set required env vars before importing # Set required env vars before importing
os.environ["UPSTREAM_BASE_URL"] = "http://test" os.environ["UPSTREAM_BASE_URL"] = "http://test"
os.environ["UPSTREAM_API_KEY"] = "test" os.environ["UPSTREAM_API_KEY"] = "test"
from routstr.algorithm import ( # noqa: E402 from routstr.algorithm import ( # noqa: E402
calculate_model_cost_score, calculate_model_cost_score,
create_model_mappings,
get_provider_penalty, get_provider_penalty,
) )
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 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.""" """Helper to create a test provider mock."""
provider = Mock() provider = Mock()
provider.provider_type = name provider.provider_type = name
provider.base_url = base_url 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 return provider
@@ -99,3 +113,80 @@ def test_get_provider_penalty_openrouter() -> None:
provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1") provider = create_test_provider("openrouter", "https://openrouter.ai/api/v1")
penalty = get_provider_penalty(provider) penalty = get_provider_penalty(provider)
assert penalty == 1.001 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
+82
View File
@@ -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"