fix(tools): bound find_files scans

This commit is contained in:
chengyongru
2026-08-25 18:40:49 +08:00
committed by chengyongru
parent e308f7fdd4
commit 649e3958c5
2 changed files with 402 additions and 63 deletions
+249 -63
View File
@@ -4,9 +4,13 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import fnmatch import fnmatch
import heapq
import os import os
import re import re
import threading
import time
from collections import deque from collections import deque
from contextlib import suppress from contextlib import suppress
from dataclasses import dataclass from dataclasses import dataclass
@@ -57,6 +61,43 @@ class _PendingContextMatch:
remaining_after: int remaining_after: int
@dataclass(slots=True)
class _FindFilesEntry:
path: Path
rel_path: str
display_path: str
name: str
is_dir: bool
class _FindFilesCancelledError(Exception):
"""Stop a worker scan after its owning async task was cancelled."""
class _FindFilesBudgetExceededError(Exception):
"""Stop an unbounded filesystem scan at its configured budget."""
@dataclass(slots=True)
class _FindFilesBudget:
cancelled: threading.Event
deadline: float
max_paths: int
scanned_paths: int = 0
def checkpoint(self) -> None:
if self.cancelled.is_set():
raise _FindFilesCancelledError
if time.monotonic() >= self.deadline:
raise _FindFilesBudgetExceededError("time")
def visit_path(self) -> None:
self.checkpoint()
self.scanned_paths += 1
if self.scanned_paths > self.max_paths:
raise _FindFilesBudgetExceededError("paths")
def _normalize_pattern(pattern: str) -> str: def _normalize_pattern(pattern: str) -> str:
return pattern.strip().replace("\\", "/") return pattern.strip().replace("\\", "/")
@@ -150,6 +191,8 @@ class _SearchTool(_FsTool):
class FindFilesTool(_SearchTool): class FindFilesTool(_SearchTool):
"""Find files by path fragment, glob, or type.""" """Find files by path fragment, glob, or type."""
_scopes = {"core", "subagent"} _scopes = {"core", "subagent"}
_MAX_SCAN_PATHS = 500_000
_MAX_SCAN_SECONDS = 30.0
@property @property
def name(self) -> str: def name(self) -> str:
@@ -211,19 +254,101 @@ class FindFilesTool(_SearchTool):
}, },
} }
def _iter_paths(self, root: Path, *, include_dirs: bool) -> Iterable[Path]: def _entry(self, path: Path, root: Path, *, is_dir: bool) -> _FindFilesEntry:
display_path = self._display_path(path, root)
return _FindFilesEntry(
path=path,
rel_path=path.relative_to(root).as_posix(),
display_path=display_path,
name=path.name,
is_dir=is_dir,
)
def _push_directory_entries(
self,
directory: Path,
root: Path,
frontier: list[tuple[str, int, _FindFilesEntry]],
sequence: int,
budget: _FindFilesBudget,
) -> int:
budget.checkpoint()
try:
with os.scandir(directory) as entries:
for raw_entry in entries:
budget.visit_path()
try:
is_dir = raw_entry.is_dir(follow_symlinks=False)
# os.walk yields special files and broken file symlinks,
# but does not descend into directory symlinks by default.
if not is_dir and raw_entry.is_symlink() and raw_entry.is_dir():
continue
except OSError:
continue
if is_dir and raw_entry.name in self._IGNORE_DIRS:
continue
entry = self._entry(Path(raw_entry.path), root, is_dir=is_dir)
sort_path = entry.display_path + ("/" if is_dir else "")
heapq.heappush(frontier, (sort_path, sequence, entry))
sequence += 1
except OSError:
# os.walk silently skips directories that cannot be listed. Preserve
# that behavior while still allowing cancellation and budget errors
# to propagate from the explicit checkpoints above.
pass
return sequence
def _iter_paths(
self,
root: Path,
*,
include_dirs: bool,
budget: _FindFilesBudget,
) -> Iterable[_FindFilesEntry]:
budget.checkpoint()
if root.is_file(): if root.is_file():
yield root budget.visit_path()
yield self._entry(root, root.parent, is_dir=False)
return return
if include_dirs: if include_dirs:
yield root yield self._entry(root, root, is_dir=True)
for dirpath, dirnames, filenames in os.walk(root):
dirnames[:] = sorted(d for d in dirnames if d not in self._IGNORE_DIRS) frontier: list[tuple[str, int, _FindFilesEntry]] = []
current = Path(dirpath) sequence = self._push_directory_entries(root, root, frontier, 0, budget)
if include_dirs and current != root: while frontier:
yield current budget.checkpoint()
for filename in sorted(filenames): _, _, entry = heapq.heappop(frontier)
yield current / filename if entry.is_dir:
if include_dirs:
yield entry
sequence = self._push_directory_entries(
entry.path,
root,
frontier,
sequence,
budget,
)
else:
yield entry
@staticmethod
def _matches_entry(
entry: _FindFilesEntry,
*,
query: str | None,
glob: str | None,
file_type: str | None,
) -> bool:
if glob and not _match_glob(entry.rel_path, entry.name, glob):
return False
if entry.is_dir:
if file_type:
return False
elif not _matches_type(entry.name, file_type):
return False
return _matches_query(entry.display_path, query)
async def execute( async def execute(
self, self,
@@ -237,66 +362,127 @@ class FindFilesTool(_SearchTool):
offset: int = 0, offset: int = 0,
**kwargs: Any, **kwargs: Any,
) -> str: ) -> str:
cancelled = threading.Event()
try: try:
target = self._resolve(path or ".") return await asyncio.to_thread(
if not target.exists(): self._execute_sync,
return ToolResult.error(f"Error: Path not found: {path}") path=path,
if not (target.is_dir() or target.is_file()): query=query,
return ToolResult.error(f"Error: Unsupported path: {path}") glob=glob,
file_type=type,
if sort not in {"path", "modified"}: include_dirs=include_dirs,
return ToolResult.error("Error: sort must be 'path' or 'modified'") sort=sort,
head_limit=head_limit,
limit = ( offset=offset,
_DEFAULT_FILE_HEAD_LIMIT cancelled=cancelled,
if head_limit is None
else None if head_limit == 0 else head_limit
) )
root = target if target.is_dir() else target.parent except asyncio.CancelledError:
matches: list[tuple[str, float]] = [] cancelled.set()
raise
for candidate in self._iter_paths(target, include_dirs=include_dirs):
if candidate.is_dir() and not include_dirs:
continue
rel_path = candidate.relative_to(root).as_posix()
display_path = self._display_path(candidate, root)
name = candidate.name
if glob and not _match_glob(rel_path, name, glob):
continue
if candidate.is_file() and not _matches_type(name, type):
continue
if candidate.is_dir() and type:
continue
if not _matches_query(display_path, query):
continue
try:
mtime = candidate.stat().st_mtime
except OSError:
mtime = 0.0
suffix = "/" if candidate.is_dir() else ""
matches.append((display_path + suffix, mtime))
if sort == "modified":
matches.sort(key=lambda item: (-item[1], item[0]))
else:
matches.sort(key=lambda item: item[0])
paths = [item[0] for item in matches]
paged, truncated = _paginate(paths, limit, offset)
if not paged:
return "No files found"
result = "\n".join(paged)
note = _pagination_note(limit, offset, truncated)
if note:
result += "\n\n" + note
return result
except PermissionError as e: except PermissionError as e:
return ToolResult.error(f"Error: {e}") return ToolResult.error(f"Error: {e}")
except Exception as e: except Exception as e:
return ToolResult.error(f"Error finding files: {e}") return ToolResult.error(f"Error finding files: {e}")
def _execute_sync(
self,
*,
path: str,
query: str | None,
glob: str | None,
file_type: str | None,
include_dirs: bool,
sort: str,
head_limit: int | None,
offset: int,
cancelled: threading.Event,
) -> str:
started_at = time.monotonic()
if cancelled.is_set():
raise _FindFilesCancelledError
target = self._resolve(path or ".")
if not target.exists():
return ToolResult.error(f"Error: Path not found: {path}")
if not (target.is_dir() or target.is_file()):
return ToolResult.error(f"Error: Unsupported path: {path}")
if sort not in {"path", "modified"}:
return ToolResult.error("Error: sort must be 'path' or 'modified'")
limit = (
_DEFAULT_FILE_HEAD_LIMIT
if head_limit is None
else None if head_limit == 0 else head_limit
)
budget = _FindFilesBudget(
cancelled=cancelled,
deadline=started_at + self._MAX_SCAN_SECONDS,
max_paths=self._MAX_SCAN_PATHS,
)
def matching_entries() -> Iterator[tuple[str, float]]:
for entry in self._iter_paths(
target,
include_dirs=include_dirs,
budget=budget,
):
if not self._matches_entry(
entry,
query=query,
glob=glob,
file_type=file_type,
):
continue
mtime = 0.0
if sort == "modified":
try:
mtime = entry.path.stat().st_mtime
except OSError:
pass
suffix = "/" if entry.is_dir else ""
yield entry.display_path + suffix, mtime
matches: list[tuple[str, float]]
try:
if sort == "modified":
if limit is None:
matches = sorted(matching_entries(), key=lambda item: (-item[1], item[0]))
else:
selection_size = offset + limit + 1
matches = heapq.nsmallest(
selection_size,
matching_entries(),
key=lambda item: (-item[1], item[0]),
)
else:
selection_size = None if limit is None else offset + limit + 1
matches = []
for match in matching_entries():
matches.append(match)
if selection_size is not None and len(matches) >= selection_size:
break
budget.checkpoint()
except _FindFilesBudgetExceededError as exc:
if str(exc) == "paths":
detail = f"{self._MAX_SCAN_PATHS} paths"
else:
detail = f"{self._MAX_SCAN_SECONDS:g} seconds"
return ToolResult.error(
f"Error: find_files scan exceeded {detail}; "
"narrow path, query, glob, or type and retry."
)
paths = [item[0] for item in matches]
paged, truncated = _paginate(paths, limit, offset)
if not paged:
return "No files found"
result = "\n".join(paged)
note = _pagination_note(limit, offset, truncated)
if note:
result += "\n\n" + note
return result
class GrepTool(_SearchTool): class GrepTool(_SearchTool):
"""Search text and document contents using a regex-like pattern.""" """Search text and document contents using a regex-like pattern."""
+153
View File
@@ -2,8 +2,10 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import os import os
import re import re
import threading
import time import time
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
@@ -18,6 +20,11 @@ from nanobot.agent.tools.web import WebSearchTool
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import WebSearchConfig from nanobot.config.schema import WebSearchConfig
from nanobot.providers.base import GenerationSettings from nanobot.providers.base import GenerationSettings
from nanobot.security.workspace_access import (
bind_workspace_scope,
default_workspace_scope,
reset_workspace_scope,
)
from nanobot.utils.llm_runtime import LLMRuntime from nanobot.utils.llm_runtime import LLMRuntime
@@ -98,6 +105,152 @@ async def test_find_files_rejects_paths_outside_workspace(tmp_path: Path) -> Non
assert result.startswith("Error:") assert result.startswith("Error:")
@pytest.mark.asyncio
async def test_find_files_worker_preserves_current_workspace_scope(tmp_path: Path) -> None:
project = tmp_path / "project"
project.mkdir()
(project / "inside.txt").write_text("ok\n", encoding="utf-8")
(tmp_path / "outside.txt").write_text("nope\n", encoding="utf-8")
tool = FindFilesTool(workspace=tmp_path, restrict_to_workspace=False)
token = bind_workspace_scope(default_workspace_scope(project, True))
try:
result = await tool.execute(path=".")
finally:
reset_workspace_scope(token)
assert result == "inside.txt"
@pytest.mark.asyncio
async def test_find_files_scan_keeps_event_loop_responsive(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
target = tmp_path / "match.txt"
target.write_text("ok\n", encoding="utf-8")
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
original_iter_paths = tool._iter_paths
started = threading.Event()
release = threading.Event()
def blocking_iter_paths(root: Path, *, include_dirs: bool, budget):
started.set()
if not release.wait(timeout=1):
raise TimeoutError("test did not release find_files traversal")
yield from original_iter_paths(root, include_dirs=include_dirs, budget=budget)
monkeypatch.setattr(tool, "_iter_paths", blocking_iter_paths)
task = asyncio.create_task(tool.execute(path="."))
try:
assert await asyncio.to_thread(started.wait, 0.5)
for _ in range(3):
await asyncio.sleep(0.01)
assert not task.done()
finally:
release.set()
assert await asyncio.wait_for(task, timeout=0.5) == "match.txt"
@pytest.mark.asyncio
async def test_find_files_cancellation_stops_worker_scan(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
started = threading.Event()
stopped = threading.Event()
def cancellable_iter_paths(root: Path, *, include_dirs: bool, budget):
del root, include_dirs
started.set()
if not budget.cancelled.wait(timeout=1):
raise TimeoutError("find_files worker did not receive cancellation")
stopped.set()
if False:
yield
monkeypatch.setattr(tool, "_iter_paths", cancellable_iter_paths)
task = asyncio.create_task(tool.execute(path="."))
assert await asyncio.to_thread(started.wait, 0.5)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(task, timeout=0.2)
assert await asyncio.to_thread(stopped.wait, 0.5)
@pytest.mark.asyncio
async def test_find_files_path_limit_stops_after_pagination_lookahead(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
paths = [tmp_path / name for name in ("a.txt", "b.txt", "c.txt")]
for path in paths:
path.write_text("ok\n", encoding="utf-8")
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
visited: list[str] = []
def ordered_iter_paths(root: Path, *, include_dirs: bool, budget):
del include_dirs
for path in paths:
budget.checkpoint()
visited.append(path.name)
yield tool._entry(path, root, is_dir=False)
monkeypatch.setattr(tool, "_iter_paths", ordered_iter_paths)
result = await tool.execute(path=".", head_limit=1)
assert result.splitlines()[0] == "a.txt"
assert "pagination: limit=1, offset=0" in result
assert visited == ["a.txt", "b.txt"]
@pytest.mark.asyncio
async def test_find_files_path_budget_counts_directories(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
(tmp_path / "one").mkdir()
(tmp_path / "two").mkdir()
monkeypatch.setattr(FindFilesTool, "_MAX_SCAN_PATHS", 1)
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
result = await tool.execute(path=".")
assert result.startswith("Error: find_files scan exceeded 1 paths")
@pytest.mark.asyncio
async def test_find_files_enforces_time_budget(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
(tmp_path / "match.txt").write_text("ok\n", encoding="utf-8")
monkeypatch.setattr(FindFilesTool, "_MAX_SCAN_SECONDS", 0.0)
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
result = await tool.execute(path=".")
assert result.startswith("Error: find_files scan exceeded 0 seconds")
@pytest.mark.asyncio
async def test_find_files_path_sort_matches_existing_lexicographic_contract(
tmp_path: Path,
) -> None:
(tmp_path / "a").mkdir()
(tmp_path / "a" / "inside.txt").write_text("ok\n", encoding="utf-8")
(tmp_path / "a+").write_text("ok\n", encoding="utf-8")
(tmp_path / "a.py").write_text("ok\n", encoding="utf-8")
tool = FindFilesTool(workspace=tmp_path, allowed_dir=tmp_path)
result = await tool.execute(path=".", head_limit=0)
assert result.splitlines() == ["a+", "a.py", "a/inside.txt"]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_grep_respects_glob_filter_and_context(tmp_path: Path) -> None: async def test_grep_respects_glob_filter_and_context(tmp_path: Path) -> None:
(tmp_path / "src").mkdir() (tmp_path / "src").mkdir()