Team Ai
Modelpublic

MrMoz33/tokioai-coder-iot

sourceHugging Facemitupdated 7d agoView on Hugging Face
0likes
guided.py229 linesDownload Raw Back to engine
1"""2TokioAI Engine -- Guided Generation3Forces the model to produce valid structured output using Ollama's format parameter.4 5Instead of hoping the model generates valid JSON, we CONSTRAIN it.6Ollama supports JSON mode natively, and we combine that with our tool schemas7to guarantee valid output.8 9For non-Ollama backends, we use post-processing + retry as fallback.10"""11 12import json13from typing import Any, Dict, List, Optional14 15from .schemas import TOOL_DEFINITIONS, TOOL_SCHEMAS, VALID_TOOL_NAMES16 17 18def get_tool_json_schema() -> Dict:19    """20    Generate a JSON schema that constrains model output to valid tool calls.21    This is used with Ollama's `format` parameter for structured output.22    """23    # Build a union schema: the model must output one of the valid tool call shapes24    tool_schemas = []25    26    for tool_def in TOOL_DEFINITIONS:27        func = tool_def["function"]28        tool_schemas.append({29            "type": "object",30            "properties": {31                "name": {32                    "type": "string",33                    "const": func["name"],34                },35                "arguments": func["parameters"],36            },37            "required": ["name", "arguments"],38            "additionalProperties": False,39        })40 41    # anyOf: model can output any one of these tool call shapes42    return {43        "type": "object",44        "anyOf": tool_schemas,45    }46 47 48def get_simple_json_schema() -> Dict:49    """50    Simplified schema: just name (enum) + arguments (object).51    More compatible with Ollama's format support.52    """53    return {54        "type": "object",55        "properties": {56            "name": {57                "type": "string",58                "enum": sorted(VALID_TOOL_NAMES),59            },60            "arguments": {61                "type": "object",62            },63        },64        "required": ["name", "arguments"],65    }66 67 68def get_text_response_schema() -> Dict:69    """70    Schema for text-only responses (no tool calls).71    Forces the model to output a simple response object.72    """73    return {74        "type": "object",75        "properties": {76            "response": {77                "type": "string",78                "description": "Your text response to the user",79            },80            "needs_tool": {81                "type": "boolean",82                "description": "Set to true if you think a tool should be used instead",83            },84        },85        "required": ["response"],86    }87 88 89def get_decision_schema() -> Dict:90    """91    Schema for decision mode: model chooses between tool call or text response.92    """93    return {94        "type": "object",95        "properties": {96            "action": {97                "type": "string",98                "enum": ["tool_call", "text_response"],99                "description": "Whether to call a tool or respond with text",100            },101            "tool_call": {102                "type": "object",103                "properties": {104                    "name": {105                        "type": "string",106                        "enum": sorted(VALID_TOOL_NAMES),107                    },108                    "arguments": {109                        "type": "object",110                    },111                },112                "description": "If action is tool_call, provide the tool call here",113            },114            "text_response": {115                "type": "string",116                "description": "If action is text_response, provide your response here",117            },118        },119        "required": ["action"],120    }121 122 123def build_ollama_request(124    model: str,125    messages: List[Dict],126    intent: str,127    force_json: bool = True,128    temperature: float = 0.3,129    max_tokens: int = 4096,130    stream: bool = False,131) -> Dict:132    """133    Build an Ollama /api/chat request with optional guided generation.134    135    Args:136        model: Ollama model name137        messages: Chat messages138        intent: Classified intent (determines schema)139        force_json: Whether to constrain output to JSON140        temperature: Sampling temperature141        max_tokens: Max tokens to generate142        stream: Whether to stream the response143        144    Returns:145        Request body dict ready for POST to /api/chat146    """147    request = {148        "model": model,149        "messages": messages,150        "stream": stream,151        "options": {152            "temperature": temperature,153            "num_predict": max_tokens,154        },155    }156 157    if force_json and intent in ("tool_call", "decision"):158        # Use Ollama's structured output159        if intent == "tool_call":160            request["format"] = get_simple_json_schema()161        elif intent == "decision":162            request["format"] = get_decision_schema()163    elif force_json and intent == "text":164        # For text, we DON'T force JSON -- let the model speak naturally165        pass166 167    # Add tools in OpenAI format (Ollama supports this)168    if intent in ("tool_call", "decision"):169        request["tools"] = TOOL_DEFINITIONS170 171    return request172 173 174def parse_guided_output(raw_output: str, intent: str) -> Dict:175    """176    Parse the guided (JSON-constrained) output from the model.177    178    Returns a normalized dict:179    - For tool_call: {"type": "tool_call", "name": str, "arguments": dict}180    - For text: {"type": "text", "response": str}181    - For decision: {"type": "tool_call"|"text", ...}182    """183    try:184        data = json.loads(raw_output)185    except json.JSONDecodeError:186        # Try to extract JSON from text187        import re188        match = re.search(r'\{.*\}', raw_output, re.DOTALL)189        if match:190            try:191                data = json.loads(match.group())192            except json.JSONDecodeError:193                return {"type": "text", "response": raw_output}194        else:195            return {"type": "text", "response": raw_output}196 197    # Decision format198    if "action" in data:199        if data["action"] == "tool_call" and "tool_call" in data:200            tc = data["tool_call"]201            return {202                "type": "tool_call",203                "name": tc.get("name", ""),204                "arguments": tc.get("arguments", {}),205            }206        elif data["action"] == "text_response":207            return {208                "type": "text",209                "response": data.get("text_response", ""),210            }211 212    # Direct tool call format213    if "name" in data and data["name"] in VALID_TOOL_NAMES:214        return {215            "type": "tool_call",216            "name": data["name"],217            "arguments": data.get("arguments", {}),218        }219 220    # Text response format221    if "response" in data:222        return {223            "type": "text",224            "response": data["response"],225        }226 227    # Unknown format -- return as text228    return {"type": "text", "response": json.dumps(data)}229