diff --git a/tests/tools/test_search_native_rg.py b/tests/tools/test_search_native_rg.py index afd6a11a14..de84ee985d 100644 --- a/tests/tools/test_search_native_rg.py +++ b/tests/tools/test_search_native_rg.py @@ -26,16 +26,29 @@ def tree(tmp_path): return tmp_path -def _ops(tree, spy): - env = LocalEnvironment(cwd=str(tree)) - real = type(env).execute.__get__(env, type(env)) +@pytest.fixture(scope="module") +def _local_env(tmp_path_factory): + """One real LocalEnvironment per module (constructing one costs ~0.8 s).""" + return LocalEnvironment(cwd=str(tmp_path_factory.mktemp("native-rg"))) - def recording(command, *a, **kw): - spy.append(command) - return real(command, *a, **kw) - env.execute = recording - return ShellFileOperations(env, cwd=str(tree)) +@pytest.fixture +def ops_factory(_local_env): + """``make(tree, spy)`` → ShellFileOperations over the shared env, every execute recorded in ``spy``.""" + real = type(_local_env).execute.__get__(_local_env, type(_local_env)) + + def make(tree, spy): + _local_env.cwd = str(tree) + + def recording(command, *a, **kw): + spy.append(command) + return real(command, *a, **kw) + + _local_env.execute = recording + return ShellFileOperations(_local_env, cwd=str(tree)) + + yield make + _local_env.__dict__.pop("execute", None) def _normalized(result): @@ -46,7 +59,7 @@ def _normalized(result): return d -def test_native_search_never_touches_the_shell_and_matches_shell_results(tree, monkeypatch): +def test_native_search_never_touches_the_shell_and_matches_shell_results(tree, ops_factory, monkeypatch): cases = [ dict(pattern="needle", path=str(tree)), # (no offset/limit slicing here: rg's parallel walk orders files @@ -59,17 +72,17 @@ def test_native_search_never_touches_the_shell_and_matches_shell_results(tree, m ] for case in cases: monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") - shell = _normalized(_ops(tree, []).search(**case)) + shell = _normalized(ops_factory(tree, []).search(**case)) monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "1") calls = [] - native = _normalized(_ops(tree, calls).search(**case)) + native = _normalized(ops_factory(tree, calls).search(**case)) assert native == shell, case # rg resolution (``command -v rg``) still goes through the shell once; the # existence probe and the rg pipeline itself must not. assert not [c for c in calls if "pipefail" in c or c.startswith("test -e")], case -def test_native_runner_honours_deadline_and_interrupt_while_rg_is_silent(tree, monkeypatch): +def test_native_runner_honours_deadline_and_interrupt_while_rg_is_silent(tree, ops_factory, monkeypatch): """A search producing no output must still stop at the deadline / on /stop (the shell path gets this from ``_wait_for_process``).""" import threading @@ -77,7 +90,7 @@ def test_native_runner_honours_deadline_and_interrupt_while_rg_is_silent(tree, m from tools import interrupt - ops = _ops(tree, []) + ops = ops_factory(tree, []) started = time.monotonic() result = ops._run_rg_native(["sh", "-c", "'sleep 30'"], 5, timeout=1) assert result.exit_code == 124 and time.monotonic() - started < 5 @@ -92,10 +105,10 @@ def test_native_runner_honours_deadline_and_interrupt_while_rg_is_silent(tree, m assert result.exit_code == 130 and time.monotonic() - started < 5 -def test_kill_switch_routes_search_back_to_the_shell(tree, monkeypatch): +def test_kill_switch_routes_search_back_to_the_shell(tree, ops_factory, monkeypatch): monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") calls = [] - result = _ops(tree, calls).search(pattern="needle", path=str(tree)) + result = ops_factory(tree, calls).search(pattern="needle", path=str(tree)) assert result.total_count == 4 assert any(c.startswith("test -e") for c in calls) assert any("pipefail" in c and "rg" in c for c in calls) diff --git a/tools/file_operations_search.py b/tools/file_operations_search.py index 296a6399e5..f111a316a0 100644 --- a/tools/file_operations_search.py +++ b/tools/file_operations_search.py @@ -387,6 +387,18 @@ class SearchMixin: # left the pipeline at 0 unless rg itself already failed. return ExecuteResult(stdout=stdout, exit_code=0 if bounded.is_set() else proc.returncode) + def _run_rg_bounded(self, words: List[str], fetch_limit: int, timeout: int, *, + merge_stderr: bool = False, native_ok: bool = True, + shell_prefix: str = "") -> ExecuteResult: + """Run an rg command (shell-quoted words) and keep the first ``fetch_limit`` + lines: natively on a local POSIX host, else through the backend shell as + `` | head -n N``. ``native_ok=False`` keeps a form the native + lane cannot express (the multi-root ``cd`` prefix); ``shell_prefix`` is shell-only.""" + if native_ok and self._native_read_enabled(): + return self._run_rg_native(words, fetch_limit, timeout, merge_stderr=merge_stderr) + stderr = "" if merge_stderr else " 2>/dev/null" + return self._exec(f"{shell_prefix}{' '.join(words)}{stderr} | head -n {fetch_limit}", timeout=timeout) + def _quote_executable(self, executable: str) -> str: """Quote an executable without leaking controller path semantics.""" if re.fullmatch(r"[A-Za-z0-9_.-]+", executable): @@ -588,10 +600,7 @@ class SearchMixin: glob_expr_probe = glob_expr probe_words = [rg, flags, "--count-matches", glob_expr_probe, self._escape_shell_arg(pattern), self._escape_native_tool_arg(path)] - if self._native_read_enabled(): - probe = self._run_rg_native(probe_words, 50, timeout=30) - else: - probe = self._exec(" ".join(probe_words) + " 2>/dev/null | head -50", timeout=30) + probe = self._run_rg_bounded(probe_words, 50, timeout=30) total, per_file = 0, [] for line in (probe.stdout or "").strip().splitlines(): p, _sep, n = line.rpartition(":") @@ -771,10 +780,8 @@ class SearchMixin: # ``--`` terminates options so a dash-prefixed root is never parsed as a flag. rg_cmd = (f"{rg} --files{sort_arg} -g {self._escape_shell_arg(glob_pattern)}" f"{exclusion_args} -- {root_args}") - if not scoped_common and self._native_read_enabled(): - result = self._run_rg_native([rg_cmd], fetch_limit, timeout=60) - else: - result = self._exec(f"set -o pipefail; {cd_prefix}{rg_cmd} 2>/dev/null | head -n {fetch_limit}", timeout=60) + result = self._run_rg_bounded([rg_cmd], fetch_limit, timeout=60, native_ok=not scoped_common, + shell_prefix=f"set -o pipefail; {cd_prefix}") stdout, limit_reason = _search_stdout_and_limit(result) all_files = [f for f in stdout.splitlines() if f] if scoped_common: @@ -830,13 +837,14 @@ class SearchMixin: (grep): bounds giant single-line matches at the pipe layer; skipped for files_only/count where lines are paths/counts.""" fetch_limit = limit + offset + (200 if context > 0 else 0) - if not line_cap and self._native_read_enabled(): - result = self._run_rg_native(cmd_parts, fetch_limit, timeout=60, merge_stderr=True) - return _parse_search_output(result, output_mode, limit, offset, context, warning=warning) - parts = cmd_parts + ["|", "head", "-n", str(fetch_limit)] - if line_cap and output_mode not in ("files_only", "count"): - parts += ["|", "cut", "-c1-2000"] - result = self._exec("set -o pipefail; " + " ".join(parts), timeout=60) + if line_cap: # grep/find pipelines: shell only, with the column cap + parts = cmd_parts + ["|", "head", "-n", str(fetch_limit)] + if output_mode not in ("files_only", "count"): + parts += ["|", "cut", "-c1-2000"] + result = self._exec("set -o pipefail; " + " ".join(parts), timeout=60) + else: + result = self._run_rg_bounded(cmd_parts, fetch_limit, timeout=60, merge_stderr=True, + shell_prefix="set -o pipefail; ") return _parse_search_output(result, output_mode, limit, offset, context, warning=warning) def _search_with_rg(self, pattern: str, path: str, file_glob: Optional[str],