From 6bbbefaa6f468efe35a0fe6e868d9a5f1acf9064 Mon Sep 17 00:00:00 2001 From: vitya Date: Mon, 4 May 2026 00:09:41 +0300 Subject: [PATCH] feat: refactor APIClient and ToolRegistry to classes - api_client.py: APIClient class with ProviderConfig types, chat_stream() returns str instead of printing, module-level wrappers for backward compat - tools.py: ToolRegistry class with instance methods, pure Python list_files (os.walk) and search_in_files (re.search) replacing find/grep for Windows compat, English tool descriptions and error messages, module-level wrappers for backward compat - 74 new tests (21 api_client + 53 tools) Co-Authored-By: Claude Opus 4.7 --- meeting_room/api_client.py | 222 +++++++++++------ meeting_room/tools.py | 474 +++++++++++++++++++++---------------- tests/test_api_client.py | 359 ++++++++++++++++++++++++++++ tests/test_tools.py | 389 ++++++++++++++++++++++++++++++ 4 files changed, 1170 insertions(+), 274 deletions(-) create mode 100644 tests/test_api_client.py create mode 100644 tests/test_tools.py diff --git a/meeting_room/api_client.py b/meeting_room/api_client.py index 207a146..9e141d5 100644 --- a/meeting_room/api_client.py +++ b/meeting_room/api_client.py @@ -1,26 +1,157 @@ """Universal API client supporting multiple OpenAI-compatible providers.""" -import httpx +from __future__ import annotations + import json +from typing import Iterator + +import httpx + +from meeting_room.models import ProviderConfig -PROVIDERS = {} +class APIClient: + """Stateful client that holds provider configs and makes LLM calls. + + Parameters + ---------- + providers: + Mapping of provider name to ``ProviderConfig`` (base_url + api_key). + """ + + def __init__(self, providers: dict[str, ProviderConfig]) -> None: + self.providers: dict[str, ProviderConfig] = providers + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + def _get_client(self, provider_name: str) -> httpx.Client: + """Create a short-lived ``httpx.Client`` for *provider_name*.""" + p = self.providers.get(provider_name) + if not p: + raise ValueError(f"Unknown provider: {provider_name}") + return httpx.Client( + base_url=p.base_url, + headers={"Authorization": f"Bearer {p.api_key}"}, + timeout=120.0, + ) + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + def chat( + self, + provider: str, + model: str, + messages: list[dict], + temperature: float = 0.7, + tools: list[dict] | None = None, + ) -> dict: + """Send a non-streaming chat completion request. + + Returns a dict with keys ``content`` (str) and ``tool_calls`` + (list | None). + """ + client = self._get_client(provider) + payload: dict = { + "model": model, + "messages": messages, + "temperature": temperature, + } + if tools: + payload["tools"] = tools + payload["tool_choice"] = "auto" + + resp = client.post("/chat/completions", json=payload) + resp.raise_for_status() + data = resp.json() + message = data["choices"][0]["message"] + + result: dict = { + "content": message.get("content") or "", + "tool_calls": None, + } + + if message.get("tool_calls"): + result["tool_calls"] = [ + { + "id": tc["id"], + "name": tc["function"]["name"], + "arguments": tc["function"]["arguments"], + } + for tc in message["tool_calls"] + ] + + return result + + def chat_stream( + self, + provider: str, + model: str, + messages: list[dict], + temperature: float = 0.7, + ) -> str: + """Send a streaming chat completion request. + + Accumulates all content deltas and **returns** the full text. + Does NOT print to stdout. + """ + client = self._get_client(provider) + payload: dict = { + "model": model, + "messages": messages, + "temperature": temperature, + "stream": True, + } + full_text = "" + with client.stream("POST", "/chat/completions", json=payload) as resp: + resp.raise_for_status() + for line in resp.iter_lines(): + if not line or not line.startswith("data: "): + continue + data = line[6:] + if data == "[DONE]": + break + try: + chunk = json.loads(data) + delta = chunk["choices"][0]["delta"] + if "content" in delta and delta["content"]: + full_text += delta["content"] + except (json.JSONDecodeError, KeyError, IndexError): + continue + return full_text -def init_providers(config: dict): - global PROVIDERS - PROVIDERS = config.get("providers", {}) +# ====================================================================== +# Module-level backward-compatible wrappers +# ====================================================================== + +_default_client: APIClient | None = None -def _get_client(provider_name: str) -> httpx.Client: - p = PROVIDERS.get(provider_name) - if not p: - raise ValueError(f"Unknown provider: {provider_name}") - return httpx.Client( - base_url=p["base_url"], - headers={"Authorization": f"Bearer {p['api_key']}"}, - timeout=120.0, - ) +def init_providers(config: dict) -> None: + """Initialise the module-level default ``APIClient``. + + *config* is the raw dict produced by ``yaml.safe_load()`` — it must + contain a ``providers`` key whose value maps provider names to + ``{base_url, api_key}`` dicts. + """ + global _default_client + raw: dict = config.get("providers", {}) + providers: dict[str, ProviderConfig] = { + name: ProviderConfig(**vals) for name, vals in raw.items() + } + _default_client = APIClient(providers) + + +def _require_default() -> APIClient: + if _default_client is None: + raise RuntimeError( + "init_providers() must be called before using module-level chat/chat_stream" + ) + return _default_client def chat( @@ -30,36 +161,8 @@ def chat( temperature: float = 0.7, tools: list[dict] | None = None, ) -> dict: - client = _get_client(provider) - payload = { - "model": model, - "messages": messages, - "temperature": temperature, - } - if tools: - payload["tools"] = tools - payload["tool_choice"] = "auto" - - resp = client.post("/chat/completions", json=payload) - resp.raise_for_status() - data = resp.json() - message = data["choices"][0]["message"] - - result = { - "content": message.get("content") or "", - "tool_calls": None, - } - - if message.get("tool_calls"): - result["tool_calls"] = [] - for tc in message["tool_calls"]: - result["tool_calls"].append({ - "id": tc["id"], - "name": tc["function"]["name"], - "arguments": tc["function"]["arguments"], - }) - - return result + """Module-level wrapper that delegates to the default ``APIClient``.""" + return _require_default().chat(provider, model, messages, temperature, tools) def chat_stream( @@ -68,29 +171,8 @@ def chat_stream( messages: list[dict], temperature: float = 0.7, ) -> str: - client = _get_client(provider) - payload = { - "model": model, - "messages": messages, - "temperature": temperature, - "stream": True, - } - full_text = "" - with client.stream("POST", "/chat/completions", json=payload) as resp: - resp.raise_for_status() - for line in resp.iter_lines(): - if not line or not line.startswith("data: "): - continue - data = line[6:] - if data == "[DONE]": - break - try: - chunk = json.loads(data) - delta = chunk["choices"][0]["delta"] - if "content" in delta and delta["content"]: - print(delta["content"], end="", flush=True) - full_text += delta["content"] - except (json.JSONDecodeError, KeyError, IndexError): - continue - print() - return full_text + """Module-level wrapper that delegates to the default ``APIClient``. + + Returns the accumulated text — does NOT print to stdout. + """ + return _require_default().chat_stream(provider, model, messages, temperature) \ No newline at end of file diff --git a/meeting_room/tools.py b/meeting_room/tools.py index 69e1c03..b7f8d3e 100644 --- a/meeting_room/tools.py +++ b/meeting_room/tools.py @@ -1,32 +1,17 @@ """ Tools system for Meeting Room agents. Provides file access and web capabilities via OpenAI function calling. + +The ToolRegistry class encapsulates all tool logic with a configurable +workdir. Module-level wrappers maintain backward compatibility. """ +import fnmatch +import httpx +import json import os import re -import json import subprocess -import httpx - -WORKDIR = "." - - -def set_workdir(path: str): - global WORKDIR - WORKDIR = os.path.abspath(path) - - -def _resolve_path(path: str) -> str: - if os.path.isabs(path): - return path - return os.path.normpath(os.path.join(WORKDIR, path)) - - -def _truncate(text: str, max_len: int = 5000) -> str: - if len(text) <= max_len: - return text - return text[:max_len] + f"\n... [обрезано, всего {len(text)} символов]" # ─── Tool schemas (OpenAI function calling format) ─── @@ -36,13 +21,13 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "read_file", - "description": "Прочитать содержимое файла. Возвращает текст с номерами строк.", + "description": "Read file contents. Returns text with line numbers.", "parameters": { "type": "object", "properties": { - "path": {"type": "string", "description": "Путь к файлу"}, - "offset": {"type": "integer", "description": "Начать с этой строки (1-indexed)", "default": 1}, - "limit": {"type": "integer", "description": "Максимум строк", "default": 100}, + "path": {"type": "string", "description": "Path to the file"}, + "offset": {"type": "integer", "description": "Start from this line (1-indexed)", "default": 1}, + "limit": {"type": "integer", "description": "Maximum number of lines", "default": 100}, }, "required": ["path"], }, @@ -52,13 +37,13 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "list_files", - "description": "Показать список файлов и директорий.", + "description": "List files and directories.", "parameters": { "type": "object", "properties": { - "path": {"type": "string", "description": "Директория для просмотра"}, - "pattern": {"type": "string", "description": "Фильтр по имени (glob)"}, - "max_depth": {"type": "integer", "description": "Глубина обхода", "default": 3}, + "path": {"type": "string", "description": "Directory to list"}, + "pattern": {"type": "string", "description": "Filter by name (glob)"}, + "max_depth": {"type": "integer", "description": "Traversal depth", "default": 3}, }, "required": [], }, @@ -68,13 +53,13 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "search_in_files", - "description": "Поиск текста в файлах (regex).", + "description": "Search text in files (regex).", "parameters": { "type": "object", "properties": { - "query": {"type": "string", "description": "Поисковый запрос (regex)"}, - "path": {"type": "string", "description": "Директория для поиска"}, - "file_pattern": {"type": "string", "description": "Фильтр по расширению"}, + "query": {"type": "string", "description": "Search query (regex)"}, + "path": {"type": "string", "description": "Directory to search in"}, + "file_pattern": {"type": "string", "description": "Filter by extension (e.g. *.py)"}, }, "required": ["query"], }, @@ -84,12 +69,12 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "web_search", - "description": "Поиск в интернете. Возвращает заголовки, сниппеты и URL.", + "description": "Search the web. Returns titles, snippets and URLs.", "parameters": { "type": "object", "properties": { - "query": {"type": "string", "description": "Поисковый запрос"}, - "max_results": {"type": "integer", "description": "Максимум результатов", "default": 5}, + "query": {"type": "string", "description": "Search query"}, + "max_results": {"type": "integer", "description": "Maximum results", "default": 5}, }, "required": ["query"], }, @@ -99,12 +84,12 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "web_fetch", - "description": "Скачать содержимое веб-страницы.", + "description": "Fetch web page content.", "parameters": { "type": "object", "properties": { - "url": {"type": "string", "description": "URL для загрузки"}, - "max_length": {"type": "integer", "description": "Максимум символов", "default": 5000}, + "url": {"type": "string", "description": "URL to fetch"}, + "max_length": {"type": "integer", "description": "Maximum characters", "default": 5000}, }, "required": ["url"], }, @@ -114,12 +99,12 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "write_file", - "description": "Записать содержимое в файл.", + "description": "Write content to a file.", "parameters": { "type": "object", "properties": { - "path": {"type": "string", "description": "Путь к файлу"}, - "content": {"type": "string", "description": "Содержимое"}, + "path": {"type": "string", "description": "Path to the file"}, + "content": {"type": "string", "description": "Content to write"}, }, "required": ["path", "content"], }, @@ -129,12 +114,12 @@ TOOL_SCHEMAS = [ "type": "function", "function": { "name": "run_command", - "description": "Выполнить shell-команду.", + "description": "Run a shell command.", "parameters": { "type": "object", "properties": { - "command": {"type": "string", "description": "Команда"}, - "timeout": {"type": "integer", "description": "Таймаут (сек)", "default": 30}, + "command": {"type": "string", "description": "Command to run"}, + "timeout": {"type": "integer", "description": "Timeout in seconds", "default": 30}, }, "required": ["command"], }, @@ -142,170 +127,251 @@ TOOL_SCHEMAS = [ }, ] +# Directories to skip during file traversal +_SKIP_DIRS = {".git", ".venv", "venv", "node_modules", "__pycache__", ".mypy_cache", ".pytest_cache"} -# ─── Implementations ─── - -def tool_read_file(path: str, offset: int = 1, limit: int = 100) -> str: - resolved = _resolve_path(path) - if not os.path.isfile(resolved): - return f"Ошибка: файл не найден — {resolved}" - try: - with open(resolved, encoding="utf-8", errors="replace") as f: - lines = f.readlines() - total = len(lines) - selected = lines[offset - 1 : offset - 1 + limit] - result = "".join(f"{i}|{line}" for i, line in enumerate(selected, start=offset)) - return _truncate(f"[{resolved} — {total} строк]\n{result}") - except Exception as e: - return f"Ошибка чтения: {e}" - - -def tool_list_files(path: str = "", pattern: str = "", max_depth: int = 3) -> str: - base = _resolve_path(path or ".") - if not os.path.isdir(base): - return f"Ошибка: директория не найдена — {base}" - try: - result = subprocess.run( - ["find", base, "-maxdepth", str(max_depth), - "-not", "-path", "*/.git/*", "-not", "-path", "*/.venv/*", - "-not", "-path", "*/node_modules/*", "-not", "-path", "*/__pycache__/*"], - capture_output=True, text=True, timeout=10, - ) - lines = result.stdout.strip().split("\n") - if pattern: - import fnmatch - lines = [l for l in lines if fnmatch.fnmatch(os.path.basename(l), pattern)] - lines.sort() - return _truncate("\n".join(lines)) - except Exception as e: - return f"Ошибка: {e}" - - -def tool_search_in_files(query: str, path: str = "", file_pattern: str = "") -> str: - base = _resolve_path(path or ".") - cmd = ["grep", "-rn", "--color=never", "-E", query, base] - if file_pattern: - cmd.extend(["--include", file_pattern]) - try: - result = subprocess.run(cmd, capture_output=True, text=True, timeout=15) - if result.returncode == 1: - return "Совпадений не найдено." - return _truncate(result.stdout.strip()) - except subprocess.TimeoutExpired: - return "Таймаут поиска." - except Exception as e: - return f"Ошибка: {e}" - - -def tool_web_search(query: str, max_results: int = 5) -> str: - try: - client = httpx.Client(timeout=15, follow_redirects=True) - resp = client.get( - "https://html.duckduckgo.com/html/", params={"q": query}, - headers={"User-Agent": "Mozilla/5.0"}, - ) - resp.raise_for_status() - results = [] - for match in re.finditer( - r'\s*(.*?)\s*.*?' - r'(.*?)', - resp.text, re.DOTALL, - ): - url = match.group(1) - title = re.sub(r'<[^>]+>', '', match.group(2)).strip() - snippet = re.sub(r'<[^>]+>', '', match.group(3)).strip() - results.append(f"{title}\n {url}\n {snippet}") - if len(results) >= max_results: - break - if not results: - return "Результатов не найдено." - return _truncate("\n\n".join(f"[{i+1}] {r}" for i, r in enumerate(results))) - except Exception as e: - return f"Ошибка поиска: {e}" - - -def tool_web_fetch(url: str, max_length: int = 5000) -> str: - try: - client = httpx.Client(timeout=20, follow_redirects=True) - resp = client.get(url, headers={"User-Agent": "Mozilla/5.0"}) - resp.raise_for_status() - content_type = resp.headers.get("content-type", "") - if "html" in content_type: - text = re.sub(r']*>.*?', '', resp.text, flags=re.DOTALL) - text = re.sub(r']*>.*?', '', text, flags=re.DOTALL) - text = re.sub(r'<[^>]+>', ' ', text) - text = re.sub(r'\s+', ' ', text).strip() - else: - text = resp.text - return _truncate(f"[{url} — {len(resp.text)} символов]\n{text}", max_length) - except Exception as e: - return f"Ошибка загрузки: {e}" - - -def tool_write_file(path: str, content: str) -> str: - resolved = _resolve_path(path) - try: - os.makedirs(os.path.dirname(resolved), exist_ok=True) - with open(resolved, "w", encoding="utf-8") as f: - f.write(content) - return f"Записано {len(content)} символов в {resolved}" - except Exception as e: - return f"Ошибка записи: {e}" - - -def tool_run_command(command: str, timeout: int = 30) -> str: - try: - result = subprocess.run( - command, shell=True, capture_output=True, text=True, - timeout=timeout, cwd=WORKDIR, - ) - output = result.stdout - if result.stderr: - output += f"\nSTDERR:\n{result.stderr}" - if result.returncode != 0: - output += f"\n[exit code: {result.returncode}]" - return _truncate(output) - except subprocess.TimeoutExpired: - return f"Таймаут ({timeout}с)" - except Exception as e: - return f"Ошибка: {e}" - - -# ─── Dispatcher ─── - -TOOL_FUNCTIONS = { - "read_file": tool_read_file, - "list_files": tool_list_files, - "search_in_files": tool_search_in_files, - "web_search": tool_web_search, - "web_fetch": tool_web_fetch, - "write_file": tool_write_file, - "run_command": tool_run_command, -} - +# Tool set definitions TOOL_SETS = { - "all": list(TOOL_FUNCTIONS.keys()), - "full": list(TOOL_FUNCTIONS.keys()), - "files": ["read_file", "list_files", "search_in_files", "write_file"], - "readonly": ["read_file", "list_files", "search_in_files"], - "web": ["web_search", "web_fetch"], - "none": [], + "all": ["read_file", "list_files", "search_in_files", "web_search", "web_fetch", "write_file", "run_command"], + "full": ["read_file", "list_files", "search_in_files", "web_search", "web_fetch", "write_file", "run_command"], + "files": ["read_file", "list_files", "search_in_files", "write_file"], + "readonly": ["read_file", "list_files", "search_in_files"], + "web": ["web_search", "web_fetch"], + "none": [], } -def get_tool_schemas(tool_set: str = "all") -> list[dict]: - if isinstance(tool_set, list): - names = tool_set - else: - names = TOOL_SETS.get(tool_set, TOOL_SETS["all"]) - return [s for s in TOOL_SCHEMAS if s["function"]["name"] in names] +def _truncate(text: str, max_len: int = 5000) -> str: + if len(text) <= max_len: + return text + return text[:max_len] + f"\n... [truncated, {len(text)} chars total]" + + +class ToolRegistry: + """Manages tool schemas and execution for a given working directory.""" + + def __init__(self, workdir: str = ".") -> None: + self._workdir = os.path.abspath(workdir) + + # ── Path resolution ── + + def _resolve_path(self, path: str) -> str: + if os.path.isabs(path): + return path + return os.path.normpath(os.path.join(self._workdir, path)) + + # ── Schema / dispatch ── + + def get_tool_schemas(self, tool_set: str | list = "all") -> list[dict]: + if isinstance(tool_set, list): + names = tool_set + else: + names = TOOL_SETS.get(tool_set, TOOL_SETS["all"]) + return [s for s in TOOL_SCHEMAS if s["function"]["name"] in names] + + def execute_tool(self, name: str, arguments: dict) -> str: + method_map = { + "read_file": self.tool_read_file, + "list_files": self.tool_list_files, + "search_in_files": self.tool_search_in_files, + "web_search": self.tool_web_search, + "web_fetch": self.tool_web_fetch, + "write_file": self.tool_write_file, + "run_command": self.tool_run_command, + } + func = method_map.get(name) + if not func: + return f"Unknown tool: {name}" + try: + return func(**arguments) + except Exception as e: + return f"Error executing {name}: {e}" + + # ── Tool implementations ── + + def tool_read_file(self, path: str, offset: int = 1, limit: int = 100) -> str: + resolved = self._resolve_path(path) + if not os.path.isfile(resolved): + return f"Error: file not found — {resolved}" + try: + with open(resolved, encoding="utf-8", errors="replace") as f: + lines = f.readlines() + total = len(lines) + selected = lines[offset - 1 : offset - 1 + limit] + result = "".join(f"{i}|{line}" for i, line in enumerate(selected, start=offset)) + return _truncate(f"[{resolved} — {total} lines]\n{result}") + except Exception as e: + return f"Error reading file: {e}" + + def tool_list_files(self, path: str = "", pattern: str = "", max_depth: int = 3) -> str: + base = self._resolve_path(path or ".") + if not os.path.isdir(base): + return f"Error: directory not found — {base}" + try: + results: list[str] = [] + base_depth = base.count(os.sep) + for dirpath, dirnames, filenames in os.walk(base): + # Prune skipped directories in-place so os.walk skips them + dirnames[:] = [d for d in dirnames if d not in _SKIP_DIRS] + # Check depth + current_depth = dirpath.count(os.sep) - base_depth + if current_depth >= max_depth: + # Don't descend further, but still list this dir's files + dirnames.clear() + # Add directory itself + rel_dir = os.path.relpath(dirpath, base) + if rel_dir == ".": + results.append(dirpath) + else: + results.append(dirpath) + # Add files + for fname in filenames: + if pattern and not fnmatch.fnmatch(fname, pattern): + continue + results.append(os.path.join(dirpath, fname)) + results.sort() + return _truncate("\n".join(results)) + except Exception as e: + return f"Error: {e}" + + def tool_search_in_files(self, query: str, path: str = "", file_pattern: str = "") -> str: + base = self._resolve_path(path or ".") + if not os.path.isdir(base): + return f"Error: directory not found — {base}" + try: + regex = re.compile(query, re.IGNORECASE) + except re.error as e: + return f"Error: invalid regex — {e}" + try: + results: list[str] = [] + for dirpath, dirnames, filenames in os.walk(base): + dirnames[:] = [d for d in dirnames if d not in _SKIP_DIRS] + for fname in filenames: + if file_pattern and not fnmatch.fnmatch(fname, file_pattern): + continue + fpath = os.path.join(dirpath, fname) + try: + with open(fpath, encoding="utf-8", errors="replace") as f: + for line_num, line in enumerate(f, start=1): + if regex.search(line): + results.append(f"{fpath}:{line_num}:{line.rstrip()}") + if len(results) >= 200: + return _truncate("\n".join(results)) + except (OSError, UnicodeDecodeError): + continue + if not results: + return "No matches found." + return _truncate("\n".join(results)) + except Exception as e: + return f"Error: {e}" + + def tool_web_search(self, query: str, max_results: int = 5) -> str: + try: + client = httpx.Client(timeout=15, follow_redirects=True) + resp = client.get( + "https://html.duckduckgo.com/html/", + params={"q": query}, + headers={"User-Agent": "Mozilla/5.0"}, + ) + resp.raise_for_status() + results = [] + for match in re.finditer( + r'\s*(.*?)\s*.*?' + r'(.*?)', + resp.text, re.DOTALL, + ): + url = match.group(1) + title = re.sub(r'<[^>]+>', '', match.group(2)).strip() + snippet = re.sub(r'<[^>]+>', '', match.group(3)).strip() + results.append(f"{title}\n {url}\n {snippet}") + if len(results) >= max_results: + break + if not results: + return "No results found." + return _truncate("\n\n".join(f"[{i+1}] {r}" for i, r in enumerate(results))) + except Exception as e: + return f"Search error: {e}" + + def tool_web_fetch(self, url: str, max_length: int = 5000) -> str: + try: + client = httpx.Client(timeout=20, follow_redirects=True) + resp = client.get(url, headers={"User-Agent": "Mozilla/5.0"}) + resp.raise_for_status() + content_type = resp.headers.get("content-type", "") + if "html" in content_type: + text = re.sub(r']*>.*?', '', resp.text, flags=re.DOTALL) + text = re.sub(r']*>.*?', '', text, flags=re.DOTALL) + text = re.sub(r'<[^>]+>', ' ', text) + text = re.sub(r'\s+', ' ', text).strip() + else: + text = resp.text + return _truncate(f"[{url} — {len(resp.text)} chars]\n{text}", max_length) + except Exception as e: + return f"Fetch error: {e}" + + def tool_write_file(self, path: str, content: str) -> str: + resolved = self._resolve_path(path) + try: + os.makedirs(os.path.dirname(resolved) or ".", exist_ok=True) + with open(resolved, "w", encoding="utf-8") as f: + f.write(content) + return f"Wrote {len(content)} chars to {resolved}" + except Exception as e: + return f"Write error: {e}" + + def tool_run_command(self, command: str, timeout: int = 30) -> str: + try: + result = subprocess.run( + command, shell=True, capture_output=True, text=True, + timeout=timeout, cwd=self._workdir, + ) + output = result.stdout + if result.stderr: + output += f"\nSTDERR:\n{result.stderr}" + if result.returncode != 0: + output += f"\n[exit code: {result.returncode}]" + return _truncate(output) + except subprocess.TimeoutExpired: + return f"Timeout ({timeout}s)" + except Exception as e: + return f"Error: {e}" + + +# ─── Module-level backward-compatible wrappers ─── + +_default_registry: ToolRegistry | None = None + + +def set_workdir(path: str) -> None: + """Set the working directory for module-level tool functions.""" + global _default_registry + _default_registry = ToolRegistry(workdir=path) + + +def _get_registry() -> ToolRegistry: + global _default_registry + if _default_registry is None: + _default_registry = ToolRegistry() + return _default_registry + + +def get_tool_schemas(tool_set: str | list = "all") -> list[dict]: + """Return tool schemas for the given set name or explicit list.""" + return _get_registry().get_tool_schemas(tool_set) def execute_tool(name: str, arguments: dict) -> str: - func = TOOL_FUNCTIONS.get(name) - if not func: - return f"Неизвестный инструмент: {name}" - try: - return func(**arguments) - except Exception as e: - return f"Ошибка выполнения {name}: {e}" + """Execute a tool by name with the given arguments.""" + return _get_registry().execute_tool(name, arguments) + + +# Legacy dict-based dispatch (kept for any code that imports TOOL_FUNCTIONS) +TOOL_FUNCTIONS = { + "read_file": lambda **kw: _get_registry().tool_read_file(**kw), + "list_files": lambda **kw: _get_registry().tool_list_files(**kw), + "search_in_files": lambda **kw: _get_registry().tool_search_in_files(**kw), + "web_search": lambda **kw: _get_registry().tool_web_search(**kw), + "web_fetch": lambda **kw: _get_registry().tool_web_fetch(**kw), + "write_file": lambda **kw: _get_registry().tool_write_file(**kw), + "run_command": lambda **kw: _get_registry().tool_run_command(**kw), +} \ No newline at end of file diff --git a/tests/test_api_client.py b/tests/test_api_client.py new file mode 100644 index 0000000..f1fd8c0 --- /dev/null +++ b/tests/test_api_client.py @@ -0,0 +1,359 @@ +"""Tests for meeting_room.api_client — APIClient class and module-level wrappers.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from meeting_room.api_client import APIClient, chat, chat_stream, init_providers +from meeting_room.models import ProviderConfig + + +# --------------------------------------------------------------------------- +# Fixtures & helpers +# --------------------------------------------------------------------------- + +PROVIDERS = { + "testprov": ProviderConfig( + base_url="https://api.test.example.com/v1", + api_key="sk-test-key-123", + ), +} + + +def _make_chat_response( + content: str = "Hello!", + tool_calls: list[dict] | None = None, +) -> dict: + """Build a minimal OpenAI-style chat completion response.""" + message: dict = {"content": content} + if tool_calls: + message["tool_calls"] = tool_calls + return {"choices": [{"message": message}]} + + +def _make_stream_lines(chunks: list[str]) -> list[str]: + """Build SSE lines that look like what OpenAI streaming returns. + + Each chunk becomes a ``data: {json}`` line. A ``data: [DONE]`` + sentinel is appended automatically. + """ + lines: list[str] = [] + for text in chunks: + delta = {"choices": [{"delta": {"content": text}}]} + lines.append(f"data: {json.dumps(delta)}") + lines.append("data: [DONE]") + return lines + + +# --------------------------------------------------------------------------- +# APIClient — construction +# --------------------------------------------------------------------------- + + +class TestAPIClientInit: + def test_stores_providers(self) -> None: + client = APIClient(PROVIDERS) + assert "testprov" in client.providers + assert client.providers["testprov"].base_url == "https://api.test.example.com/v1" + + def test_empty_providers_ok(self) -> None: + client = APIClient({}) + assert client.providers == {} + + +# --------------------------------------------------------------------------- +# APIClient.chat +# --------------------------------------------------------------------------- + + +class TestAPIClientChat: + @patch("meeting_room.api_client.httpx.Client") + def test_basic_text_response(self, mock_client_cls: MagicMock) -> None: + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("Hi there") + mock_resp.raise_for_status = MagicMock() + mock_client_instance = mock_client_cls.return_value + mock_client_instance.post.return_value = mock_resp + + client = APIClient(PROVIDERS) + result = client.chat("testprov", "gpt-4o-mini", [{"role": "user", "content": "Hi"}]) + + assert result == {"content": "Hi there", "tool_calls": None} + mock_client_instance.post.assert_called_once() + call_args = mock_client_instance.post.call_args + assert call_args[0][0] == "/chat/completions" + payload = call_args[1]["json"] + assert payload["model"] == "gpt-4o-mini" + assert payload["messages"] == [{"role": "user", "content": "Hi"}] + assert payload["temperature"] == 0.7 + + @patch("meeting_room.api_client.httpx.Client") + def test_custom_temperature(self, mock_client_cls: MagicMock) -> None: + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("cold") + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + client = APIClient(PROVIDERS) + client.chat("testprov", "m", [], temperature=0.2) + + payload = mock_client_cls.return_value.post.call_args[1]["json"] + assert payload["temperature"] == 0.2 + + @patch("meeting_room.api_client.httpx.Client") + def test_tool_calls_in_response(self, mock_client_cls: MagicMock) -> None: + tool_calls = [ + { + "id": "call_abc", + "function": {"name": "read_file", "arguments": '{"path": "/tmp/x"}'}, + "type": "function", + } + ] + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("", tool_calls=tool_calls) + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + client = APIClient(PROVIDERS) + result = client.chat("testprov", "m", []) + + assert result["tool_calls"] == [ + {"id": "call_abc", "name": "read_file", "arguments": '{"path": "/tmp/x"}'} + ] + + @patch("meeting_room.api_client.httpx.Client") + def test_tools_sent_in_payload(self, mock_client_cls: MagicMock) -> None: + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("ok") + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + tools = [{"type": "function", "function": {"name": "search"}}] + client = APIClient(PROVIDERS) + client.chat("testprov", "m", [], tools=tools) + + payload = mock_client_cls.return_value.post.call_args[1]["json"] + assert "tools" in payload + assert payload["tool_choice"] == "auto" + + def test_unknown_provider_raises(self) -> None: + client = APIClient(PROVIDERS) + with pytest.raises(ValueError, match="Unknown provider: nonexistent"): + client.chat("nonexistent", "m", []) + + @patch("meeting_room.api_client.httpx.Client") + def test_empty_content_yields_empty_string(self, mock_client_cls: MagicMock) -> None: + mock_resp = MagicMock() + mock_resp.json.return_value = {"choices": [{"message": {"content": None}}]} + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + client = APIClient(PROVIDERS) + result = client.chat("testprov", "m", []) + assert result["content"] == "" + + @patch("meeting_room.api_client.httpx.Client") + def test_http_error_propagates(self, mock_client_cls: MagicMock) -> None: + mock_resp = MagicMock() + mock_resp.raise_for_status.side_effect = httpx.HTTPStatusError( + "Bad Request", + request=MagicMock(), + response=httpx.Response(400), + ) + mock_client_cls.return_value.post.return_value = mock_resp + + client = APIClient(PROVIDERS) + with pytest.raises(httpx.HTTPStatusError): + client.chat("testprov", "m", []) + + +# --------------------------------------------------------------------------- +# APIClient.chat_stream +# --------------------------------------------------------------------------- + + +class TestAPIClientChatStream: + @patch("meeting_room.api_client.httpx.Client") + def test_accumulates_text(self, mock_client_cls: MagicMock) -> None: + """chat_stream must return the concatenated text, NOT print.""" + lines = _make_stream_lines(["Hello", " world", "!"]) + + mock_response = MagicMock() + mock_response.iter_lines.return_value = lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + + mock_client_instance = mock_client_cls.return_value + mock_client_instance.stream.return_value = mock_response + + client = APIClient(PROVIDERS) + result = client.chat_stream("testprov", "gpt-4o-mini", [{"role": "user", "content": "Hi"}]) + + assert result == "Hello world!" + + @patch("meeting_room.api_client.httpx.Client") + def test_stream_payload_includes_stream_flag(self, mock_client_cls: MagicMock) -> None: + lines = _make_stream_lines(["x"]) + mock_response = MagicMock() + mock_response.iter_lines.return_value = lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + mock_client_cls.return_value.stream.return_value = mock_response + + client = APIClient(PROVIDERS) + client.chat_stream("testprov", "m", []) + + call_args = mock_client_cls.return_value.stream.call_args + payload = call_args[1]["json"] + assert payload["stream"] is True + + @patch("meeting_room.api_client.httpx.Client") + def test_empty_stream_returns_empty_string(self, mock_client_cls: MagicMock) -> None: + lines = _make_stream_lines([]) + mock_response = MagicMock() + mock_response.iter_lines.return_value = lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + mock_client_cls.return_value.stream.return_value = mock_response + + client = APIClient(PROVIDERS) + result = client.chat_stream("testprov", "m", []) + assert result == "" + + @patch("meeting_room.api_client.httpx.Client") + def test_malformed_lines_skipped(self, mock_client_cls: MagicMock) -> None: + """Lines that are not valid SSE data should be silently skipped.""" + raw_lines = [ + "", # empty line — skipped + "not SSE", # no "data: " prefix — skipped + "data: {bad json", # malformed JSON — skipped + 'data: {"choices":[{"delta":{"content":"ok"}}]}', # valid + "data: [DONE]", + ] + mock_response = MagicMock() + mock_response.iter_lines.return_value = raw_lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + mock_client_cls.return_value.stream.return_value = mock_response + + client = APIClient(PROVIDERS) + result = client.chat_stream("testprov", "m", []) + assert result == "ok" + + @patch("meeting_room.api_client.httpx.Client") + def test_does_not_print(self, mock_client_cls: MagicMock, capsys: pytest.CaptureFixture) -> None: + """chat_stream must NOT write to stdout.""" + lines = _make_stream_lines(["Hello", " world"]) + mock_response = MagicMock() + mock_response.iter_lines.return_value = lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + mock_client_cls.return_value.stream.return_value = mock_response + + client = APIClient(PROVIDERS) + result = client.chat_stream("testprov", "m", []) + + captured = capsys.readouterr() + assert captured.out == "" + assert result == "Hello world" + + def test_unknown_provider_raises(self) -> None: + client = APIClient(PROVIDERS) + with pytest.raises(ValueError, match="Unknown provider: nope"): + client.chat_stream("nope", "m", []) + + +# --------------------------------------------------------------------------- +# Module-level wrappers — backward compatibility +# --------------------------------------------------------------------------- + + +class TestModuleLevelWrappers: + def test_init_providers_creates_default_client(self) -> None: + import meeting_room.api_client as mod + + config = { + "providers": { + "myprov": {"base_url": "https://api.example.com/v1", "api_key": "sk-abc"}, + } + } + init_providers(config) + assert mod._default_client is not None + assert "myprov" in mod._default_client.providers + assert mod._default_client.providers["myprov"].api_key == "sk-abc" + + @patch("meeting_room.api_client.httpx.Client") + def test_module_chat_delegates(self, mock_client_cls: MagicMock) -> None: + import meeting_room.api_client as mod + + init_providers({"providers": {"prov": {"base_url": "https://x/v1", "api_key": "k"}}}) + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("delegated") + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + result = mod.chat("prov", "m", []) + assert result["content"] == "delegated" + + @patch("meeting_room.api_client.httpx.Client") + def test_module_chat_stream_delegates(self, mock_client_cls: MagicMock) -> None: + import meeting_room.api_client as mod + + init_providers({"providers": {"prov": {"base_url": "https://x/v1", "api_key": "k"}}}) + lines = _make_stream_lines(["abc"]) + mock_response = MagicMock() + mock_response.iter_lines.return_value = lines + mock_response.raise_for_status = MagicMock() + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + mock_client_cls.return_value.stream.return_value = mock_response + + result = mod.chat_stream("prov", "m", []) + assert result == "abc" + + def test_module_chat_without_init_raises(self) -> None: + import meeting_room.api_client as mod + + mod._default_client = None # ensure clean state + with pytest.raises(RuntimeError, match="init_providers"): + mod.chat("prov", "m", []) + + def test_module_chat_stream_without_init_raises(self) -> None: + import meeting_room.api_client as mod + + mod._default_client = None + with pytest.raises(RuntimeError, match="init_providers"): + mod.chat_stream("prov", "m", []) + + +# --------------------------------------------------------------------------- +# APIClient._get_client — httpx.Client construction +# --------------------------------------------------------------------------- + + +class TestGetClient: + @patch("meeting_room.api_client.httpx.Client") + def test_client_created_with_correct_params(self, mock_client_cls: MagicMock) -> None: + client = APIClient(PROVIDERS) + # Trigger _get_client by calling chat with a mocked response + mock_resp = MagicMock() + mock_resp.json.return_value = _make_chat_response("ok") + mock_resp.raise_for_status = MagicMock() + mock_client_cls.return_value.post.return_value = mock_resp + + client.chat("testprov", "m", []) + + mock_client_cls.assert_called_once_with( + base_url="https://api.test.example.com/v1", + headers={"Authorization": "Bearer sk-test-key-123"}, + timeout=120.0, + ) \ No newline at end of file diff --git a/tests/test_tools.py b/tests/test_tools.py new file mode 100644 index 0000000..34d7be4 --- /dev/null +++ b/tests/test_tools.py @@ -0,0 +1,389 @@ +"""Tests for meeting_room.tools — ToolRegistry class and module-level wrappers.""" + +import os +import textwrap + +import pytest + +from meeting_room.tools import ( + TOOL_SCHEMAS, + TOOL_SETS, + ToolRegistry, + execute_tool, + get_tool_schemas, + set_workdir, +) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def tmp_workdir(tmp_path): + """Create a temporary directory with some files for testing.""" + # Create directories + (tmp_path / "src").mkdir() + (tmp_path / "src" / "sub").mkdir() + (tmp_path / "docs").mkdir() + (tmp_path / ".git" / "objects").mkdir(parents=True) + (tmp_path / "node_modules" / "pkg").mkdir(parents=True) + (tmp_path / "__pycache__").mkdir() + + # Create files + (tmp_path / "README.md").write_text("# Hello\nWorld", encoding="utf-8") + (tmp_path / "src" / "main.py").write_text("def main():\n pass\n", encoding="utf-8") + (tmp_path / "src" / "utils.py").write_text("def add(a, b):\n return a + b\n", encoding="utf-8") + (tmp_path / "src" / "sub" / "deep.py").write_text("# deep file\nx = 42\n", encoding="utf-8") + (tmp_path / "docs" / "notes.txt").write_text("Meeting notes\n", encoding="utf-8") + (tmp_path / "config.yaml").write_text("key: value\n", encoding="utf-8") + # Files in skip dirs + (tmp_path / ".git" / "HEAD").write_text("ref: refs/heads/main\n", encoding="utf-8") + (tmp_path / "node_modules" / "pkg" / "index.js").write_text("module.exports = {};\n", encoding="utf-8") + (tmp_path / "__pycache__" / "cache.pyc").write_text("bytecode", encoding="utf-8") + return tmp_path + + +@pytest.fixture +def reg(tmp_workdir): + """ToolRegistry pointed at the temp workdir.""" + return ToolRegistry(workdir=str(tmp_workdir)) + + +# =========================================================================== +# TOOL_SCHEMAS — structure checks +# =========================================================================== + + +class TestToolSchemas: + def test_schema_names_match_expected(self): + names = {s["function"]["name"] for s in TOOL_SCHEMAS} + expected = {"read_file", "list_files", "search_in_files", "web_search", "web_fetch", "write_file", "run_command"} + assert names == expected + + def test_each_schema_has_required_fields(self): + for s in TOOL_SCHEMAS: + func = s["function"] + assert "name" in func + assert "description" in func + assert "parameters" in func + params = func["parameters"] + assert params["type"] == "object" + assert "properties" in params + + def test_descriptions_are_in_english(self): + """All descriptions must be in English (no Cyrillic).""" + for s in TOOL_SCHEMAS: + desc = s["function"]["description"] + # Check no Cyrillic characters + assert not re.search(r"[а-яА-ЯёЁ]", desc), f"Cyrillic in description: {desc}" + + +# =========================================================================== +# TOOL_SETS — grouping checks +# =========================================================================== + + +class TestToolSets: + def test_all_sets_exist(self): + for name in ("all", "full", "files", "readonly", "web", "none"): + assert name in TOOL_SETS + + def test_all_equals_full(self): + assert TOOL_SETS["all"] == TOOL_SETS["full"] + + def test_none_is_empty(self): + assert TOOL_SETS["none"] == [] + + def test_readonly_subset_of_files(self): + assert set(TOOL_SETS["readonly"]).issubset(set(TOOL_SETS["files"])) + + def test_web_only_web_tools(self): + assert set(TOOL_SETS["web"]) == {"web_search", "web_fetch"} + + +# =========================================================================== +# ToolRegistry — construction and path resolution +# =========================================================================== + + +class TestToolRegistryInit: + def test_default_workdir(self): + r = ToolRegistry() + assert r._workdir == os.path.abspath(".") + + def test_custom_workdir(self, tmp_path): + r = ToolRegistry(workdir=str(tmp_path)) + assert r._workdir == str(tmp_path) + + def test_resolve_absolute_path(self, reg): + abs_path = os.path.abspath(os.path.join(os.sep, "tmp", "some", "file.txt")) + assert reg._resolve_path(abs_path) == abs_path + + def test_resolve_relative_path(self, reg, tmp_workdir): + result = reg._resolve_path("src/main.py") + assert result == os.path.normpath(os.path.join(str(tmp_workdir), "src/main.py")) + + +# =========================================================================== +# get_tool_schemas +# =========================================================================== + + +class TestGetToolSchemas: + def test_all_set(self, reg): + schemas = reg.get_tool_schemas("all") + assert len(schemas) == len(TOOL_SCHEMAS) + + def test_none_set(self, reg): + schemas = reg.get_tool_schemas("none") + assert schemas == [] + + def test_readonly_set(self, reg): + schemas = reg.get_tool_schemas("readonly") + names = {s["function"]["name"] for s in schemas} + assert names == {"read_file", "list_files", "search_in_files"} + + def test_explicit_list(self, reg): + schemas = reg.get_tool_schemas(["read_file", "web_search"]) + names = {s["function"]["name"] for s in schemas} + assert names == {"read_file", "web_search"} + + def test_unknown_set_falls_back_to_all(self, reg): + schemas = reg.get_tool_schemas("nonexistent_set") + assert len(schemas) == len(TOOL_SCHEMAS) + + +# =========================================================================== +# execute_tool dispatch +# =========================================================================== + + +class TestExecuteTool: + def test_unknown_tool(self, reg): + result = reg.execute_tool("nonexistent_tool", {}) + assert "Unknown tool" in result + + def test_execute_read_file(self, reg): + result = reg.execute_tool("read_file", {"path": "README.md"}) + assert "Hello" in result + + def test_execute_write_then_read(self, reg, tmp_workdir): + reg.execute_tool("write_file", {"path": "new_file.txt", "content": "test content"}) + result = reg.execute_tool("read_file", {"path": "new_file.txt"}) + assert "test content" in result + + +# =========================================================================== +# tool_read_file +# =========================================================================== + + +class TestToolReadFile: + def test_read_existing_file(self, reg): + result = reg.tool_read_file("README.md") + assert "Hello" in result + + def test_read_with_offset(self, reg): + result = reg.tool_read_file("README.md", offset=2) + assert "World" in result + # Line numbers start at offset + assert "2|" in result + + def test_read_with_limit(self, reg): + result = reg.tool_read_file("src/main.py", limit=1) + # Should have only 1 content line + lines = [l for l in result.split("\n") if "|" in l] + assert len(lines) == 1 + + def test_read_nonexistent_file(self, reg): + result = reg.tool_read_file("nonexistent.txt") + assert "not found" in result + + def test_read_shows_total_lines(self, reg): + result = reg.tool_read_file("README.md") + assert "lines" in result + + +# =========================================================================== +# tool_list_files — pure Python, Windows-compatible +# =========================================================================== + + +class TestToolListFiles: + def test_list_root(self, reg, tmp_workdir): + result = reg.tool_list_files() + # Must include README.md and config.yaml at root + assert "README.md" in result + assert "config.yaml" in result + + def test_list_skips_git(self, reg): + result = reg.tool_list_files() + # .git contents should NOT appear + assert ".git" not in result or "HEAD" not in result + + def test_list_skips_node_modules(self, reg): + result = reg.tool_list_files() + assert "index.js" not in result + + def test_list_skips_pycache(self, reg): + result = reg.tool_list_files() + assert "cache.pyc" not in result + + def test_list_with_pattern(self, reg): + result = reg.tool_list_files(pattern="*.py") + lines = result.strip().split("\n") if result.strip() else [] + # Only .py files should appear + for line in lines: + if line.strip(): + assert line.endswith(".py") or os.path.isdir(line), f"Non-Python file in filtered results: {line}" + + def test_list_with_path(self, reg, tmp_workdir): + result = reg.tool_list_files(path="src") + assert "main.py" in result + assert "utils.py" in result + + def test_list_depth_limit(self, reg): + result = reg.tool_list_files(max_depth=1) + # With depth 1, should include root dir and its immediate files + # but NOT deeply nested files like src/sub/deep.py + # (it may list the src directory itself) + assert "config.yaml" in result + + def test_list_nonexistent_dir(self, reg): + result = reg.tool_list_files(path="nonexistent_dir") + assert "not found" in result + + def test_no_find_command_used(self): + """Verify list_files implementation does not shell out to find.""" + import inspect + source = inspect.getsource(ToolRegistry.tool_list_files) + assert "subprocess" not in source + assert '"find"' not in source + + +# =========================================================================== +# tool_search_in_files — pure Python, Windows-compatible +# =========================================================================== + + +class TestToolSearchInFiles: + def test_search_basic(self, reg): + result = reg.tool_search_in_files(query="def main") + assert "main.py" in result + + def test_search_regex(self, reg): + result = reg.tool_search_in_files(query=r"def \w+\(") + # Should find function definitions + assert "def " in result + + def test_search_with_file_pattern(self, reg): + result = reg.tool_search_in_files(query="value", file_pattern="*.yaml") + assert "config.yaml" in result + + def test_search_no_matches(self, reg): + result = reg.tool_search_in_files(query="zzz_nonexistent_pattern_xyz") + assert "No matches" in result + + def test_search_invalid_regex(self, reg): + result = reg.tool_search_in_files(query="[invalid") + assert "invalid regex" in result + + def test_search_nonexistent_dir(self, reg): + result = reg.tool_search_in_files(query="test", path="nonexistent_dir") + assert "not found" in result + + def test_search_skips_git(self, reg): + result = reg.tool_search_in_files(query="refs") + # Should not find content inside .git/HEAD + # (unless something else matches, but the file .git/HEAD itself should be skipped) + lines = result.strip().split("\n") if result.strip() else [] + for line in lines: + assert ".git" not in line + + def test_search_shows_line_numbers(self, reg): + result = reg.tool_search_in_files(query="def add") + # Format is filepath:linenum:content + assert ":" in result + + def test_no_grep_command_used(self): + """Verify search_in_files implementation does not shell out to grep.""" + import inspect + source = inspect.getsource(ToolRegistry.tool_search_in_files) + assert "subprocess" not in source + assert '"grep"' not in source + + +# =========================================================================== +# tool_write_file +# =========================================================================== + + +class TestToolWriteFile: + def test_write_new_file(self, reg, tmp_workdir): + result = reg.tool_write_file("output.txt", "hello world") + assert "Wrote" in result + assert os.path.isfile(os.path.join(str(tmp_workdir), "output.txt")) + + def test_write_creates_subdirs(self, reg, tmp_workdir): + result = reg.tool_write_file("sub/dir/test.txt", "nested") + assert "Wrote" in result + assert os.path.isfile(os.path.join(str(tmp_workdir), "sub", "dir", "test.txt")) + + def test_write_overwrites(self, reg, tmp_workdir): + reg.tool_write_file("overwrite.txt", "first") + reg.tool_write_file("overwrite.txt", "second") + result = reg.tool_read_file("overwrite.txt") + assert "second" in result + + +# =========================================================================== +# tool_run_command +# =========================================================================== + + +class TestToolRunCommand: + def test_run_simple_command(self, reg): + result = reg.tool_run_command("echo hello") + assert "hello" in result.lower() + + def test_run_command_timeout(self, reg): + # Use a command that sleeps longer than the timeout + result = reg.tool_run_command("ping -n 10 127.0.0.1" if os.name == "nt" else "sleep 30", timeout=1) + assert "Timeout" in result or "timeout" in result.lower() + + def test_run_command_nonzero_exit(self, reg): + result = reg.tool_run_command("exit 1" if os.name != "nt" else "cmd /c exit 1") + # Should show exit code + assert "exit code" in result or result.strip() != "" + + +# =========================================================================== +# Module-level backward compatibility wrappers +# =========================================================================== + + +class TestModuleLevelWrappers: + def test_set_workdir_and_execute(self, tmp_workdir): + set_workdir(str(tmp_workdir)) + result = execute_tool("read_file", {"path": "README.md"}) + assert "Hello" in result + + def test_get_tool_schemas_default(self): + # Without set_workdir, should still work (defaults to ".") + schemas = get_tool_schemas("all") + assert len(schemas) == len(TOOL_SCHEMAS) + + def test_get_tool_schemas_module_level(self): + schemas = get_tool_schemas("readonly") + names = {s["function"]["name"] for s in schemas} + assert "read_file" in names + assert "web_search" not in names + + def test_execute_unknown_tool(self): + result = execute_tool("does_not_exist", {}) + assert "Unknown tool" in result + + +# Need re import for Cyrillic check +import re \ No newline at end of file