dheraingoud's picture
feat: sync upstream commits up to f17c92bc
a1bab2d
Raw
History Blame Contribute Delete
6.62 kB
"""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"),
)
@dataclass(frozen=True, slots=True)
class ResolvedModel:
original_model: str
provider_id: str
provider_model: str
provider_model_ref: str
reasoning_preference: ReasoningPreference
@dataclass(frozen=True, slots=True)
class RoutedMessagesRequest:
request: MessagesRequest
resolved: ResolvedModel
reasoning: ReasoningPolicy
@dataclass(frozen=True, slots=True)
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,
)
@staticmethod
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
@staticmethod
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)