fix model alias problem for anthropic

This commit is contained in:
Shroominic
2025-11-11 12:43:02 +08:00
parent cfe03d6dcb
commit b7e4fbf739
6 changed files with 43 additions and 65 deletions
+5 -1
View File
@@ -258,7 +258,11 @@ def create_model_mappings(
unique_models[base_id] = unique_model
# Get all aliases for this model
aliases = resolve_model_alias(model_to_use.id, model_to_use.canonical_slug)
aliases = resolve_model_alias(
model_to_use.id,
model_to_use.canonical_slug,
alias_ids=model_to_use.alias_ids,
)
# Add prefixed alias if applicable
if upstream_prefix and "/" not in model_to_use.id:
+3 -7
View File
@@ -25,7 +25,7 @@ class CostDataError(BaseModel):
code: str
async def calculate_cost(
async def calculate_cost( # todo: can be sync
response_data: dict, max_cost: int, session: AsyncSession
) -> CostData | MaxCostData | CostDataError:
"""
@@ -78,13 +78,9 @@ async def calculate_cost(
extra={"model": response_model},
)
from ..proxy import get_upstreams
from ..upstream import get_model_with_override
from ..proxy import get_model_instance
upstreams = get_upstreams()
model_obj = await get_model_with_override(
response_model, upstreams, session=session
)
model_obj = get_model_instance(response_model)
if not model_obj:
logger.error(
+2 -4
View File
@@ -104,11 +104,9 @@ async def get_max_cost_for_model(
return max(settings.min_request_msat, default_cost_msats)
if not model_obj:
from ..proxy import get_upstreams
from ..upstream import get_model_with_override
from ..proxy import get_model_instance
upstreams = get_upstreams()
model_obj = await get_model_with_override(model, upstreams, session)
model_obj = get_model_instance(model)
if not model_obj:
fallback_msats = settings.fixed_cost_per_request * 1000
+2
View File
@@ -60,6 +60,7 @@ class Model(BaseModel):
enabled: bool = True
upstream_provider_id: int | None = None
canonical_slug: str | None = None
alias_ids: list[str] | None = None
def __hash__(self) -> int:
return hash(self.id)
@@ -409,6 +410,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
alias_ids=model.alias_ids,
)
except Exception as e:
logger.error(
+6 -50
View File
@@ -24,7 +24,9 @@ from .upstreams.generic import GenericUpstreamProvider
logger = get_logger(__name__)
def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> list[str]:
def resolve_model_alias(
model_id: str, canonical_slug: str | None = None, alias_ids: list[str] | None = None
) -> list[str]:
"""Resolve model ID to all possible aliases.
Returns list of aliases including canonical slug and variations without provider prefix.
@@ -66,6 +68,9 @@ def resolve_model_alias(model_id: str, canonical_slug: str | None = None) -> lis
if canonical_base not in aliases:
aliases.append(canonical_base)
if alias_ids:
aliases.extend(alias_ids)
return aliases
@@ -120,55 +125,6 @@ async def get_all_models_with_overrides(
return list(all_models.values())
async def get_model_with_override(
model_id: str,
upstreams: list[UpstreamProvider],
session: AsyncSession,
) -> Model | None:
"""Get a specific model from providers with database override applied.
Resolves model aliases automatically (e.g., both "gpt-5-mini" and "openai/gpt-5-mini").
Args:
model_id: Model identifier (with or without provider prefix)
upstreams: List of upstream provider instances
Returns:
Model object or None if not found
"""
from sqlmodel import select
from .payment.models import _row_to_model
aliases = resolve_model_alias(model_id)
for alias in aliases:
result = await session.exec(
select(ModelRow).where(
ModelRow.id == alias,
ModelRow.upstream_provider_id.isnot(None), # type: ignore
ModelRow.enabled,
)
)
override_row = result.first()
if override_row:
provider = await session.get(
UpstreamProviderRow, override_row.upstream_provider_id
)
provider_fee = provider.provider_fee if provider else 1.01
return _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
for alias in aliases:
for upstream in upstreams:
model = upstream.get_cached_model_by_id(alias)
if model and model.enabled:
return model
return None
async def refresh_upstreams_models_periodically(
upstreams: list[UpstreamProvider],
) -> None:
+25 -3
View File
@@ -1626,6 +1626,7 @@ class UpstreamProvider:
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
alias_ids=model.alias_ids,
)
(
@@ -1648,6 +1649,7 @@ class UpstreamProvider:
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
alias_ids=model.alias_ids,
)
async def fetch_models(self) -> list[Model]:
@@ -1737,13 +1739,33 @@ class AnthropicUpstreamProvider(UpstreamProvider):
)
def transform_model_name(self, model_id: str) -> str:
"""Strip 'anthropic/' prefix for Anthropic API compatibility."""
return model_id.removeprefix("anthropic/")
"""Strip 'anthropic/' prefix for Anthropic API compatibility and transform model names."""
if model_id.startswith("anthropic/"):
model_id = model_id[len("anthropic/") :]
fixed_transforms = {
"claude-haiku-4.5": "claude-haiku-4-5-20251001",
"claude-sonnet-4.5": "claude-sonnet-4-5-20250929",
"claude-opus-4.1": "claude-opus-4-1-20250805",
"claude-opus-4": "claude-opus-4-20250514",
"claude-sonnet-4": "claude-sonnet-4-20250514",
"claude-3.5-haiku": "claude-3-5-haiku-20241022",
"claude-3-haiku": "claude-3-haiku-20240307",
"claude-haiku-4-5": "claude-haiku-4-5-20251001",
"claude-sonnet-4-5": "claude-sonnet-4-5-20250929",
"claude-opus-4-1": "claude-opus-4-1-20250805",
"claude-3-5-haiku": "claude-3-5-haiku-20241022",
}
if model_id in fixed_transforms:
model_id = fixed_transforms[model_id]
return model_id
async def fetch_models(self) -> list[Model]:
"""Fetch Anthropic models from OpenRouter API filtered by anthropic source."""
models_data = await async_fetch_openrouter_models(source_filter="anthropic")
return [Model(**model) for model in models_data] # type: ignore
models = [Model(**model) for model in models_data] # type: ignore
for model in models:
model.alias_ids = [self.transform_model_name(model.id)]
return models
class AzureUpstreamProvider(UpstreamProvider):