From b7e4fbf7392901583776f94a90f66cbc2842b3ba Mon Sep 17 00:00:00 2001 From: Shroominic Date: Tue, 11 Nov 2025 12:43:02 +0800 Subject: [PATCH] fix model alias problem for anthropic --- routstr/algorithm.py | 6 +++- routstr/payment/cost_caculation.py | 10 ++---- routstr/payment/helpers.py | 6 ++-- routstr/payment/models.py | 2 ++ routstr/upstream.py | 56 ++++-------------------------- routstr/upstreams/upstream.py | 28 +++++++++++++-- 6 files changed, 43 insertions(+), 65 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 66f25d2b..316b9f35 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -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: diff --git a/routstr/payment/cost_caculation.py b/routstr/payment/cost_caculation.py index 4efe13ff..f8eb4ffb 100644 --- a/routstr/payment/cost_caculation.py +++ b/routstr/payment/cost_caculation.py @@ -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( diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index e4893d48..5744e37a 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -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 diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 0f4a4c0b..99888642 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -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( diff --git a/routstr/upstream.py b/routstr/upstream.py index c1744585..4d22930b 100644 --- a/routstr/upstream.py +++ b/routstr/upstream.py @@ -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: diff --git a/routstr/upstreams/upstream.py b/routstr/upstreams/upstream.py index 9d3b1e35..f2591d83 100644 --- a/routstr/upstreams/upstream.py +++ b/routstr/upstreams/upstream.py @@ -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):