mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix: harden azure routing and model override mapping
This commit is contained in:
+46
-18
@@ -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
@@ -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
@@ -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",
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
Reference in New Issue
Block a user