fix: face group name read consistency, sync_file_status fix, cleanup ghost records, identity_agent replaced with face_dedup
- 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
This commit is contained in:
@@ -0,0 +1,790 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user