fix(search): serialize filename walks by root
This commit is contained in:
238
tests/tools/test_search_files_cpu_windows.py
Normal file
238
tests/tools/test_search_files_cpu_windows.py
Normal file
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user