Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
model_router.py384 linesDownload Raw Back to api
1"""Model routing for Claude-compatible requests."""2 3from __future__ import annotations4 5from dataclasses import dataclass6 7from loguru import logger8 9from config.provider_ids import SUPPORTED_PROVIDER_IDS10from config.settings import Settings11from core.model_capabilities import find_best_model_for_task12from core.session_tracker import SessionTracker13from core.task_detector import TaskDetector14from providers.rate_limit import GlobalRateLimiter15 16from .gateway_model_ids import decode_gateway_model_id17from .models.anthropic import MessagesRequest, TokenCountRequest18 19# Default NIM models to include in auto routing (in order of preference)20DEFAULT_NIM_AUTO_MODELS = [21    "nvidia_nim/qwen/qwen3-coder-480b-a35b-instruct",22    "nvidia_nim/z-ai/glm4.7",23    "nvidia_nim/stepfun-ai/step-3.5-flash",24    "nvidia_nim/mistralai/mistral-large-3-675b-instruct-2512",25    "nvidia_nim/abacusai/dracarys-llama-3.1-70b-instruct",26    "nvidia_nim/bytedance/seed-oss-36b-instruct",27    "nvidia_nim/mistralai/mistral-nemotron",28]29 30 31@dataclass(frozen=True, slots=True)32class ResolvedModel:33    original_model: str34    provider_id: str35    provider_model: str36    provider_model_ref: str37    thinking_enabled: bool38 39 40@dataclass(frozen=True, slots=True)41class RoutedMessagesRequest:42    request: MessagesRequest43    resolved: ResolvedModel44 45 46@dataclass(frozen=True, slots=True)47class RoutedTokenCountRequest:48    request: TokenCountRequest49    resolved: ResolvedModel50 51 52class ModelRouter:53    """Resolve incoming Claude model names to configured provider/model pairs."""54 55    def __init__(self, settings: Settings):56        self._settings = settings57 58    def _is_auto(self, model_name: str) -> bool:59        """Return whether the model name refers to the virtual 'auto' model."""60        name_lower = model_name.lower()61        return name_lower == "auto" or name_lower == "anthropic/auto"62 63    def _normalize_candidate_ref(self, raw_ref: str) -> str | None:64        """Normalize auto candidate refs to ``provider/model`` when possible."""65        candidate = (raw_ref or "").strip()66        if not candidate:67            return None68 69        provider_id, separator, remainder = candidate.partition("/")70        if separator and provider_id in SUPPORTED_PROVIDER_IDS and remainder:71            return f"{provider_id}/{remainder}"72 73        # Treat bare model ids and vendor/model ids as NVIDIA NIM models.74        return f"nvidia_nim/{candidate}"75 76    def resolve(self, claude_model_name: str) -> ResolvedModel:77        # Special virtual model 'auto' maps to the configured default MODEL and78        # enables provider-side fallbacks. Resolve it to the configured model79        # while preserving the original requested name.80        if self._is_auto(claude_model_name):81            # If the user configured an explicit AUTO_MODEL_ORDER, try each82            # provider/model pair in order and pick the first provider that is83            # plausibly configured. Fall back to the single configured MODEL.84            order_csv = (self._settings.auto_model_order or "").strip()85            if order_csv:86                for cand in [c.strip() for c in order_csv.split(",") if c.strip()]:87                    if "/" not in cand:88                        # assume vendor-prefixed entries; skip malformed89                        continue90                    provider_id = Settings.parse_provider_type(cand)91                    provider_model = Settings.parse_model_name(cand)92                    if self._settings.provider_is_configured(provider_id):93                        thinking_enabled = self._settings.resolve_thinking(94                            claude_model_name95                        )96                        return ResolvedModel(97                            original_model=claude_model_name,98                            provider_id=provider_id,99                            provider_model=provider_model,100                            provider_model_ref=cand,101                            thinking_enabled=thinking_enabled,102                        )103            # No explicit order matched or none configured — fall back to default MODEL104            provider_model_ref = self._settings.model105            provider_id = Settings.parse_provider_type(provider_model_ref)106            provider_model = Settings.parse_model_name(provider_model_ref)107            thinking_enabled = self._settings.resolve_thinking(claude_model_name)108            return ResolvedModel(109                original_model=claude_model_name,110                provider_id=provider_id,111                provider_model=provider_model,112                provider_model_ref=provider_model_ref,113                thinking_enabled=thinking_enabled,114            )115 116        (117            direct_provider_id,118            direct_provider_model,119            force_thinking_enabled,120        ) = self._direct_provider_model(claude_model_name)121        if direct_provider_id is not None and direct_provider_model is not None:122            thinking_enabled = (123                force_thinking_enabled124                if force_thinking_enabled is not None125                else self._settings.resolve_thinking(direct_provider_model)126            )127            logger.debug(128                "MODEL DIRECT: '{}' -> provider='{}' model='{}' thinking={}",129                claude_model_name,130                direct_provider_id,131                direct_provider_model,132                thinking_enabled,133            )134            return ResolvedModel(135                original_model=claude_model_name,136                provider_id=direct_provider_id,137                provider_model=direct_provider_model,138                provider_model_ref=claude_model_name,139                thinking_enabled=thinking_enabled,140            )141 142        provider_model_ref = self._settings.resolve_model(claude_model_name)143        thinking_enabled = self._settings.resolve_thinking(claude_model_name)144        provider_id = Settings.parse_provider_type(provider_model_ref)145        provider_model = Settings.parse_model_name(provider_model_ref)146        if provider_model != claude_model_name:147            logger.debug(148                "MODEL MAPPING: '{}' -> '{}'", claude_model_name, provider_model149            )150        return ResolvedModel(151            original_model=claude_model_name,152            provider_id=provider_id,153            provider_model=provider_model,154            provider_model_ref=provider_model_ref,155            thinking_enabled=thinking_enabled,156        )157 158    def resolve_candidates(self, claude_model_name: str) -> list[ResolvedModel]:159        """Resolve a model name to a prioritized list of candidates.160 161        Used by the 'auto' routing logic to implement provider-side failover.162        Considers session load for fair resource sharing across multiple clients.163 164        Priority order:165        1. AUTO_MODEL_ORDER (if configured)166        2. MODEL (primary)167        3. NVIDIA NIM fallback models (if configured, or DEFAULT_NIM_AUTO_MODELS)168        4. MODEL_OPUS, MODEL_SONNET, MODEL_HAIKU169        """170        if not self._is_auto(claude_model_name):171            return [self.resolve(claude_model_name)]172 173        healthy_candidates: list[ResolvedModel] = []174        blocked_candidates: list[ResolvedModel] = []175        seen: set[str] = set()176        session_tracker = SessionTracker.get_instance()177 178        def add_candidate(ref: str | None, source: str) -> None:179            normalized_ref = self._normalize_candidate_ref(ref or "")180            if normalized_ref is None or normalized_ref in seen:181                return182            provider_id = Settings.parse_provider_type(normalized_ref)183            provider_model = Settings.parse_model_name(normalized_ref)184            if self._settings.provider_is_configured(provider_id):185                seen.add(normalized_ref)186                resolved = ResolvedModel(187                    original_model=claude_model_name,188                    provider_id=provider_id,189                    provider_model=provider_model,190                    provider_model_ref=normalized_ref,191                    thinking_enabled=self._settings.resolve_thinking(claude_model_name),192                )193 194                limiter = GlobalRateLimiter.get_scoped_instance(provider_id)195                is_blocked = limiter.is_blocked()196 197                # For Zen provider, never consider it blocked (no rate limits)198                if provider_id == "zen":199                    is_blocked = False200 201                # Check model health (recent failures)202                is_healthy = limiter.is_healthy(normalized_ref)203 204                if is_blocked or not is_healthy:205                    reason = "BLOCKED" if is_blocked else "UNHEALTHY"206                    logger.debug(207                        "Routing: candidate '{}' (from {}) is {} (health={})",208                        normalized_ref,209                        source,210                        reason,211                        is_healthy,212                    )213                    blocked_candidates.append(resolved)214                else:215                    # Smart ordering: Zen (no rate limits) gets priority, then by load216                    logger.debug(217                        "Routing: added candidate '{}' (from {})",218                        normalized_ref,219                        source,220                    )221                    healthy_candidates.append(resolved)222 223            else:224                logger.debug(225                    "Routing: candidate '{}' (from {}) is NOT CONFIGURED",226                    normalized_ref,227                    source,228                )229 230        # 1. AUTO_MODEL_ORDER (user-configured priority)231        order_csv = (self._settings.auto_model_order or "").strip()232        if order_csv:233            for cand in [c.strip() for c in order_csv.split(",") if c.strip()]:234                add_candidate(cand, "AUTO_MODEL_PRIORITY")235 236        # 2. Primary MODEL237        add_candidate(self._settings.model, "MODEL")238 239        # 3. NVIDIA Fallbacks - use configured or defaults240        nim_csv = (self._settings.nvidia_nim_fallback_models or "").strip()241        if nim_csv:242            for cand in [c.strip() for c in nim_csv.split(",") if c.strip()]:243                add_candidate(cand, "NVIDIA_NIM_FALLBACK_MODELS")244        else:245            # Use default NIM models when no explicit fallback configured246            for cand in DEFAULT_NIM_AUTO_MODELS:247                add_candidate(cand, "DEFAULT_NIM_AUTO_MODELS")248 249        # 4. Model-specific overrides250        add_candidate(self._settings.model_opus, "MODEL_OPUS")251        add_candidate(self._settings.model_sonnet, "MODEL_SONNET")252        add_candidate(self._settings.model_haiku, "MODEL_HAIKU")253 254        # Smart ordering: Zen goes first (no rate limits), then sort by load255        def provider_priority(c: ResolvedModel) -> tuple:256            # Priority: zen > others, then by active request count257            is_zen = 0 if c.provider_id == "zen" else 1258            active = session_tracker._provider_active.get(c.provider_id, 0)259            return (is_zen, active)260 261        healthy_candidates.sort(key=provider_priority)262 263        all_candidates = healthy_candidates + blocked_candidates264        logger.info(265            "Routing: resolved '{}' to {} candidates: {}",266            claude_model_name,267            len(all_candidates),268            ", ".join(c.provider_model_ref for c in all_candidates),269        )270        return all_candidates271 272    def _direct_provider_model(273        self, model_name: str274    ) -> tuple[str | None, str | None, bool | None]:275        decoded = decode_gateway_model_id(model_name)276        if decoded is not None:277            if decoded.provider_id not in SUPPORTED_PROVIDER_IDS:278                return None, None, None279            return (280                decoded.provider_id,281                decoded.provider_model,282                decoded.force_thinking_enabled,283            )284 285        provider_id, separator, provider_model = model_name.partition("/")286        if not separator:287            return None, None, None288        if provider_id not in SUPPORTED_PROVIDER_IDS:289            return None, None, None290        if not provider_model:291            return None, None, None292        return provider_id, provider_model, None293 294    def resolve_messages_request(295        self, request: MessagesRequest296    ) -> RoutedMessagesRequest:297        """Return an internal routed request context."""298        resolved = self.resolve(request.model)299        routed = request.model_copy(deep=True)300        routed.model = resolved.provider_model301        return RoutedMessagesRequest(request=routed, resolved=resolved)302 303    def resolve_token_count_request(304        self, request: TokenCountRequest305    ) -> RoutedTokenCountRequest:306        """Return an internal token-count request context."""307        resolved = self.resolve(request.model)308        routed = request.model_copy(309            update={"model": resolved.provider_model}, deep=True310        )311        return RoutedTokenCountRequest(request=routed, resolved=resolved)312 313    def resolve_with_task_awareness(314        self,315        claude_model_name: str,316        messages: list,317    ) -> ResolvedModel:318        """Resolve model with task-based capability matching.319 320        For 'auto' model, detects task requirements and routes to best-capable model.321        """322        if not self._is_auto(claude_model_name):323            return self.resolve(claude_model_name)324 325        # Detect what capabilities are needed326        detector = TaskDetector()327        requirements = detector.detect_requirements(messages)328 329        logger.info(330            "Task-aware routing: detected requirements={} confidence={:.2f}",331            requirements.required_capabilities,332            requirements.confidence,333        )334 335        # Get available candidates336        candidates = self.resolve_candidates(claude_model_name)337 338        if not candidates:339            # Fallback to default340            return self.resolve(claude_model_name)341 342        # If confidence is low or only general text needed, use load-based selection343        if requirements.confidence < 0.7 or (344            not requirements.requires_vision345            and not requirements.requires_coding346            and not requirements.requires_reasoning347        ):348            logger.debug(349                "Task-aware routing: low confidence, using load-based selection"350            )351            return candidates[0]352 353        # Find best model matching required capabilities354        required_caps = set()355        if requirements.requires_coding:356            required_caps.add("coding")357        if requirements.requires_reasoning:358            required_caps.add("reasoning")359        if requirements.requires_vision:360            required_caps.add("vision")361 362        if required_caps:363            model_refs = [c.provider_model_ref for c in candidates]364            best = find_best_model_for_task(required_caps, model_refs)365            if best:366                # Find the matching candidate367                for cand in candidates:368                    if cand.provider_model_ref == best.model_ref:369                        logger.info(370                            "Task-aware routing: selected {} for capabilities={}",371                            best.model_ref,372                            required_caps,373                        )374                        return cand375 376        # Default to first candidate (load-balanced)377        return candidates[0]378 379    def get_routing_hint(self, messages: list) -> str:380        """Get a hint about what kind of model would be best."""381        detector = TaskDetector()382        requirements = detector.detect_requirements(messages)383        return detector.get_priority_hint(requirements)384