Spaces:
Sleeping
Sleeping
| """Model routing for Claude-compatible requests.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from loguru import logger | |
| from free_claude_code.application.errors import UnknownProviderError | |
| from free_claude_code.config.model_refs import parse_model_name, parse_provider_type | |
| from free_claude_code.config.provider_catalog import ( | |
| PROVIDER_CATALOG, | |
| SUPPORTED_PROVIDER_IDS, | |
| ) | |
| from free_claude_code.config.reasoning import ReasoningPreference | |
| from free_claude_code.config.settings import Settings | |
| from free_claude_code.core.anthropic import MessagesRequest, TokenCountRequest | |
| from free_claude_code.core.gateway_model_ids import decode_gateway_model_id | |
| from free_claude_code.core.reasoning import ReasoningPolicy | |
| from .reasoning import resolve_reasoning_policy | |
| _ROUTE_SETTINGS = ( | |
| ("fable", "model_fable", "reasoning_fable"), | |
| ("opus", "model_opus", "reasoning_opus"), | |
| ("haiku", "model_haiku", "reasoning_haiku"), | |
| ("sonnet", "model_sonnet", "reasoning_sonnet"), | |
| ) | |
| class ResolvedModel: | |
| original_model: str | |
| provider_id: str | |
| provider_model: str | |
| provider_model_ref: str | |
| reasoning_preference: ReasoningPreference | |
| class RoutedMessagesRequest: | |
| request: MessagesRequest | |
| resolved: ResolvedModel | |
| reasoning: ReasoningPolicy | |
| class RoutedTokenCountRequest: | |
| request: TokenCountRequest | |
| resolved: ResolvedModel | |
| class ModelRouter: | |
| """Resolve incoming Claude model names to configured provider/model pairs.""" | |
| def __init__(self, settings: Settings): | |
| self._settings = settings | |
| def resolve(self, claude_model_name: str) -> ResolvedModel: | |
| ( | |
| direct_provider_id, | |
| direct_provider_model, | |
| force_reasoning_off, | |
| ) = self._direct_provider_model(claude_model_name) | |
| if direct_provider_id is not None and direct_provider_model is not None: | |
| reasoning_preference = ( | |
| ReasoningPreference.OFF | |
| if force_reasoning_off | |
| else self._settings.reasoning_policy | |
| ) | |
| logger.debug( | |
| "MODEL DIRECT: '{}' -> provider='{}' model='{}' reasoning={}", | |
| claude_model_name, | |
| direct_provider_id, | |
| direct_provider_model, | |
| reasoning_preference.value, | |
| ) | |
| return ResolvedModel( | |
| original_model=claude_model_name, | |
| provider_id=direct_provider_id, | |
| provider_model=direct_provider_model, | |
| provider_model_ref=claude_model_name, | |
| reasoning_preference=reasoning_preference, | |
| ) | |
| provider_model_ref = self._resolve_model_ref(claude_model_name) | |
| reasoning_preference = self._resolve_reasoning_preference(claude_model_name) | |
| provider_id = parse_provider_type(provider_model_ref) | |
| self._validate_provider_id(provider_id) | |
| provider_model = parse_model_name(provider_model_ref) | |
| if provider_model != claude_model_name: | |
| logger.debug( | |
| "MODEL MAPPING: '{}' -> '{}'", claude_model_name, provider_model | |
| ) | |
| return ResolvedModel( | |
| original_model=claude_model_name, | |
| provider_id=provider_id, | |
| provider_model=provider_model, | |
| provider_model_ref=provider_model_ref, | |
| reasoning_preference=reasoning_preference, | |
| ) | |
| def _validate_provider_id(provider_id: str) -> None: | |
| if provider_id not in PROVIDER_CATALOG: | |
| raise UnknownProviderError.for_provider(provider_id, PROVIDER_CATALOG) | |
| def _direct_provider_model( | |
| self, model_name: str | |
| ) -> tuple[str | None, str | None, bool]: | |
| decoded = decode_gateway_model_id(model_name) | |
| if decoded is not None: | |
| if decoded.provider_id not in SUPPORTED_PROVIDER_IDS: | |
| return None, None, False | |
| return ( | |
| decoded.provider_id, | |
| decoded.provider_model, | |
| decoded.force_reasoning_off, | |
| ) | |
| provider_id, separator, provider_model = model_name.partition("/") | |
| if not separator: | |
| return None, None, False | |
| if provider_id not in SUPPORTED_PROVIDER_IDS: | |
| return None, None, False | |
| if not provider_model: | |
| return None, None, False | |
| return provider_id, provider_model, False | |
| def _resolve_model_ref(self, claude_model_name: str) -> str: | |
| """Resolve a Claude model name to the configured provider/model ref.""" | |
| route = self._matched_route(claude_model_name) | |
| if route is not None: | |
| model = getattr(self._settings, route[1]) | |
| if isinstance(model, str): | |
| return model | |
| return self._settings.model | |
| def _resolve_reasoning_preference( | |
| self, claude_model_name: str | |
| ) -> ReasoningPreference: | |
| """Resolve a route override without inspecting the provider model.""" | |
| route = self._matched_route(claude_model_name) | |
| if route is not None: | |
| preference = getattr(self._settings, route[2]) | |
| if preference is not ReasoningPreference.INHERIT: | |
| return preference | |
| return self._settings.reasoning_policy | |
| def _matched_route(model_name: str) -> tuple[str, str, str] | None: | |
| normalized = model_name.lower() | |
| return next( | |
| (route for route in _ROUTE_SETTINGS if route[0] in normalized), | |
| None, | |
| ) | |
| def resolve_messages_request( | |
| self, request: MessagesRequest | |
| ) -> RoutedMessagesRequest: | |
| """Return an internal routed request context.""" | |
| resolved = self.resolve(request.model) | |
| routed = request.model_copy(deep=True) | |
| routed.model = resolved.provider_model | |
| return RoutedMessagesRequest( | |
| request=routed, | |
| resolved=resolved, | |
| reasoning=resolve_reasoning_policy( | |
| routed, | |
| resolved.reasoning_preference, | |
| ), | |
| ) | |
| def resolve_token_count_request( | |
| self, request: TokenCountRequest | |
| ) -> RoutedTokenCountRequest: | |
| """Return an internal token-count request context.""" | |
| resolved = self.resolve(request.model) | |
| routed = request.model_copy( | |
| update={"model": resolved.provider_model}, deep=True | |
| ) | |
| return RoutedTokenCountRequest(request=routed, resolved=resolved) | |