Files
rp/rp/core/tool_selector.py
T

317 lines
12 KiB
Python
Raw Normal View History

import logging
import re
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, List, Optional, Set
logger = logging.getLogger("rp")
class ToolCategory(Enum):
FILESYSTEM = "filesystem"
SHELL = "shell"
DATABASE = "database"
WEB = "web"
PYTHON = "python"
EDITOR = "editor"
MEMORY = "memory"
AGENT = "agent"
REASONING = "reasoning"
@dataclass
class ToolSelection:
tool: str
category: ToolCategory
reason: str
priority: int = 0
arguments_hint: Dict[str, Any] = field(default_factory=dict)
parallelizable: bool = True
@dataclass
class SelectionDecision:
decisions: List[ToolSelection]
execution_pattern: str
reasoning: str
TOOL_METADATA = {
'run_command': {
'category': ToolCategory.SHELL,
'indicators': ['run', 'execute', 'command', 'shell', 'bash', 'terminal'],
'capabilities': ['system_commands', 'process_management', 'file_operations'],
'parallelizable': True
},
'read_file': {
'category': ToolCategory.FILESYSTEM,
'indicators': ['read', 'view', 'show', 'display', 'content', 'cat'],
'capabilities': ['file_reading', 'inspection'],
'parallelizable': True
},
'write_file': {
'category': ToolCategory.FILESYSTEM,
'indicators': ['write', 'create', 'save', 'generate', 'output'],
'capabilities': ['file_creation', 'file_modification'],
'parallelizable': False
},
'list_directory': {
'category': ToolCategory.FILESYSTEM,
'indicators': ['list', 'ls', 'directory', 'folder', 'files'],
'capabilities': ['directory_listing', 'exploration'],
'parallelizable': True
},
'search_replace': {
'category': ToolCategory.EDITOR,
'indicators': ['replace', 'substitute', 'change', 'update', 'modify'],
'capabilities': ['text_modification', 'refactoring'],
'parallelizable': False
},
'glob_files': {
'category': ToolCategory.FILESYSTEM,
'indicators': ['find', 'search', 'glob', 'pattern', 'match'],
'capabilities': ['file_search', 'pattern_matching'],
'parallelizable': True
},
'grep': {
'category': ToolCategory.FILESYSTEM,
'indicators': ['grep', 'search', 'find', 'pattern', 'content'],
'capabilities': ['content_search', 'pattern_matching'],
'parallelizable': True
},
'http_fetch': {
'category': ToolCategory.WEB,
'indicators': ['fetch', 'http', 'url', 'api', 'request', 'download'],
'capabilities': ['web_requests', 'api_calls'],
'parallelizable': True
},
'web_search': {
'category': ToolCategory.WEB,
'indicators': ['search', 'web', 'internet', 'google', 'lookup'],
'capabilities': ['web_search', 'information_retrieval'],
'parallelizable': True
},
'python_exec': {
'category': ToolCategory.PYTHON,
'indicators': ['python', 'calculate', 'compute', 'script', 'code'],
'capabilities': ['code_execution', 'computation'],
'parallelizable': False
},
'db_query': {
'category': ToolCategory.DATABASE,
'indicators': ['database', 'sql', 'query', 'select', 'table'],
'capabilities': ['database_queries', 'data_retrieval'],
'parallelizable': True
},
'search_knowledge': {
'category': ToolCategory.MEMORY,
'indicators': ['remember', 'recall', 'knowledge', 'memory', 'stored'],
'capabilities': ['memory_retrieval', 'context_recall'],
'parallelizable': True
},
'add_knowledge_entry': {
'category': ToolCategory.MEMORY,
'indicators': ['remember', 'store', 'save', 'note', 'important'],
'capabilities': ['memory_storage', 'knowledge_management'],
'parallelizable': False
}
}
class ToolSelector:
def __init__(self):
self.tool_metadata = TOOL_METADATA
self.selection_history: List[SelectionDecision] = []
def select(self, request: str, context: Dict[str, Any]) -> SelectionDecision:
request_lower = request.lower()
decisions = []
is_filesystem_heavy = self._is_filesystem_heavy(request_lower)
needs_file_read = self._needs_file_read(request_lower, context)
needs_file_write = self._needs_file_write(request_lower)
is_complex = self._is_complex_decision(request_lower)
needs_web = self._needs_web_access(request_lower)
needs_execution = self._needs_code_execution(request_lower)
needs_memory = self._needs_memory_access(request_lower)
reasoning_parts = []
if is_filesystem_heavy:
decisions.append(ToolSelection(
tool='run_command',
category=ToolCategory.SHELL,
reason='Filesystem operations are more efficient via shell commands',
priority=1
))
reasoning_parts.append("Task involves filesystem operations - shell commands preferred")
if needs_file_read:
decisions.append(ToolSelection(
tool='read_file',
category=ToolCategory.FILESYSTEM,
reason='Content inspection required',
priority=2
))
reasoning_parts.append("Need to read file contents")
if needs_file_write:
decisions.append(ToolSelection(
tool='write_file',
category=ToolCategory.FILESYSTEM,
reason='File creation or modification needed',
priority=3,
parallelizable=False
))
reasoning_parts.append("Need to write or modify files")
if is_complex:
decisions.append(ToolSelection(
tool='think',
category=ToolCategory.REASONING,
reason='Complex decision requires analysis',
priority=0,
parallelizable=False
))
reasoning_parts.append("Complex decision - using think tool for analysis")
if needs_web:
decisions.append(ToolSelection(
tool='http_fetch',
category=ToolCategory.WEB,
reason='Web access required',
priority=2
))
reasoning_parts.append("Need to access web resources")
if needs_execution:
decisions.append(ToolSelection(
tool='python_exec',
category=ToolCategory.PYTHON,
reason='Code execution or computation needed',
priority=2,
parallelizable=False
))
reasoning_parts.append("Need to execute code")
if needs_memory:
decisions.append(ToolSelection(
tool='search_knowledge',
category=ToolCategory.MEMORY,
reason='Memory/knowledge access needed',
priority=1
))
reasoning_parts.append("Need to access stored knowledge")
execution_pattern = self._determine_execution_pattern(decisions)
decision = SelectionDecision(
decisions=decisions,
execution_pattern=execution_pattern,
reasoning=" | ".join(reasoning_parts) if reasoning_parts else "No specific tools identified"
)
self.selection_history.append(decision)
return decision
def _is_filesystem_heavy(self, request: str) -> bool:
indicators = [
'file', 'files', 'directory', 'directories', 'folder', 'folders',
'find', 'search', 'list', 'delete', 'remove', 'move', 'copy',
'rename', 'organize', 'sort', 'count', 'size', 'disk'
]
matches = sum(1 for ind in indicators if ind in request)
return matches >= 2
def _needs_file_read(self, request: str, context: Dict[str, Any]) -> bool:
read_indicators = [
'read', 'view', 'show', 'display', 'content', 'what', 'check',
'inspect', 'review', 'analyze', 'look at', 'open'
]
return any(ind in request for ind in read_indicators)
def _needs_file_write(self, request: str) -> bool:
write_indicators = [
'write', 'create', 'save', 'generate', 'make', 'add',
'update', 'modify', 'change', 'edit', 'fix'
]
return any(ind in request for ind in write_indicators)
def _is_complex_decision(self, request: str) -> bool:
complexity_indicators = [
'best', 'optimal', 'compare', 'choose', 'decide', 'trade-off',
'vs', 'versus', 'which', 'should i', 'recommend', 'suggest',
'multiple', 'several', 'options', 'alternatives'
]
matches = sum(1 for ind in complexity_indicators if ind in request)
return matches >= 2
def _needs_web_access(self, request: str) -> bool:
web_indicators = [
'http', 'https', 'url', 'api', 'fetch', 'download',
'web', 'internet', 'online', 'website'
]
return any(ind in request for ind in web_indicators)
def _needs_code_execution(self, request: str) -> bool:
code_indicators = [
'calculate', 'compute', 'run python', 'execute', 'script',
'eval', 'result of', 'what is'
]
return any(ind in request for ind in code_indicators)
def _needs_memory_access(self, request: str) -> bool:
memory_indicators = [
'remember', 'recall', 'stored', 'knowledge', 'previous',
'earlier', 'before', 'told you', 'mentioned'
]
return any(ind in request for ind in memory_indicators)
def _determine_execution_pattern(self, decisions: List[ToolSelection]) -> str:
if not decisions:
return 'none'
parallelizable = [d for d in decisions if d.parallelizable]
sequential = [d for d in decisions if not d.parallelizable]
if len(sequential) > 0 and len(parallelizable) > 0:
return 'mixed'
elif len(sequential) > 0:
return 'sequential'
elif len(parallelizable) > 1:
return 'parallel'
return 'sequential'
def get_tool_for_task(self, task_type: str) -> Optional[str]:
task_tool_map = {
'find_files': 'glob_files',
'search_content': 'grep',
'read_file': 'read_file',
'write_file': 'write_file',
'execute_command': 'run_command',
'web_request': 'http_fetch',
'web_search': 'web_search',
'compute': 'python_exec',
'database': 'db_query',
'remember': 'add_knowledge_entry',
'recall': 'search_knowledge'
}
return task_tool_map.get(task_type)
def suggest_parallelization(self, tool_calls: List[Dict[str, Any]]) -> Dict[str, List[Dict[str, Any]]]:
parallelizable = []
sequential = []
for call in tool_calls:
tool_name = call.get('function', {}).get('name', '')
metadata = self.tool_metadata.get(tool_name, {})
if metadata.get('parallelizable', True):
parallelizable.append(call)
else:
sequential.append(call)
return {
'parallel': parallelizable,
'sequential': sequential
}
def get_statistics(self) -> Dict[str, Any]:
if not self.selection_history:
return {'total_selections': 0}
tool_usage = {}
pattern_usage = {}
for decision in self.selection_history:
for sel in decision.decisions:
tool_usage[sel.tool] = tool_usage.get(sel.tool, 0) + 1
pattern_usage[decision.execution_pattern] = pattern_usage.get(decision.execution_pattern, 0) + 1
return {
'total_selections': len(self.selection_history),
'tool_usage': tool_usage,
'pattern_usage': pattern_usage,
'most_used_tool': max(tool_usage.items(), key=lambda x: x[1])[0] if tool_usage else None
}