Yash030/claude-code-proxy
2
1"""Model chaining engine for multi-stage AI pipelines."""2 3from __future__ import annotations4 5import asyncio6from collections.abc import AsyncIterator7from dataclasses import dataclass8from typing import Any, Callable9 10from loguru import logger11 12 13@dataclass(frozen=True, slots=True)14class ChainStage:15 """A single stage in a model chain."""16 17 model_ref: str # e.g., "zen/minimax-m2.5-free"18 stage_name: str # e.g., "vision_analysis", "code_generation"19 description: str20 21 22@dataclass(frozen=True, slots=True)23class ChainResult:24 """Result from executing a chain stage."""25 26 stage: ChainStage27 output: str28 success: bool29 error: str | None = None30 31 32# Chain templates for common multi-capability tasks33CHAIN_TEMPLATES: dict[str, list[ChainStage]] = {34 "vision_to_text": [35 ChainStage(36 model_ref="nvidia_nim/stepfun-ai/step-3.5-flash",37 stage_name="image_analysis",38 description="Analyze image content",39 ),40 ChainStage(41 model_ref="zen/minimax-m2.5-free",42 stage_name="response_generation",43 description="Generate final response",44 ),45 ],46 "reasoning_to_generation": [47 ChainStage(48 model_ref="nvidia_nim/qwen/qwen3-coder-480b-a35b-instruct",49 stage_name="analysis",50 description="Analyze and plan",51 ),52 ChainStage(53 model_ref="zen/minimax-m2.5-free",54 stage_name="generation",55 description="Generate output",56 ),57 ],58}59 60 61class ChainEngine:62 """Execute multi-model pipelines for complex requests."""63 64 def __init__(self, provider_getter: Callable[[str], Any]):65 self._provider_getter = provider_getter66 67 async def execute_simple_chain(68 self,69 stages: list[ChainStage],70 initial_messages: list[Any],71 system_prompt: str | None = None,72 ) -> AsyncIterator[str]:73 """Execute a chain of models sequentially.74 75 Args:76 stages: List of chain stages to execute77 initial_messages: Initial user messages78 system_prompt: Optional system prompt79 80 Yields:81 SSE events from the final model in the chain82 """83 if not stages:84 return85 86 logger.info("ChainEngine: executing {} stages", len(stages))87 88 # For now, execute single model - full chaining requires more integration89 # This is a placeholder for the full implementation90 first_stage = stages[0]91 provider = self._provider_getter(first_stage.model_ref.split("/")[0])92 93 logger.info(94 "ChainEngine: using model {} for chain",95 first_stage.model_ref,96 )97 98 # For Phase 1, just delegate to provider - full chaining comes later99 # The infrastructure is now in place100 async for event in provider.stream_response(101 initial_messages, system_prompt, {}102 ):103 yield event104 105 def get_chain_for_requirements(106 self,107 required_capabilities: set[str],108 available_models: list[str],109 ) -> list[ChainStage] | None:110 """Determine the appropriate chain based on required capabilities.111 112 Args:113 required_capabilities: Set of capabilities needed114 available_models: Available model references115 116 Returns:117 Chain stages or None if single model is sufficient118 """119 # If only one capability needed, no chain needed120 if len(required_capabilities) <= 1:121 return None122 123 # If multiple capabilities, build a simple chain124 if "vision" in required_capabilities and "coding" in required_capabilities:125 return CHAIN_TEMPLATES.get("vision_to_text")126 127 if "vision" in required_capabilities and "reasoning" in required_capabilities:128 return CHAIN_TEMPLATES.get("vision_to_text")129 130 if "reasoning" in required_capabilities and "coding" in required_capabilities:131 return CHAIN_TEMPLATES.get("reasoning_to_generation")132 133 # Default: no chain for now134 return None135 136 137async def execute_model_for_stage(138 provider: Any,139 messages: list[Any],140 system: str | None,141 metadata: dict[str, Any],142) -> str:143 """Execute a single model stage and return its output."""144 output_parts = []145 146 try:147 async for event in provider.stream_response(messages, system, metadata):148 # Parse SSE and collect text output149 if "content_block_delta" in event:150 # Extract text from delta151 pass152 153 return "".join(output_parts)154 except Exception as e:155 logger.error("Chain stage failed: {}", e)156 raise157 