diff --git a/tests/tools/test_search_files_cpu_windows.py b/tests/tools/test_search_files_cpu_windows.py new file mode 100644 index 0000000000..f9ba10e94b --- /dev/null +++ b/tests/tools/test_search_files_cpu_windows.py @@ -0,0 +1,238 @@ +"""Concurrency admission tests for expensive filename walks.""" + +from concurrent.futures import ThreadPoolExecutor +import threading +import types + +import pytest + +from tools.environments.local import LocalEnvironment +from tools.file_operations import ( + _ACTIVE_FILENAME_SEARCH_ROOTS, + _FILENAME_SEARCH_ADMISSION, + _normalized_filename_search_root, + SearchResult, + ShellFileOperations, +) +from tools.interrupt import set_interrupt + + +class RemoteEnvironment: + is_local = False + cwd = "/workspace" + + def execute(self, command, **kwargs): + raise AssertionError(f"unexpected backend command: {command}") + + +def _operations(env, scan): + operations = ShellFileOperations(env) + operations._resolve_command = lambda command: "/usr/bin/rg" if command == "rg" else None + operations._search_files_rg = types.MethodType(scan, operations) + return operations + + +def test_same_backend_class_and_root_serialize_five_filename_walks(): + entered = threading.Event() + release = threading.Event() + counter_lock = threading.Lock() + active = 0 + maximum_active = 0 + completed = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal active, maximum_active, completed + with counter_lock: + active += 1 + maximum_active = max(maximum_active, active) + entered.set() + assert release.wait(5) + with counter_lock: + active -= 1 + completed += 1 + return SearchResult(files=[str(path)], total_count=1) + + operations = [_operations(RemoteEnvironment(), scan) for _ in range(5)] + with ThreadPoolExecutor(max_workers=5) as pool: + futures = [ + pool.submit(operation._search_files, "*.py", "/repo", 50, 0) + for operation in operations + ] + assert entered.wait(5) + release.set() + results = [future.result(timeout=5) for future in futures] + + assert all(result.error is None for result in results) + assert completed == 5 + assert maximum_active == 1 + + +def test_different_roots_can_enter_filename_walks_together(): + both_entered = threading.Barrier(2) + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + both_entered.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(RemoteEnvironment(), scan) + with ThreadPoolExecutor(max_workers=2) as pool: + futures = [ + pool.submit(first._search_files, "*.py", "/one", 50, 0), + pool.submit(second._search_files, "*.py", "/two", 50, 0), + ] + assert [future.result(timeout=5).error for future in futures] == [None, None] + + +def test_different_backend_classes_can_walk_the_same_root_together(): + class OtherRemoteEnvironment(RemoteEnvironment): + pass + + both_entered = threading.Barrier(2) + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + both_entered.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(OtherRemoteEnvironment(), scan) + with ThreadPoolExecutor(max_workers=2) as pool: + futures = [ + pool.submit(first._search_files, "*.py", "/same", 50, 0), + pool.submit(second._search_files, "*.py", "/same", 50, 0), + ] + assert [future.result(timeout=5).error for future in futures] == [None, None] + + +def test_overlapping_multi_root_sets_are_claimed_atomically(monkeypatch): + first_entered = threading.Event() + release_first = threading.Event() + second_waiting = threading.Event() + lock = threading.Lock() + active = 0 + maximum_active = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal active, maximum_active + with lock: + active += 1 + maximum_active = max(maximum_active, active) + if path == ["/a", "/b"]: + first_entered.set() + if path == ["/a", "/b"]: + assert release_first.wait(5) + with lock: + active -= 1 + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(RemoteEnvironment(), scan) + original_wait = _FILENAME_SEARCH_ADMISSION.wait + + def observed_wait(timeout=None): + second_waiting.set() + return original_wait(timeout) + + monkeypatch.setattr(_FILENAME_SEARCH_ADMISSION, "wait", observed_wait) + with ThreadPoolExecutor(max_workers=2) as pool: + first_future = pool.submit(first._search_files, "*.py", ["/a", "/b"], 50, 0) + assert first_entered.wait(5) + second_future = pool.submit(second._search_files, "*.py", ["/b", "/c"], 50, 0) + assert second_waiting.wait(5) + release_first.set() + assert first_future.result(timeout=5).error is None + assert second_future.result(timeout=5).error is None + + assert maximum_active == 1 + + +def test_interrupted_waiter_returns_without_dispatch_or_late_dispatch(monkeypatch): + holder_entered = threading.Event() + release_holder = threading.Event() + waiter_waiting = threading.Event() + waiter_tid = [] + dispatches = [] + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + dispatches.append(threading.get_ident()) + holder_entered.set() + assert release_holder.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + holder = _operations(RemoteEnvironment(), scan) + waiter = _operations(RemoteEnvironment(), scan) + + original_wait = _FILENAME_SEARCH_ADMISSION.wait + + def observed_wait(timeout=None): + waiter_waiting.set() + return original_wait(timeout) + + monkeypatch.setattr(_FILENAME_SEARCH_ADMISSION, "wait", observed_wait) + + def run_waiter(): + waiter_tid.append(threading.get_ident()) + return waiter._search_files("*.py", "/repo", 50, 0) + + with ThreadPoolExecutor(max_workers=2) as pool: + holder_future = pool.submit(holder._search_files, "*.py", "/repo", 50, 0) + assert holder_entered.wait(5) + waiter_future = pool.submit(run_waiter) + assert waiter_waiting.wait(5) + set_interrupt(True, waiter_tid[0]) + try: + interrupted = waiter_future.result(timeout=5) + assert "interrupted" in (interrupted.error or "").lower() + assert len(dispatches) == 1 + release_holder.set() + assert holder_future.result(timeout=5).error is None + assert len(dispatches) == 1 + finally: + set_interrupt(False, waiter_tid[0]) + release_holder.set() + + +@pytest.mark.parametrize("raised", [Exception, KeyboardInterrupt, SystemExit, BaseException]) +def test_admission_releases_after_every_base_exception_path(raised): + attempts = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise raised("engine failed") + return SearchResult(files=[str(path)], total_count=1) + + operations = _operations(RemoteEnvironment(), scan) + with pytest.raises(raised, match="engine failed"): + operations._search_files("*.py", "/repo", 50, 0) + + result = operations._search_files("*.py", "/repo", 50, 0) + assert result.error is None + assert attempts == 2 + assert _ACTIVE_FILENAME_SEARCH_ROOTS == set() + + +def test_remote_roots_are_normalized_lexically_against_backend_cwd(monkeypatch): + env = RemoteEnvironment() + monkeypatch.setattr( + "tools.file_operations.os.path.abspath", + lambda path: (_ for _ in ()).throw(AssertionError("controller resolution used")), + ) + + relative = _normalized_filename_search_root(env, "repo/../repo", "/controller") + absolute = _normalized_filename_search_root(env, "/workspace/repo", "/controller") + + assert relative == "/workspace/repo" + assert absolute == relative + + +@pytest.mark.windows_only +def test_windows_local_root_spellings_share_one_normalized_key(): + env = LocalEnvironment.__new__(LocalEnvironment) + env.cwd = "C:/Repo" + + native = _normalized_filename_search_root(env, r"C:\Repo\src\..", "C:/ignored") + msys = _normalized_filename_search_root(env, "/c/Repo", "C:/ignored") + + assert native == msys diff --git a/tools/file_operations.py b/tools/file_operations.py index 6f71cd8ac1..a4adeb40f6 100644 --- a/tools/file_operations.py +++ b/tools/file_operations.py @@ -36,6 +36,7 @@ import difflib import hashlib import json import logging +import threading import unicodedata from abc import ABC, abstractmethod from dataclasses import dataclass, field @@ -49,6 +50,7 @@ from agent.file_safety import ( get_write_denied_error, is_write_denied as _shared_is_write_denied, ) +from tools import interrupt as tool_interrupt logger = logging.getLogger(__name__) @@ -70,6 +72,70 @@ _MACOS_TCC_PROTECTED_HOME_DIRS = ( ) +_FILENAME_SEARCH_ADMISSION = threading.Condition() +_ACTIVE_FILENAME_SEARCH_ROOTS: set[tuple[str, str, str]] = set() +_FILENAME_SEARCH_WAIT_SECONDS = 0.05 + + +def _normalized_filename_search_root(env: Any, root: str, fallback_cwd: str) -> str: + """Normalize a filename-walk root without resolving remote paths locally.""" + from tools.environments.local import LocalEnvironment, _IS_WINDOWS, _msys_to_windows_path + + cwd = getattr(env, "cwd", None) or fallback_cwd + if isinstance(env, LocalEnvironment): + if _IS_WINDOWS: + root = _msys_to_windows_path(root) + cwd = _msys_to_windows_path(cwd) + if not os.path.isabs(root): + root = os.path.join(cwd, root) + return os.path.normcase(os.path.abspath(os.path.normpath(root))) + + if not posixpath.isabs(root): + root = posixpath.join(cwd, root) + return posixpath.normpath(root) + + +def _filename_search_root_keys( + env: Any, roots: List[str], fallback_cwd: str +) -> tuple[tuple[str, str, str], ...]: + """Return unique backend/root admission keys in deterministic order.""" + env_type = type(env) + return tuple(sorted({ + ( + env_type.__module__, + env_type.__qualname__, + _normalized_filename_search_root(env, root, fallback_cwd), + ) + for root in roots + })) + + +def _acquire_filename_search_roots( + keys: tuple[tuple[str, str, str], ...], +) -> bool: + """Atomically claim every key, polling for thread-scoped interruption.""" + with _FILENAME_SEARCH_ADMISSION: + while any(key in _ACTIVE_FILENAME_SEARCH_ROOTS for key in keys): + if tool_interrupt.is_interrupted(): + return False + _FILENAME_SEARCH_ADMISSION.wait(_FILENAME_SEARCH_WAIT_SECONDS) + if tool_interrupt.is_interrupted(): + return False + if tool_interrupt.is_interrupted(): + return False + _ACTIVE_FILENAME_SEARCH_ROOTS.update(keys) + return True + + +def _release_filename_search_roots( + keys: tuple[tuple[str, str, str], ...], +) -> None: + """Release a completed walk and leave no idle per-root state behind.""" + with _FILENAME_SEARCH_ADMISSION: + _ACTIVE_FILENAME_SEARCH_ROOTS.difference_update(keys) + _FILENAME_SEARCH_ADMISSION.notify_all() + + def _macos_protected_search_exclusions( path: str, *, @@ -3547,16 +3613,11 @@ class ShellFileOperations(FileOperations): return None if target == "files": - # A file search across several roots is one global rg traversal so + # A file search across several roots is one global traversal so # modified ordering and pagination are exact across the whole set. - if self._has_command("rg"): - resolved = self._resolve_command("rg") or "rg" - merged = self._search_files_rg( - pattern.split("/")[-1], existing, limit, offset, order, - rg_executable=resolved, - ) - else: - merged = self._search_files(pattern, existing, limit, offset, order) + # Route every engine through _search_files so root admission wraps + # the actual rg/find invocation for this multi-root request. + merged = self._search_files(pattern, existing, limit, offset, order) else: merged = SearchResult() for root in existing: @@ -3708,17 +3769,34 @@ class ShellFileOperations(FileOperations): else: search_pattern = pattern.split('/')[-1] + roots = [path] if isinstance(path, str) else path + # Prefer ripgrep: bounded parallel traversal with ignore semantics. + # Resolve the engine and exact-order capability before admission so a + # queued request does not occupy a root while doing command discovery. if self._has_command("rg"): - return self._search_files_rg( - search_pattern, path, limit, offset, order, - rg_executable=self._resolve_command("rg") or "rg", - ) + rg_executable = self._resolve_command("rg") or "rg" + if order == "modified": + capability_error = self._modified_rg_capability_error(rg_executable) + if capability_error: + return SearchResult(error=capability_error) + keys = _filename_search_root_keys(self.env, roots, self.cwd) + if not _acquire_filename_search_roots(keys): + return SearchResult(error=( + "File search was interrupted while waiting for another filename " + "search on the same root. Retry when ready." + )) + try: + return self._search_files_rg( + search_pattern, path, limit, offset, order, + rg_executable=rg_executable, + ) + finally: + _release_filename_search_roots(keys) # A local find traversal rooted at/above the user's home or at a # filesystem root can consume minutes and prompt on protected paths. # Refuse before invoking find. Controller paths never classify remotes. - roots = [path] if isinstance(path, str) else path if any(self._is_broad_local_search_root(root) for root in roots): return SearchResult(error=( "Broad local file search without ripgrep is disabled because " @@ -3773,7 +3851,16 @@ class ShellFileOperations(FileOperations): + f" -print 2>/dev/null | head -n {fetch_limit}" ) - result = self._exec(cmd, timeout=60) + keys = _filename_search_root_keys(self.env, roots, self.cwd) + if not _acquire_filename_search_roots(keys): + return SearchResult(error=( + "File search was interrupted while waiting for another filename " + "search on the same root. Retry when ready." + )) + try: + result = self._exec(cmd, timeout=60) + finally: + _release_filename_search_roots(keys) stdout, limit_reason = _search_stdout_and_limit(result) # Parse before classifying exit 141: with pipefail, a bounded producer