Team Ai
Modelpublic

MrMoz33/tokioai-coder-iot

sourceHugging Facemitupdated 7d agoView on Hugging Face
0likes
pipeline.py477 linesDownload Raw Back to engine
1"""2TokioAI Engine v2.0 -- Main Pipeline3Orchestrates: Router -> Memory -> Guided Generation -> Verifier -> Output4 5The pipeline wraps any Ollama-compatible model and transforms it into6an elite AI agent with multi-step tool calling.7 8Key change from v1: when force_intent is None, the model freely chooses9between tool calls and text responses. This enables investigation chains.10"""11 12import json13import sys14import os15import re16import time17import requests18from typing import Any, Callable, Dict, List, Optional, Tuple19 20from .router import classify_intent, build_system_prompt, IntentType21from .memory import EngineMemory22from .guided import build_ollama_request, parse_guided_output23from .verifier import Verifier, RetryStrategy, VerificationResult24from .schemas import TOOL_DEFINITIONS, tools_to_prompt_text, validate_tool_call25 26 27class EngineConfig:28    """Pipeline configuration."""29    30    def __init__(31        self,32        ollama_host: str = None,33        model_name: str = "tokioai-coder",34        max_retries: int = 2,35        temperature: float = 0.3,36        max_tokens: int = 4096,37        use_guided_gen: bool = True,38        use_memory: bool = True,39        memory_path: str = None,40        verbose: bool = False,41    ):42        self.ollama_host = ollama_host or os.getenv(43            "TOKIOAI_ENGINE_HOST",44            os.getenv("OLLAMA_HOST", "http://localhost:11434")45        )46        self.model_name = model_name47        self.max_retries = max_retries48        self.temperature = temperature49        self.max_tokens = max_tokens50        self.use_guided_gen = use_guided_gen51        self.use_memory = use_memory52        self.memory_path = memory_path53        self.verbose = verbose54 55 56class PipelineResult:57    """Result from the pipeline."""58    59    def __init__(60        self,61        output_type: str,  # "tool_call", "text"62        tool_calls: List[Dict] = None,63        text: str = "",64        intent: str = "",65        confidence: float = 0.0,66        attempts: int = 1,67        latency_ms: int = 0,68        from_memory: bool = False,69    ):70        self.output_type = output_type71        self.tool_calls = tool_calls or []72        self.text = text73        self.intent = intent74        self.confidence = confidence75        self.attempts = attempts76        self.latency_ms = latency_ms77        self.from_memory = from_memory78 79    def to_openai_message(self) -> Dict:80        """Convert to OpenAI-style assistant message."""81        if self.output_type == "tool_call" and self.tool_calls:82            return {83                "role": "assistant",84                "content": None,85                "tool_calls": [86                    {87                        "id": f"call_{i}",88                        "type": "function",89                        "function": {90                            "name": tc["name"],91                            "arguments": json.dumps(tc["arguments"]),92                        },93                    }94                    for i, tc in enumerate(self.tool_calls)95                ],96            }97        return {98            "role": "assistant",99            "content": self.text,100        }101 102    def __repr__(self):103        if self.output_type == "tool_call":104            names = [tc["name"] for tc in self.tool_calls]105            return f"PipelineResult(tool_calls={names}, attempts={self.attempts}, {self.latency_ms}ms)"106        return f"PipelineResult(text={self.text[:80]!r}, attempts={self.attempts}, {self.latency_ms}ms)"107 108 109class Pipeline:110    """111    The TokioAI Engine pipeline.112    113    Usage:114        engine = Pipeline(EngineConfig(model_name="tokioai-coder"))115        result = engine.process("check disk space")116        # result.tool_calls = [{"name": "execute_local", "arguments": {"command": "df -h"}}]117    """118 119    def __init__(self, config: EngineConfig = None):120        self.config = config or EngineConfig()121        self.verifier = Verifier()122        self.retry_strategy = RetryStrategy(max_retries=self.config.max_retries)123        self.memory = EngineMemory(self.config.memory_path) if self.config.use_memory else None124        self.conversation_history: List[Dict] = []125        self._stats = {126            "total_requests": 0,127            "tool_calls": 0,128            "text_responses": 0,129            "retries": 0,130            "failures": 0,131            "total_latency_ms": 0,132        }133 134    def process(135        self,136        user_message: str,137        conversation_history: List[Dict] = None,138        force_intent: str = None,139    ) -> PipelineResult:140        """141        Process a user message through the full pipeline.142        143        Args:144            user_message: The user's input145            conversation_history: Previous messages (optional)146            force_intent: Override intent classification147                - "text": force text-only response (no tool calls allowed)148                - "tool_call": force tool call149                - None: model decides freely (default for investigation rounds)150            151        Returns:152            PipelineResult with tool calls or text response153        """154        start_time = time.time()155        self._stats["total_requests"] += 1156 157        history = conversation_history or self.conversation_history158 159        # === STAGE 1: ROUTE ===160        if force_intent:161            intent, confidence = force_intent, 1.0162        else:163            intent, confidence = classify_intent(user_message, history)164 165        if self.config.verbose:166            print(f"[ENGINE] Intent: {intent} (confidence: {confidence:.2f})")167 168        # === PURE TEXT SHORT-CIRCUIT ===169        # For greetings and thanks, strip conversation history to prevent170        # context bleeding (e.g. "hola" after HA investigation should NOT talk about HA).171        # BUT NOT for affirmations ("dale", "si", "ok") which need context.172        import re as _re173        _GREETING_PATS = [174            r'^(hi|hello|hey|hola|que\s+tal|buenos?\s+d[ií]as?|buenas?\s+(?:tardes?|noches?))\s*[.!?]*$',175            r'^(thanks?|thank\s+you|gracias|thx|de\s+nada)\s*[.!?]*$',176            r'^(como\s+(va|estas?|andas?|te\s+va)|how\s+are\s+you)\s*[.!?]*$',177            r'^(todo\s+bien|bien\s+y\s+tu)\s*[.!?]*$',178            r'\b(who\s+are\s+you|what\s+can\s+you|tell\s+me\s+about\s+yourself)\b',179        ]180        _msg_lower = user_message.lower().strip()181        is_pure_text = (182            intent == "text" and confidence >= 0.9183            and any(_re.match(p, _msg_lower, _re.IGNORECASE) for p in _GREETING_PATS)184        )185 186        # === STAGE 2: MEMORY LOOKUP ===187        few_shot = []188        if self.memory and self.config.use_memory and not is_pure_text:189            few_shot = self.memory.retrieve(user_message, intent, top_k=3)190            if self.config.verbose and few_shot:191                print(f"[ENGINE] Memory: {len(few_shot)} relevant examples found")192 193        # === STAGE 3: BUILD PROMPT ===194        system_prompt = build_system_prompt(195            intent=intent,196            few_shot_examples=few_shot,197        )198 199        # Build messages200        messages = [{"role": "system", "content": system_prompt}]201 202        # Add conversation history -- but for pure text (greetings etc),203        # only include the LAST 2 messages to avoid context bleeding204        if is_pure_text:205            # For greetings/thanks: minimal or no history206            # This prevents "hola" from generating an HA status report207            # because the previous conversation was about HA208            if self.config.verbose:209                print(f"[ENGINE] Pure text detected, stripping history ({len(history)} msgs)")210        else:211            for msg in history[-20:]:212                messages.append(msg)213 214        # Add current user message215        messages.append({"role": "user", "content": user_message})216 217        # === STAGE 4: GENERATE ===218        # Use guided generation ONLY for definite tool_call intent on first request219        # For investigation rounds (force_intent=None with history), let model choose220        use_guided = self.config.use_guided_gen and intent == "tool_call"221 222        result = self._generate_with_retries(223            messages=messages,224            intent=intent,225            use_guided=use_guided,226            force_text_only=(force_intent == "text"),227        )228 229        # === STAGE 5: RECORD ===230        latency_ms = int((time.time() - start_time) * 1000)231        result.intent = intent232        result.confidence = confidence233        result.latency_ms = latency_ms234 235        self._stats["total_latency_ms"] += latency_ms236 237        if result.output_type == "tool_call":238            self._stats["tool_calls"] += 1239        else:240            self._stats["text_responses"] += 1241 242        # Save to memory (only successful results)243        if self.memory and result.output_type in ("tool_call", "text"):244            tool_name = result.tool_calls[0]["name"] if result.tool_calls else None245            output_text = (246                json.dumps(result.tool_calls[0]) if result.tool_calls247                else result.text[:500]248            )249            self.memory.add(250                user_input=user_message,251                model_output=output_text,252                intent=intent,253                success=True,254                tool_name=tool_name,255                latency_ms=latency_ms,256            )257 258        # Update conversation history259        self.conversation_history.append({"role": "user", "content": user_message})260        self.conversation_history.append(result.to_openai_message())261 262        return result263 264    def add_tool_result(self, tool_name: str, result: str):265        """Add a tool execution result to the conversation."""266        self.conversation_history.append({267            "role": "tool",268            "content": result,269            "name": tool_name,270        })271 272    def _generate_with_retries(273        self,274        messages: List[Dict],275        intent: str,276        use_guided: bool,277        force_text_only: bool = False,278    ) -> PipelineResult:279        """280        Generate model output with retry logic on verification failure.281        282        Args:283            messages: Full message chain284            intent: Classified intent285            use_guided: Whether to use Ollama guided generation (JSON schema)286            force_text_only: If True, NEVER return tool calls (depth limit reached)287        """288 289        for attempt in range(self.config.max_retries + 1):290            # Build request291            if use_guided:292                request = build_ollama_request(293                    model=self.config.model_name,294                    messages=messages,295                    intent=intent,296                    force_json=True,297                    temperature=self.config.temperature,298                    max_tokens=self.config.max_tokens,299                )300            else:301                # Free-form generation: model can output tool calls OR text302                request = {303                    "model": self.config.model_name,304                    "messages": messages,305                    "stream": False,306                    "options": {307                        "temperature": self.config.temperature,308                        "num_predict": self.config.max_tokens,309                    },310                }311                # Include tools so Ollama can use native tool calling312                # UNLESS we're forcing text-only mode313                if not force_text_only and intent != "text":314                    request["tools"] = TOOL_DEFINITIONS315 316            # Call Ollama317            try:318                raw_output, native_tool_calls = self._call_ollama(request)319            except Exception as e:320                if self.config.verbose:321                    print(f"[ENGINE] Ollama error: {e}")322                self._stats["failures"] += 1323                return PipelineResult(324                    output_type="text",325                    text=f"Error calling model: {e}",326                    attempts=attempt + 1,327                )328 329            if self.config.verbose:330                print(f"[ENGINE] Raw output (attempt {attempt+1}): {raw_output[:200]}")331                if native_tool_calls:332                    print(f"[ENGINE] Native tool calls: {native_tool_calls}")333 334            # === NATIVE TOOL CALLS (Ollama returned tool_calls in response) ===335            if native_tool_calls and not force_text_only:336                valid_calls = []337                for tc in native_tool_calls:338                    is_valid, err = validate_tool_call(tc["name"], tc["arguments"])339                    if is_valid:340                        valid_calls.append(tc)341                    elif self.config.verbose:342                        print(f"[ENGINE] Invalid native tool call: {err}")343                344                if valid_calls:345                    return PipelineResult(346                        output_type="tool_call",347                        tool_calls=valid_calls,348                        attempts=attempt + 1,349                    )350 351            # === GUIDED OUTPUT (JSON constrained) ===352            if use_guided and intent in ("tool_call", "decision"):353                parsed = parse_guided_output(raw_output, intent)354                if parsed["type"] == "tool_call" and not force_text_only:355                    name = parsed.get("name", "")356                    args = parsed.get("arguments", {})357                    is_valid, err = validate_tool_call(name, args)358                    if is_valid:359                        return PipelineResult(360                            output_type="tool_call",361                            tool_calls=[{"name": name, "arguments": args}],362                            attempts=attempt + 1,363                        )364                    elif self.config.verbose:365                        print(f"[ENGINE] Guided validation failed: {err}")366                elif parsed["type"] == "text":367                    return PipelineResult(368                        output_type="text",369                        text=parsed.get("response", raw_output),370                        attempts=attempt + 1,371                    )372 373            # === FREE-FORM VERIFICATION ===374            verification = self.verifier.verify(raw_output, expected_intent=intent)375 376            if verification.valid:377                if verification.output_type == "tool_call" and not force_text_only:378                    # Model wants to call a tool -- allow it379                    return PipelineResult(380                        output_type="tool_call",381                        tool_calls=verification.tool_calls,382                        attempts=attempt + 1,383                    )384                else:385                    # Text response (or forced text-only)386                    text = verification.text or raw_output387                    # Clean up: if model tried to output tool call JSON but we're388                    # in text-only mode, don't show the raw JSON389                    if force_text_only and verification.output_type == "tool_call":390                        text = (391                            "Based on my investigation, I was unable to gather "392                            "additional information. Here is what I found so far."393                        )394                    return PipelineResult(395                        output_type="text",396                        text=text,397                        attempts=attempt + 1,398                    )399 400            # Failed verification -- retry?401            if self.retry_strategy.should_retry(attempt, verification):402                self._stats["retries"] += 1403                retry_msg = self.retry_strategy.get_retry_message(attempt, verification)404                messages.append({"role": "assistant", "content": raw_output})405                messages.append({"role": "user", "content": retry_msg})406                if self.config.verbose:407                    print(f"[ENGINE] Retrying (attempt {attempt+2})...")408                continue409            else:410                # Give up -- return whatever we got411                self._stats["failures"] += 1412                return PipelineResult(413                    output_type="text",414                    text=verification.text or raw_output,415                    attempts=attempt + 1,416                )417 418        # Should never reach here419        return PipelineResult(output_type="text", text="Engine error: max retries exceeded")420 421    def _call_ollama(self, request: Dict):422        """423        Call Ollama API and return (text_output, native_tool_calls).424        425        Returns:426            Tuple of (str, list):427                - text content from the response428                - list of native tool calls (if any), each as {name, arguments}429        """430        url = f"{self.config.ollama_host}/api/chat"431        432        try:433            resp = requests.post(url, json=request, timeout=120)434            if self.config.verbose and resp.status_code != 200:435                print(f"[ENGINE] Ollama HTTP {resp.status_code}: {resp.text[:500]}")436            resp.raise_for_status()437            data = resp.json()438        except requests.exceptions.ConnectionError:439            raise RuntimeError(f"Cannot connect to Ollama at {self.config.ollama_host}")440        except requests.exceptions.Timeout:441            raise RuntimeError("Ollama request timed out (120s)")442        except Exception as e:443            raise RuntimeError(f"Ollama request failed: {e}")444 445        message = data.get("message", {})446        content = message.get("content", "")447 448        # Check for native tool calls from Ollama449        native_tool_calls = []450        tool_calls_raw = message.get("tool_calls", [])451        if tool_calls_raw:452            for tc in tool_calls_raw:453                func = tc.get("function", {})454                native_tool_calls.append({455                    "name": func.get("name", ""),456                    "arguments": func.get("arguments", {}),457                })458 459        return content, native_tool_calls460 461    def get_stats(self) -> Dict:462        """Return pipeline statistics."""463        total = self._stats["total_requests"] or 1464        return {465            "total_requests": self._stats["total_requests"],466            "tool_calls": self._stats["tool_calls"],467            "text_responses": self._stats["text_responses"],468            "retries": self._stats["retries"],469            "failures": self._stats["failures"],470            "avg_latency_ms": self._stats["total_latency_ms"] // total,471            "tool_call_rate": f"{100 * self._stats['tool_calls'] // total}%",472        }473 474    def reset(self):475        """Reset conversation history."""476        self.conversation_history = []477