Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
chain_engine.py157 linesDownload Raw Back to core
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