fix(search): serialize filename walks by root

This commit is contained in:
Royalaid
2026-08-28 19:38:36 -07:00
committed by kshitij
parent af98771b45
commit 9e47ac17bd
2 changed files with 340 additions and 15 deletions

View 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

View File

@@ -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