mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
fix model alias problem for anthropic
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user