#!/usr/bin/env python3 """ Momentry Tool Calling Module ============================ Supports sequential multi-tool execution using Ollama API. Tools: - query_postgres: PostgreSQL database queries - search_qdrant: Vector similarity search - execute_bash: Bash command execution - call_api: HTTP API calls Version: 1.1.0 Updated: 2026-07-26 Changes: - Added embedding server health check - Enhanced bash safety patterns - Added logging support - Environment variable for Qdrant collection """ import json import logging import os import re import subprocess import requests from typing import Any, Callable, Dict, List, Optional, Tuple from dataclasses import dataclass from datetime import datetime, date # Configure logging logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' ) logger = logging.getLogger('tool_caller') class DateTimeEncoder(json.JSONEncoder): """Custom JSON encoder for datetime objects""" def default(self, obj): if isinstance(obj, (datetime, date)): return obj.isoformat() return super().default(obj) @dataclass class ToolResult: """Result from tool execution""" success: bool data: Any error: Optional[str] = None execution_time_ms: float = 0 class ToolRegistry: """Registry for available tools""" def __init__(self): self._tools: Dict[str, Dict[str, Any]] = {} self._executors: Dict[str, Callable] = {} def register( self, name: str, description: str, parameters: Dict[str, Any], executor: Callable[[Dict[str, Any]], ToolResult] ): """Register a new tool""" self._tools[name] = { "type": "function", "function": { "name": name, "description": description, "parameters": parameters } } self._executors[name] = executor logger.debug(f"Registered tool: {name}") def get_definitions(self) -> List[Dict[str, Any]]: """Get all tool definitions for API calls""" return list(self._tools.values()) def has_tool(self, name: str) -> bool: """Check if tool exists""" return name in self._executors def execute(self, name: str, arguments: Dict[str, Any]) -> ToolResult: """Execute a tool by name""" if name not in self._executors: logger.error(f"Unknown tool: {name}") return ToolResult( success=False, data=None, error=f"Unknown tool: {name}" ) start_time = datetime.now() logger.info(f"Executing tool: {name} with args: {json.dumps(arguments, ensure_ascii=False)[:200]}") try: result = self._executors[name](arguments) elapsed = (datetime.now() - start_time).total_seconds() * 1000 result.execution_time_ms = elapsed if result.success: logger.info(f"Tool {name} succeeded in {elapsed:.1f}ms") else: logger.error(f"Tool {name} failed: {result.error}") return result except Exception as e: elapsed = (datetime.now() - start_time).total_seconds() * 1000 logger.error(f"Tool {name} exception: {e}") return ToolResult( success=False, data=None, error=str(e), execution_time_ms=elapsed ) class OllamaToolCaller: """Tool caller using Ollama API with improved stability""" def __init__( self, base_url: str = "http://localhost:11434", model: str = "llama3.1:8b", max_iterations: int = 10, max_tool_calls: int = 5 ): self.base_url = base_url self.model = model self.max_iterations = max_iterations self.max_tool_calls = max_tool_calls self.registry = ToolRegistry() # System prompt optimized for tool calling self.system_prompt = """You are a data assistant with access to tools. CORE RULES: 1. When you receive a tool result, you MUST provide a final answer - do NOT call more tools 2. Call exactly ONE tool per response 3. Answer based on tool results, not assumptions 4. If tool fails, report the error and stop IMPORTANT: After receiving ANY tool result, respond with a clear answer to the user.""" # Track tool calls for loop prevention self._tool_call_history: List[str] = [] logger.info(f"Initialized OllamaToolCaller with model={model}") def chat(self, messages: List[Dict[str, Any]]) -> Dict[str, Any]: """Send chat request to Ollama""" url = f"{self.base_url}/api/chat" payload = { "model": self.model, "messages": messages, "tools": self.registry.get_definitions(), "stream": False } logger.debug(f"Sending chat request to {url}") response = requests.post(url, json=payload, timeout=120) response.raise_for_status() return response.json() def execute_tool_call(self, tool_call: Dict[str, Any]) -> ToolResult: """Execute a single tool call with flexible argument parsing""" func = tool_call.get("function", {}) name = func.get("name", "") # Handle both dict and string arguments arguments = func.get("arguments", {}) if isinstance(arguments, str): try: arguments = json.loads(arguments) except json.JSONDecodeError: arguments = {} # Normalize arguments based on tool name arguments = self._normalize_arguments(name, arguments) return self.registry.execute(name, arguments) def _normalize_arguments(self, tool_name: str, args: Dict[str, Any]) -> Dict[str, Any]: """Normalize arguments to match expected format""" if not args: return args # For query_postgres, ensure we have 'query' parameter if tool_name == "query_postgres": if "query" not in args: for key in ["query_text", "sql", "sql_query", "statement"]: if key in args: args["query"] = args[key] break # For search_qdrant, ensure we have 'query_text' parameter elif tool_name == "search_qdrant": if "query_text" not in args: for key in ["query", "search", "text"]: if key in args: args["query_text"] = args[key] break return args def _get_tool_key(self, name: str, arguments: Any) -> str: """Generate unique key for tool call""" if isinstance(arguments, str): try: arguments = json.loads(arguments) except: arguments = {"raw": arguments} return f"{name}:{json.dumps(arguments, sort_keys=True)}" def _is_duplicate_call(self, tool_key: str) -> bool: """Check if this tool call was already made""" is_dup = tool_key in self._tool_call_history if is_dup: logger.warning(f"Duplicate tool call detected: {tool_key[:50]}") return is_dup def _add_to_history(self, tool_key: str): """Add tool call to history""" self._tool_call_history.append(tool_key) def _extract_tool_from_text(self, text: str) -> Optional[Tuple[str, Dict[str, Any]]]: """Extract tool call from text content (when model outputs JSON instead of tool_calls)""" if not text or '{"name":' not in text: return None try: start_idx = text.find('{"name":') if start_idx == -1: return None # Find the end of the JSON object depth = 0 end_idx = start_idx for i in range(start_idx, len(text)): if text[i] == '{': depth += 1 elif text[i] == '}': depth -= 1 if depth == 0: end_idx = i + 1 break if end_idx > start_idx: json_str = text[start_idx:end_idx] # Fix common JSON issues json_str = json_str.replace('\\ ', ' ') json_str = json_str.replace('\\n', '\n') tool_data = json.loads(json_str) if 'name' in tool_data: name = tool_data['name'] params = tool_data.get('parameters', tool_data.get('arguments', {})) if self.registry.has_tool(name): logger.debug(f"Extracted tool from text: {name}") return (name, params) except Exception as e: logger.debug(f"Failed to extract tool from text: {e}") return None def _force_answer(self, messages: List[Dict[str, Any]], tool_results: str) -> str: """Force the model to provide a final answer based on tool results""" messages.append({ "role": "user", "content": f"Based on the tool results above, please provide your final answer now. Do not call any more tools.\n\nTool results:\n{tool_results}" }) url = f"{self.base_url}/api/chat" payload = { "model": self.model, "messages": messages, "stream": False } try: logger.info("Forcing final answer from model") response = requests.post(url, json=payload, timeout=120) response.raise_for_status() return response.json().get("message", {}).get("content", "Unable to generate answer.") except Exception as e: logger.error(f"Failed to force answer: {e}") return f"Tool execution completed. Results: {tool_results}" def run(self, user_query: str) -> str: """ Run tool calling loop with user query. Args: user_query: The user's question or request Returns: Final text response from the model """ logger.info(f"Starting tool call loop for query: {user_query[:100]}") self._tool_call_history = [] messages = [ {"role": "system", "content": self.system_prompt}, {"role": "user", "content": user_query} ] tool_results_collected = [] for iteration in range(self.max_iterations): logger.info(f"Iteration {iteration + 1}/{self.max_iterations}") # Call LLM response = self.chat(messages) message = response.get("message", {}) content = message.get("content", "") tool_calls = message.get("tool_calls", []) # If no tool calls in response, check if content has text tool call if not tool_calls: extracted = self._extract_tool_from_text(content) if extracted: name, params = extracted tool_key = self._get_tool_key(name, params) if not self._is_duplicate_call(tool_key): self._add_to_history(tool_key) tool_call = {"function": {"name": name, "arguments": params}} tool_result = self.execute_tool_call(tool_call) # Build result string with error handling if tool_result.success: result_str = json.dumps({ "success": True, "data": tool_result.data }, ensure_ascii=False, cls=DateTimeEncoder) else: result_str = json.dumps({ "success": False, "error": tool_result.error, "suggestion": "Try a different tool or rephrase your query." }, ensure_ascii=False) messages.append({"role": "assistant", "content": "", "tool_calls": [tool_call]}) messages.append({"role": "user", "content": f"Tool result: {result_str}"}) tool_results_collected.append(f"{name}: {result_str[:500]}") continue # No tool call - return content if content: logger.info(f"Returning final answer after {iteration + 1} iterations") return content else: return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "No tool was called") # Execute first tool call (sequential) tool_call = tool_calls[0] func = tool_call.get("function", {}) tool_name = func.get("name", "") arguments = func.get("arguments", {}) tool_key = self._get_tool_key(tool_name, arguments) # Check for duplicate if self._is_duplicate_call(tool_key): return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "Tool already executed") # Track and execute self._add_to_history(tool_key) tool_result = self.execute_tool_call(tool_call) # Check if we've hit the tool call limit if len(self._tool_call_history) >= self.max_tool_calls: if tool_result.success: result_str = json.dumps({ "success": True, "data": tool_result.data }, ensure_ascii=False, cls=DateTimeEncoder) else: result_str = json.dumps({ "success": False, "error": tool_result.error, "suggestion": "Try a different tool or rephrase your query." }, ensure_ascii=False) tool_results_collected.append(f"{tool_name}: {result_str[:500]}") return self._force_answer(messages, "\n".join(tool_results_collected)) # Add to messages messages.append({"role": "assistant", "content": "", "tool_calls": [tool_call]}) if tool_result.success: result_str = json.dumps({ "success": True, "data": tool_result.data }, ensure_ascii=False, cls=DateTimeEncoder) else: result_str = json.dumps({ "success": False, "error": tool_result.error, "suggestion": "Try a different tool or rephrase your query." }, ensure_ascii=False) messages.append({"role": "user", "content": f"Tool result: {result_str}"}) tool_results_collected.append(f"{tool_name}: {result_str[:500]}") # Max iterations reached logger.warning(f"Max iterations ({self.max_iterations}) reached") return self._force_answer(messages, "\n".join(tool_results_collected) if tool_results_collected else "Max iterations reached") def register_default_tools(self, db_url: str = None, qdrant_url: str = None): """Register default tools with connection strings""" # Get default collection from environment variable default_collection = os.environ.get("QDRANT_DEFAULT_COLLECTION", "momentry_rule1") logger.info(f"Using default Qdrant collection: {default_collection}") # PostgreSQL tool def query_postgres(args: Dict[str, Any]) -> ToolResult: import psycopg2 import psycopg2.extras query = args.get("query", "") if not query: return ToolResult(success=False, data=None, error="No query provided") logger.info(f"Executing PostgreSQL query: {query[:100]}") conn = psycopg2.connect(db_url or "postgres://accusys@localhost:5432/momentry") try: with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur: cur.execute(query) if cur.description: rows = cur.fetchall() logger.info(f"Query returned {len(rows)} rows") return ToolResult( success=True, data={"rows": [dict(r) for r in rows], "row_count": len(rows)} ) else: conn.commit() logger.info(f"Query affected {cur.rowcount} rows") return ToolResult( success=True, data={"affected_rows": cur.rowcount} ) finally: conn.close() self.registry.register( name="query_postgres", description="Execute SQL query on PostgreSQL database. Use SELECT for queries, INSERT/UPDATE/DELETE for modifications.", parameters={ "type": "object", "properties": { "query": { "type": "string", "description": "SQL query to execute" } }, "required": ["query"] }, executor=query_postgres ) # Qdrant tool with health check def search_qdrant(args: Dict[str, Any]) -> ToolResult: import requests as req collection = args.get("collection", default_collection) query_text = args.get("query_text", "") limit = args.get("limit", 10) # Get Qdrant API key from args or environment api_key = args.get("api_key") or os.environ.get("QDRANT_API_KEY", "Test3200Test3200Test3200") if not query_text: return ToolResult(success=False, data=None, error="No query text provided") logger.info(f"Searching Qdrant collection '{collection}' for: {query_text[:50]}") # Check embedding server health try: embed_health = req.get("http://localhost:11436/health", timeout=5) if embed_health.status_code != 200: return ToolResult( success=False, data=None, error=f"Embedding server not healthy (status={embed_health.status_code})" ) except req.exceptions.ConnectionError: return ToolResult( success=False, data=None, error="Embedding server not available at http://localhost:11436" ) except Exception as e: return ToolResult( success=False, data=None, error=f"Embedding server health check failed: {str(e)}" ) # Get embedding try: embed_url = "http://localhost:11436/v1/embeddings" embed_resp = req.post(embed_url, json={"input": query_text}, timeout=30) embed_resp.raise_for_status() embed_data = embed_resp.json() # Handle both response formats if "data" in embed_data: embedding = embed_data["data"][0]["embedding"] elif "embeddings" in embed_data: embedding = embed_data["embeddings"][0] else: return ToolResult( success=False, data=None, error=f"Unexpected embedding response format: {list(embed_data.keys())}" ) logger.info(f"Generated embedding (dim={len(embedding)})") except Exception as e: return ToolResult( success=False, data=None, error=f"Embedding generation failed: {str(e)}" ) # Check Qdrant collection exists qdrant_url_base = qdrant_url or "http://localhost:6333" headers = {"api-key": api_key} try: collections_url = f"{qdrant_url_base}/collections" coll_resp = req.get(collections_url, headers=headers, timeout=10) coll_resp.raise_for_status() collections = [c["name"] for c in coll_resp.json().get("result", {}).get("collections", [])] logger.info(f"Available Qdrant collections: {collections}") if collection not in collections: return ToolResult( success=False, data=None, error=f"Collection '{collection}' not found. Available: {', '.join(collections)}" ) except Exception as e: return ToolResult( success=False, data=None, error=f"Qdrant connection failed: {str(e)}" ) # Execute search search_url = f"{qdrant_url_base}/collections/{collection}/points/search" search_payload = {"vector": embedding, "limit": limit, "with_payload": True} try: search_resp = req.post(search_url, json=search_payload, headers=headers, timeout=30) search_resp.raise_for_status() results = search_resp.json().get("result", []) logger.info(f"Qdrant returned {len(results)} matches") except Exception as e: return ToolResult( success=False, data=None, error=f"Qdrant search failed: {str(e)}" ) return ToolResult( success=True, data={ "matches": [ {"id": r.get("id"), "score": r.get("score"), "payload": r.get("payload", {})} for r in results ], "match_count": len(results) } ) self.registry.register( name="search_qdrant", description=f"Search for similar vectors in Qdrant collection. Default collection: {default_collection}", parameters={ "type": "object", "properties": { "collection": { "type": "string", "description": f"Qdrant collection name (default: {default_collection})", "default": default_collection }, "query_text": { "type": "string", "description": "Search query text (will be embedded)" }, "limit": { "type": "integer", "description": "Maximum number of results", "default": 10 } }, "required": ["query_text"] }, executor=search_qdrant ) # Bash tool with enhanced safety def execute_bash(args: Dict[str, Any]) -> ToolResult: command = args.get("command", "") timeout = args.get("timeout", 30) if not command: return ToolResult(success=False, data=None, error="No command provided") # Limit command length if len(command) > 2000: return ToolResult(success=False, data=None, error="Command too long (max 2000 chars)") logger.info(f"Executing bash command: {command[:100]}") # Enhanced safety patterns blocked_patterns = [ # File system destruction "rm -rf /", "rm -rf /*", "mkfs", "dd if=", "> /dev/", # Permission escalation "sudo ", "su -", "chmod 777", "chown root", # Remote execution "curl | bash", "curl | sh", "wget | sh", "wget | bash", "curl http", "wget http", # Fork bomb / DoS ":(){", "fork", "kill -9 1", # Network listeners "nc -l", "netcat -l", "socat", # Process killing "killall", "pkill", # Disk operations "fdisk", "parted", ] command_lower = command.lower() for pattern in blocked_patterns: if pattern in command_lower: logger.warning(f"Blocked dangerous command pattern: {pattern}") return ToolResult( success=False, data=None, error=f"Blocked dangerous command pattern: {pattern}" ) try: result = subprocess.run( command, shell=True, capture_output=True, text=True, timeout=timeout ) logger.info(f"Command completed with return code: {result.returncode}") return ToolResult( success=result.returncode == 0, data={ "stdout": result.stdout[:5000], "stderr": result.stderr[:2000], "returncode": result.returncode } ) except subprocess.TimeoutExpired: logger.error(f"Command timed out after {timeout}s") return ToolResult( success=False, data=None, error=f"Command timed out after {timeout}s" ) self.registry.register( name="execute_bash", description="Execute a bash command on the system. Use for file operations, system checks, etc.", parameters={ "type": "object", "properties": { "command": { "type": "string", "description": "Bash command to execute" }, "timeout": { "type": "integer", "description": "Timeout in seconds", "default": 30 } }, "required": ["command"] }, executor=execute_bash ) # API tool def call_api(args: Dict[str, Any]) -> ToolResult: import requests as req url = args.get("url", "") method = args.get("method", "GET").upper() headers = args.get("headers", {}) data = args.get("data") if not url: return ToolResult(success=False, data=None, error="No URL provided") logger.info(f"Calling API: {method} {url[:100]}") try: if method == "GET": resp = req.get(url, headers=headers, timeout=30) elif method == "POST": resp = req.post(url, json=data, headers=headers, timeout=30) elif method == "PUT": resp = req.put(url, json=data, headers=headers, timeout=30) elif method == "DELETE": resp = req.delete(url, headers=headers, timeout=30) else: return ToolResult( success=False, data=None, error=f"Unsupported method: {method}" ) logger.info(f"API response: {resp.status_code}") return ToolResult( success=resp.status_code < 400, data={ "status_code": resp.status_code, "headers": dict(resp.headers), "body": resp.text[:10000] } ) except Exception as e: logger.error(f"API call failed: {e}") return ToolResult(success=False, data=None, error=str(e)) self.registry.register( name="call_api", description="Call an external HTTP API endpoint. Use for REST API interactions.", parameters={ "type": "object", "properties": { "url": { "type": "string", "description": "API endpoint URL" }, "method": { "type": "string", "enum": ["GET", "POST", "PUT", "DELETE"], "description": "HTTP method", "default": "GET" }, "headers": { "type": "object", "description": "HTTP headers" }, "data": { "type": "object", "description": "Request body for POST/PUT" } }, "required": ["url"] }, executor=call_api ) def main(): """Example usage""" caller = OllamaToolCaller() caller.register_default_tools() query = "How many videos are in the database?" print(f"Query: {query}") print("-" * 50) result = caller.run(query) print(f"Result: {result}") if __name__ == "__main__": main()