39a2cbc65b
- get_face_groups_handler: COALESCE(tp.name, tn.label) for name consistency - sync_file_status: compare JSON vs pre_chunks (not chunk table) - face consistency: compare frames.len() not total_faces - cleanup 2 ghost records with NULL file_name/file_path - replace identity_agent with face_dedup in pipeline stages - remove identity_agent_api.rs and all references - update required_processors to match actual processors - update AGENTS.md with team responsibilities - add Studio pipeline changes documentation
791 lines
30 KiB
Python
791 lines
30 KiB
Python
#!/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()
|