From ddae612ffd9cc48b4c7f64e50fe790be4f1ff192 Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 13:43:10 -0700 Subject: [PATCH 1/8] feat: add Workspace file tools and WorkspaceLLMTaskWorker for jailed file access Adds planai.tools.filesystem (Workspace, make_file_tools, hash_files) so an LLM can read/write/edit/list/grep files inside a sandboxed per-job directory, and planai.workspace_task.WorkspaceLLMTaskWorker/WorkspaceTask to wire those tools into a CachedLLMTaskWorker whose working directory is discovered from task provenance. LLMTaskWorker gains get_tools()/get_cache_salt() hooks and a max_tool_rounds field (forwarded to generate_pydantic only when tools are in use), and CachedTaskWorker gains a _cache_hit_is_valid() hook so a cache hit missing its expected output files is treated as a miss and re-executed. Bumps version to 0.7.0 and documents the new worker in usage.rst. Co-Authored-By: Claude Fable 5.1 --- docs/source/usage.rst | 50 ++++ pyproject.toml | 2 +- src/planai/__init__.py | 7 + src/planai/_version.py | 2 +- src/planai/cached_task.py | 29 ++ src/planai/dispatcher.py | 1 - src/planai/llm_task.py | 55 +++- src/planai/provenance.py | 1 - src/planai/tools/__init__.py | 18 ++ src/planai/tools/filesystem.py | 384 ++++++++++++++++++++++++++ src/planai/workspace_task.py | 152 ++++++++++ tests/planai/test_cached_task.py | 26 ++ tests/planai/test_llm_task.py | 99 +++++++ tests/planai/test_utils.py | 2 +- tests/planai/test_workspace_task.py | 218 +++++++++++++++ tests/planai/tools/test_filesystem.py | 318 +++++++++++++++++++++ 16 files changed, 1358 insertions(+), 6 deletions(-) create mode 100644 src/planai/tools/__init__.py create mode 100644 src/planai/tools/filesystem.py create mode 100644 src/planai/workspace_task.py create mode 100644 tests/planai/test_workspace_task.py create mode 100644 tests/planai/tools/test_filesystem.py diff --git a/docs/source/usage.rst b/docs/source/usage.rst index aefb87d..035fb18 100644 --- a/docs/source/usage.rst +++ b/docs/source/usage.rst @@ -209,6 +209,56 @@ Example: initial_input = Task1WorkItem(data="start") main_graph.run(initial_tasks=[(subgraph_worker, initial_input)]) +Letting an LLM Work with Files +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +``WorkspaceLLMTaskWorker`` (in ``planai.workspace_task``) is a ``CachedLLMTaskWorker`` that +gives the LLM a set of file tools -- ``read_file``, ``write_file``, ``edit_file``, +``list_files``, and ``grep_files`` -- jailed to a per-job working directory. This is useful +when many jobs run concurrently, each operating on its own directory, and you want an LLM to +read, write, or search files without risking access outside of the directory assigned to it. + +The working directory is discovered automatically from the task's provenance chain: include a +``WorkspaceTask`` (or any ``Task`` subclass with a string ``workspace`` attribute) upstream, and +``WorkspaceLLMTaskWorker.get_workspace()`` will find the nearest one. All file paths the LLM +uses are relative to that directory; absolute paths, ``..`` components, and symlinks that +escape the directory are rejected. + +.. code-block:: python + + from planai import Graph, WorkspaceLLMTaskWorker, WorkspaceTask, llm_from_config + + class CodeReviewer(WorkspaceLLMTaskWorker): + prompt = "Review the code in this repository and write your findings to review.md." + output_types = [ReviewResult] + + def expected_output_files(self, task): + # a cache hit whose files are missing (e.g. a fresh checkout) is re-executed + return ["review.md"] + + llm = llm_from_config(provider="openai", model_name="gpt-4o") + reviewer = CodeReviewer(llm=llm, max_tool_rounds=40) + + graph = Graph(name="Review Workflow") + graph.add_workers(reviewer) + graph.set_entry(reviewer) + graph.set_exit(reviewer) + graph.run( + initial_tasks=[(reviewer, WorkspaceTask(workspace="/jobs/job-123"))] + ) + +Because a cache hit replays only the published output tasks and not any files the tools wrote +on a previous run, override ``expected_output_files()`` to list the workspace-relative paths +the worker is expected to produce; if any are missing, the cache entry is treated as a miss and +the worker re-executes. Set ``input_globs`` to fold the content of matching workspace files +into the cache key so that changing an input file also invalidates the cache. Pass +``read_only=True`` to omit the ``write_file`` and ``edit_file`` tools. + +The underlying pieces are reusable on their own: ``Workspace`` is the sandboxing primitive, +``make_file_tools(workspace, ...)`` builds the ``llm_interface`` ``Tool`` objects for any +``LLMTaskWorker`` (via the ``get_tools()`` hook), and ``hash_files(workspace, globs)`` computes +a stable content hash for a set of glob patterns. + Best Practices -------------- diff --git a/pyproject.toml b/pyproject.toml index ba18e8b..31523b1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "planai" -version = "0.6.1" +version = "0.7.0" description = "A simple framework for coordinating classical compute and LLM-based tasks." authors = ["Niels Provos "] license = "Apache-2.0" diff --git a/src/planai/__init__.py b/src/planai/__init__.py index 172baad..66f51be 100644 --- a/src/planai/__init__.py +++ b/src/planai/__init__.py @@ -25,7 +25,9 @@ from .provenance import ProvenanceChain from .pydantic_dict_wrapper import PydanticDictWrapper from .task import Task, TaskWorker +from .tools import Workspace, hash_files, make_file_tools from .user_input import UserInputRequest +from .workspace_task import WorkspaceLLMTaskWorker, WorkspaceTask # Limit what gets imported with "from planai import *" __all__ = [ @@ -50,4 +52,9 @@ "UserInputRequest", "tool", "Tool", + "Workspace", + "make_file_tools", + "hash_files", + "WorkspaceLLMTaskWorker", + "WorkspaceTask", ] diff --git a/src/planai/_version.py b/src/planai/_version.py index dc89a63..f4e9355 100644 --- a/src/planai/_version.py +++ b/src/planai/_version.py @@ -14,4 +14,4 @@ """Version information for PlanAI.""" -__version__ = "0.6.1" +__version__ = "0.7.0" diff --git a/src/planai/cached_task.py b/src/planai/cached_task.py index 6ac432b..159b469 100644 --- a/src/planai/cached_task.py +++ b/src/planai/cached_task.py @@ -50,8 +50,19 @@ def _pre_consume_work(self, task: Task): logging.error("Error getting data from cache %s: %s", cache_key, str(e)) result = None + cache_hit = False + cached_results = None if result is not None: cached_results, _ = result + cache_hit = self._cache_hit_is_valid(task, cached_results) + if not cache_hit: + logging.info( + "Cache hit for %s with key: %s is no longer valid; re-executing", + self.name, + cache_key, + ) + + if cache_hit: logging.info("Cache hit for %s with key: %s", self.name, cache_key) self._publish_cached_results(cached_results, task) else: @@ -66,6 +77,24 @@ def _pre_consume_work(self, task: Task): self.post_consume_work(task) + def _cache_hit_is_valid( + self, task: Task, cached_results: List[Tuple[str, Task]] + ) -> bool: + """ + Hook for subclasses to reject a cache hit based on external state that isn't + captured by the cache key itself (e.g. output files on disk that a previous + run's cache entry references but that are no longer present). Defaults to + always accepting the cache hit. + + Args: + task (Task): The input task that produced a cache hit. + cached_results (List[Tuple[str, Task]]): The cached (consumer_name, output_task) pairs. + + Returns: + bool: True if the cached results should be used, False to treat this as a cache miss. + """ + return True + def pre_consume_work(self, task: Task): """ This method is called before consuming the work item. It will be called even if the task has been cached. diff --git a/src/planai/dispatcher.py b/src/planai/dispatcher.py index a53edd6..7793388 100644 --- a/src/planai/dispatcher.py +++ b/src/planai/dispatcher.py @@ -54,7 +54,6 @@ across various workers effectively. """ - import logging import random import threading diff --git a/src/planai/llm_task.py b/src/planai/llm_task.py index 82941b2..2387407 100644 --- a/src/planai/llm_task.py +++ b/src/planai/llm_task.py @@ -85,6 +85,13 @@ class LLMTaskWorker(BaseLLMTaskWorker): default=False, description="Whether to use XML format for the data input to the LLM", ) + max_tool_rounds: Optional[int] = Field( + default=None, + description=( + "Maximum number of tool-calling rounds to allow the LLM. Only forwarded to " + "generate_pydantic when tools are actually in use for a given call." + ), + ) def __init__(self, **data): super().__init__(**data) @@ -124,6 +131,38 @@ def _format_task(self, task: Task) -> str: else task.model_dump_xml() ) + def get_tools(self, task: Task) -> Optional[List[Tool]]: + """ + Returns the tools that should be made available to the LLM for this task. + Defaults to the static ``tools`` field, but subclasses can override this to + provide dynamic, per-task tools (e.g. file tools bound to a task-specific + working directory). + + Args: + task (Task): The input task. + + Returns: + Optional[List[Tool]]: The tools to make available, or None/empty for no tools. + """ + return self.tools + + def get_cache_salt(self, task: Task) -> Optional[str]: + """ + Returns an optional salt to forward as ``cache_salt`` to generate_pydantic. + Only used by _invoke_llm when it returns a non-None value (and tools are in + use), so the default LLMTaskWorker behavior of not passing cache_salt is + unchanged. Subclasses (e.g. WorkspaceLLMTaskWorker) can override this to + invalidate the LLM's own response cache when external state, such as + workspace files, changes. + + Args: + task (Task): The input task. + + Returns: + Optional[str]: The cache salt, or None to omit it. + """ + return None + def _invoke_llm(self, task: Task): # allow subclasses to customize the prompt based on the input task task_prompt = self.format_prompt(task) @@ -141,6 +180,19 @@ def extra_validation_with_task(response: BaseModel): assert isinstance(response, Task) return self.extra_validation(response, task) + tools = self.get_tools(task) + + # only forward these new, tool-related kwargs when tools are actually in + # use for this call, so that callers/mocks that don't expect them (and + # don't use tools) keep working unchanged. + extra_kwargs: Dict[str, Any] = {} + if tools: + if self.max_tool_rounds is not None: + extra_kwargs["max_tool_rounds"] = self.max_tool_rounds + cache_salt = self.get_cache_salt(task) + if cache_salt is not None: + extra_kwargs["cache_salt"] = cache_salt + response = self.llm.generate_pydantic( prompt_template=( (PROMPT_TEMPLATE if processed_task is not None else "{instructions}") @@ -152,7 +204,7 @@ def extra_validation_with_task(response: BaseModel): ), output_schema=self._output_type(), system=self.system_prompt, - tools=self.tools if self.tools else None, + tools=tools if tools else None, task=self._format_task(processed_task), temperature=self.temperature, instructions=task_prompt, @@ -162,6 +214,7 @@ def extra_validation_with_task(response: BaseModel): debug_saver=save_debug_with_task if self.debug_mode else None, extra_validation=extra_validation_with_task, images=task.images if isinstance(task, MediaTask) else None, + **extra_kwargs, ) assert isinstance(response, Task) or response is None self.post_process(response=response, input_task=task) diff --git a/src/planai/provenance.py b/src/planai/provenance.py index 05ff4d1..7f0e599 100644 --- a/src/planai/provenance.py +++ b/src/planai/provenance.py @@ -40,7 +40,6 @@ management of task dependencies and ordered task execution. """ - import logging import sys from collections import defaultdict diff --git a/src/planai/tools/__init__.py b/src/planai/tools/__init__.py new file mode 100644 index 0000000..b10ce20 --- /dev/null +++ b/src/planai/tools/__init__.py @@ -0,0 +1,18 @@ +# Copyright 2024 Niels Provos +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tools that can be handed to LLMTaskWorker subclasses to let an LLM act on files.""" + +from .filesystem import Workspace, hash_files, make_file_tools + +__all__ = ["Workspace", "make_file_tools", "hash_files"] diff --git a/src/planai/tools/filesystem.py b/src/planai/tools/filesystem.py new file mode 100644 index 0000000..518f48e --- /dev/null +++ b/src/planai/tools/filesystem.py @@ -0,0 +1,384 @@ +# Copyright 2024 Niels Provos +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""File system tools that let an LLM read, write, and search files inside a +sandboxed per-job working directory (a :class:`Workspace`). + +These are meant to be handed to :class:`planai.llm_task.LLMTaskWorker` (or more +commonly :class:`planai.workspace_task.WorkspaceLLMTaskWorker`) as ``tools`` so +that many concurrent jobs, each with their own working directory, can safely let +an LLM operate on files without risking access outside of the assigned directory. +""" + +import hashlib +import re +from pathlib import Path +from typing import List, Union + +from llm_interface import Tool +from llm_interface.llm_tool import create_tool + +__all__ = ["Workspace", "make_file_tools", "hash_files"] + + +class Workspace: + """A sandboxed working directory that file tools are jailed to. + + All paths handed to :meth:`resolve` are interpreted as relative to the + workspace root and are rejected with a :class:`ValueError` if they would + escape that root, whether through an absolute path, a ``..`` component, or + a symlink that points outside of the root. + """ + + def __init__(self, root: Union[str, "Path"]): + root_path = Path(root).expanduser() + root_path.mkdir(parents=True, exist_ok=True) + # resolve() follows symlinks and normalizes the path so that every + # subsequent comparison against self.root is done against the real path. + self.root: Path = root_path.resolve() + + def resolve(self, rel_path: str) -> Path: + """Resolve a workspace-relative path to an absolute path inside the root. + + Args: + rel_path: A path relative to the workspace root. + + Returns: + Path: The resolved absolute path, guaranteed to be inside the workspace root. + + Raises: + ValueError: If the path is empty, absolute, contains a ``..`` component, + or resolves (following symlinks) to a location outside the workspace root. + """ + if not rel_path: + raise ValueError("Path must not be empty") + + candidate = Path(rel_path) + if candidate.is_absolute(): + raise ValueError(f"Absolute paths are not allowed: {rel_path!r}") + if ".." in candidate.parts: + raise ValueError(f"Path must not contain '..': {rel_path!r}") + + try: + # strict=False resolves symlinks and normalizes the path as far as + # possible even if the final component does not yet exist (e.g. a + # file that write_file is about to create). + resolved = (self.root / candidate).resolve(strict=False) + except OSError as e: + raise ValueError(f"Could not resolve path: {rel_path!r}") from e + + try: + resolved.relative_to(self.root) + except ValueError: + raise ValueError(f"Path escapes the workspace root: {rel_path!r}") from None + + return resolved + + +def hash_files(workspace: Union[Workspace, str, Path], globs: List[str]) -> str: + """Compute a stable hash over the content of files matching the given globs. + + The hash is a sha1 digest over the sorted (relative path, content) pairs of every + file matched by any of the glob patterns, so it changes whenever a matched file's + content (or the set of matched files) changes, and is stable across runs otherwise. + + Args: + workspace: The Workspace (or a path to one) whose files to hash. + globs: A list of glob patterns (e.g. ``["**/*.py"]``), evaluated relative to + the workspace root. + + Returns: + str: A hex sha1 digest, or an empty string if no files match any pattern. + """ + ws = workspace if isinstance(workspace, Workspace) else Workspace(workspace) + + matched = {} + for pattern in globs: + for path in ws.root.glob(pattern): + if path.is_file(): + rel = path.relative_to(ws.root).as_posix() + matched[rel] = path + + if not matched: + return "" + + digest = hashlib.sha1() + for rel in sorted(matched.keys()): + digest.update(rel.encode("utf-8")) + digest.update(b"\x00") + digest.update(matched[rel].read_bytes()) + digest.update(b"\x00") + return digest.hexdigest() + + +def make_file_tools( + workspace: Union[Workspace, str, Path], + *, + read_only: bool = False, + max_read_chars: int = 100_000, + max_list_entries: int = 500, + max_grep_matches: int = 200, +) -> List[Tool]: + """Create llm_interface Tool objects bound to a single sandboxed workspace. + + Args: + workspace: The Workspace to jail all file operations to (or a path to one, + which will be created if missing). + read_only: If True, omit the write_file and edit_file tools. + max_read_chars: Maximum number of characters read_file returns before truncating. + max_list_entries: Maximum number of entries list_files returns before truncating. + max_grep_matches: Maximum number of matches grep_files returns before truncating. + + Returns: + List[Tool]: Tool objects usable as the ``tools`` argument to an LLMTaskWorker. + """ + ws = workspace if isinstance(workspace, Workspace) else Workspace(workspace) + + def read_file(path: str, offset: int = 0, limit: int = 0) -> str: + """Read the text content of a file in the workspace, with line numbers. + + Returns the file's content formatted like ``cat -n``, i.e. each line is + prefixed with its 1-based line number. Use offset and limit to read a + large file in smaller windows. + + Args: + path: Workspace-relative path to the file to read. + offset: 1-based line number to start reading from. 0 or 1 both mean + start at the first line of the file. + limit: Maximum number of lines to return. 0 means return every line + from offset to the end of the file (subject to the character limit). + """ + try: + target = ws.resolve(path) + except ValueError as e: + return f"Error: {e}" + + if not target.exists(): + return f"Error: File not found: {path}" + if not target.is_file(): + return f"Error: Not a file: {path}" + + try: + text = target.read_text(encoding="utf-8") + except UnicodeDecodeError: + return f"Error: File is not valid UTF-8 text (binary?): {path}" + except OSError as e: + return f"Error: Could not read file {path}: {e.strerror or e}" + + lines = text.splitlines() + start = max(offset - 1, 0) if offset > 1 else 0 + end = start + limit if limit > 0 else len(lines) + selected = lines[start:end] + + numbered = "\n".join( + f"{start + i + 1:6d}\t{line}" for i, line in enumerate(selected) + ) + + if len(numbered) > max_read_chars: + numbered = ( + numbered[:max_read_chars] + + f"\n\n[Output truncated at {max_read_chars} characters. " + "Use the offset and limit parameters to read this file in smaller windows.]" + ) + return numbered + + def write_file(path: str, content: str) -> str: + """Write text content to a file in the workspace, overwriting it if it exists. + + Parent directories are created automatically as needed. + + Args: + path: Workspace-relative path to the file to write. + content: The full text content to write to the file. + """ + try: + target = ws.resolve(path) + except ValueError as e: + return f"Error: {e}" + + try: + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(content, encoding="utf-8") + except OSError as e: + return f"Error: Could not write file {path}: {e.strerror or e}" + + byte_count = len(content.encode("utf-8")) + line_count = len(content.splitlines()) + return f"Wrote {byte_count} bytes ({line_count} lines) to {path}" + + def edit_file( + path: str, old_string: str, new_string: str, replace_all: bool = False + ) -> str: + """Replace an exact snippet of text within an existing file. + + By default, old_string must occur exactly once in the file; use replace_all + to replace every occurrence instead. + + Args: + path: Workspace-relative path to the file to edit. + old_string: The exact, literal text to find in the file. Must be unique + in the file unless replace_all is true. + new_string: The text to replace old_string with. + replace_all: When true, replace every occurrence of old_string instead + of requiring exactly one match. + """ + try: + target = ws.resolve(path) + except ValueError as e: + return f"Error: {e}" + + if not target.exists(): + return f"Error: File not found: {path}" + if not target.is_file(): + return f"Error: Not a file: {path}" + if not old_string: + return "Error: old_string must not be empty" + + try: + text = target.read_text(encoding="utf-8") + except UnicodeDecodeError: + return f"Error: File is not valid UTF-8 text (binary?): {path}" + except OSError as e: + return f"Error: Could not read file {path}: {e.strerror or e}" + + count = text.count(old_string) + if count == 0: + return f"Error: old_string not found in {path} (0 matches)" + if count > 1 and not replace_all: + return ( + f"Error: old_string is not unique in {path} ({count} matches); " + "pass replace_all=true or include more surrounding context" + ) + + if replace_all: + new_text = text.replace(old_string, new_string) + replacements = count + else: + new_text = text.replace(old_string, new_string, 1) + replacements = 1 + + try: + target.write_text(new_text, encoding="utf-8") + except OSError as e: + return f"Error: Could not write file {path}: {e.strerror or e}" + + return f"Replaced {replacements} occurrence(s) in {path}" + + def list_files(path: str = ".", pattern: str = "**/*") -> str: + """List files under a directory in the workspace matching a glob pattern. + + Only files are listed (not directories), sorted by relative path, each + shown with its size in bytes. + + Args: + path: Workspace-relative directory to list. Use "." for the workspace root. + pattern: Glob pattern, relative to path, selecting which files to include, + e.g. "**/*" for every file recursively or "*.py" for top-level Python files. + """ + try: + target = ws.resolve(path) + except ValueError as e: + return f"Error: {e}" + + if not target.exists(): + return f"Error: Directory not found: {path}" + if not target.is_dir(): + return f"Error: Not a directory: {path}" + + try: + matches = sorted(p for p in target.glob(pattern) if p.is_file()) + except (re.error, ValueError) as e: + return f"Error: Invalid pattern: {e}" + + truncated = len(matches) > max_list_entries + lines = [] + for p in matches[:max_list_entries]: + rel = p.relative_to(ws.root).as_posix() + try: + size = p.stat().st_size + except OSError: + size = 0 + lines.append(f"{size:>10} {rel}") + + if not lines: + return "(no files found)" + + output = "\n".join(lines) + if truncated: + output += f"\n\n[Output truncated at {max_list_entries} entries.]" + return output + + def grep_files(pattern: str, path: str = ".", glob: str = "**/*.md") -> str: + """Search for a regular expression across text files in the workspace. + + Args: + pattern: Python regular expression to search for within each line of + each matching file. + path: Workspace-relative directory to search within. Use "." for the + workspace root. + glob: Glob pattern, relative to path, selecting which files to search, + e.g. "**/*.md" for every Markdown file recursively. + """ + try: + target = ws.resolve(path) + except ValueError as e: + return f"Error: {e}" + + if not target.exists(): + return f"Error: Directory not found: {path}" + if not target.is_dir(): + return f"Error: Not a directory: {path}" + + try: + regex = re.compile(pattern) + except re.error as e: + return f"Error: Invalid regular expression: {e}" + + try: + candidates = sorted(p for p in target.glob(glob) if p.is_file()) + except ValueError as e: + return f"Error: Invalid pattern: {e}" + + results: List[str] = [] + truncated = False + for file_path in candidates: + try: + text = file_path.read_text(encoding="utf-8") + except (UnicodeDecodeError, OSError): + continue + + rel = file_path.relative_to(ws.root).as_posix() + for lineno, line in enumerate(text.splitlines(), start=1): + if regex.search(line): + results.append(f"{rel}:{lineno}: {line}") + if len(results) >= max_grep_matches: + truncated = True + break + if truncated: + break + + if not results: + return "(no matches found)" + + output = "\n".join(results) + if truncated: + output += f"\n\n[Output truncated at {max_grep_matches} matches.]" + return output + + tools = [create_tool(read_file)] + if not read_only: + tools.append(create_tool(write_file)) + tools.append(create_tool(edit_file)) + tools.append(create_tool(list_files)) + tools.append(create_tool(grep_files)) + return tools diff --git a/src/planai/workspace_task.py b/src/planai/workspace_task.py new file mode 100644 index 0000000..7baf289 --- /dev/null +++ b/src/planai/workspace_task.py @@ -0,0 +1,152 @@ +# Copyright 2024 Niels Provos +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""A CachedLLMTaskWorker that lets an LLM read, write, and search files inside a +per-job working directory carried through task provenance.""" + +from typing import List, Optional, Tuple + +from pydantic import Field + +from .llm_task import CachedLLMTaskWorker +from .task import Task +from .tools.filesystem import Workspace, hash_files, make_file_tools + + +class WorkspaceTask(Task): + """A Task that carries the absolute path of a per-job working directory. + + Graphs that want their LLM workers to operate on files should include a + WorkspaceTask (or a Task subclass with a ``workspace: str`` attribute) + somewhere in the provenance chain; WorkspaceLLMTaskWorker.get_workspace() + will find it automatically. + """ + + workspace: str = Field( + ..., description="Absolute path to the per-job working directory." + ) + + +class WorkspaceLLMTaskWorker(CachedLLMTaskWorker): + """A CachedLLMTaskWorker that gives the LLM file tools jailed to a workspace. + + The workspace directory is discovered from the input task's provenance chain + (see get_workspace()) unless overridden. Because a cache hit replays only the + published output tasks and not any files the tools wrote on a prior run, + subclasses can declare expected_output_files() so that a cache hit whose + files are missing is treated as a cache miss and re-executed. + """ + + read_only: bool = Field( + default=False, + description="If true, only read/list/grep tools are exposed; write_file and edit_file are omitted.", + ) + max_tool_rounds: int = Field( + default=40, + description="Maximum number of tool-calling rounds to allow the LLM.", + ) + input_globs: List[str] = Field( + default_factory=list, + description=( + "Glob patterns (relative to the workspace root) whose matched files' " + "content is folded into the cache key, so the cache is invalidated " + "when those input files change." + ), + ) + max_read_chars: int = Field( + default=100_000, + description="Maximum number of characters the read_file tool returns before truncating.", + ) + + def get_workspace(self, task: Task) -> Workspace: + """ + Finds the Workspace for this task by looking for the nearest task in the + provenance chain -- the task itself first, then its input provenance, + walked nearest-first the same way Task.find_input_task() does -- that has + a string ``workspace`` attribute. Subclasses may override this, e.g. to + pull the directory from worker configuration instead. + + Args: + task (Task): The input task. + + Returns: + Workspace: The workspace for this task. + + Raises: + ValueError: If no task with a string ``workspace`` attribute is found. + """ + candidates = [task] + list(reversed(task._input_provenance)) + for candidate in candidates: + workspace_path = getattr(candidate, "workspace", None) + if isinstance(workspace_path, str) and workspace_path: + return Workspace(workspace_path) + + raise ValueError( + f"{self.name}: could not find a workspace directory in the task or " + "its input provenance; expected a WorkspaceTask (or a Task with a " + "string 'workspace' attribute) upstream" + ) + + def get_tools(self, task: Task): + file_tools = make_file_tools( + self.get_workspace(task), + read_only=self.read_only, + max_read_chars=self.max_read_chars, + ) + if self.tools: + return file_tools + list(self.tools) + return file_tools + + def get_cache_salt(self, task: Task) -> Optional[str]: + return self._get_cache_key(task) + + def extra_cache_key(self, task: Task) -> str: + workspace = self.get_workspace(task) + return hash_files(workspace, self.input_globs) + + def expected_output_files(self, task: Task) -> List[str]: + """ + Workspace-relative file paths this worker is expected to have written when + it runs. Used to bypass a cache hit whose files are no longer present. + Defaults to no expectations (any cache hit is honored). Subclasses that + have the LLM write files via the file tools should override this. + + Args: + task (Task): The input task. + + Returns: + List[str]: Workspace-relative paths expected to exist after this worker runs. + """ + return [] + + def _cache_hit_is_valid( + self, task: Task, cached_results: List[Tuple[str, Task]] + ) -> bool: + expected_files = self.expected_output_files(task) + if not expected_files: + return True + + try: + workspace = self.get_workspace(task) + except ValueError: + # nothing we can check against; fall back to the default behavior + return True + + for rel_path in expected_files: + try: + resolved = workspace.resolve(rel_path) + except ValueError: + return False + if not resolved.exists(): + return False + return True diff --git a/tests/planai/test_cached_task.py b/tests/planai/test_cached_task.py index 968c8ee..64dae19 100644 --- a/tests/planai/test_cached_task.py +++ b/tests/planai/test_cached_task.py @@ -168,6 +168,32 @@ def test_single_consumer_different_cached_consumer_name(self): self.worker._pre_consume_work(task) mock_get_consumer.assert_called_once_with(cached_result[0][1]) + def test_cache_hit_is_valid_default_true(self): + task = DummyInputTask(data="test") + cached_result = [("SinkTaskWorker", DummyOutputTask(processed_data="x"))] + self.assertTrue(self.worker._cache_hit_is_valid(task, cached_result)) + + def test_cache_hit_invalidated_by_hook_reexecutes(self): + task = DummyInputTask(data="test") + cached_result = [DummyOutputTask(processed_data="Cached: test")] + cache_key = self.worker._get_cache_key(task) + self.mock_cache.set(cache_key, [cached_result, task]) + self.mock_cache.clear_stats() + + with patch.object( + self.worker, "_cache_hit_is_valid", return_value=False + ) as mock_valid: + with patch( + "test_cached_task.DummyCachedTaskWorker.consume_work" + ) as mock_consume: + with patch.object( + self.worker, "_publish_cached_results" + ) as mock_publish: + self.worker._pre_consume_work(task) + mock_valid.assert_called_once_with(task, cached_result) + mock_consume.assert_called_once_with(task) + mock_publish.assert_not_called() + def test_two_consumers_invalid_consumer_name(self): second_sink_worker = SinkTaskWorker() self.worker.register_consumer(DummyOutputTask, second_sink_worker) diff --git a/tests/planai/test_llm_task.py b/tests/planai/test_llm_task.py index 0cbfbb2..70e296b 100644 --- a/tests/planai/test_llm_task.py +++ b/tests/planai/test_llm_task.py @@ -129,6 +129,105 @@ def test_tools_passed_to_llm_interface(self, mock_publish_work): task=output_task_payload, input_task=input_task ) + def test_get_tools_defaults_to_static_tools_field(self): + mock_tool = MagicMock(spec=LLMToolInstance) + worker = LLMTaskWorker( + llm=self.llm, + prompt="Test prompt", + output_types=[DummyOutputTask], + tools=[mock_tool], + ) + task = DummyTask(content="content") + self.assertEqual(worker.get_tools(task), [mock_tool]) + + @patch("planai.llm_task.LLMTaskWorker.publish_work") + def test_invoke_llm_uses_get_tools_hook(self, mock_publish_work): + """_invoke_llm must call self.get_tools(task) rather than reading self.tools directly.""" + mock_tool = MagicMock(spec=LLMToolInstance) + output_task_payload = DummyOutputTask(result="Tool test output") + input_task = DummyTask(content="Tool test input") + + with patch.object( + self.llm, "generate_pydantic", return_value=output_task_payload + ) as mock_generate_pydantic: + with patch( + "planai.llm_task.LLMTaskWorker.get_tools", return_value=[mock_tool] + ) as mock_get_tools: + self.worker._invoke_llm(input_task) + + mock_get_tools.assert_called_once_with(input_task) + call_args = mock_generate_pydantic.call_args + self.assertEqual(call_args.kwargs["tools"], [mock_tool]) + + @patch("planai.llm_task.LLMTaskWorker.publish_work") + def test_max_tool_rounds_forwarded_only_with_tools(self, mock_publish_work): + mock_tool = MagicMock(spec=LLMToolInstance) + output_task_payload = DummyOutputTask(result="Tool test output") + input_task = DummyTask(content="Tool test input") + + worker_with_tools = LLMTaskWorker( + llm=self.llm, + prompt="Test prompt", + output_types=[DummyOutputTask], + tools=[mock_tool], + max_tool_rounds=7, + ) + + with patch.object( + self.llm, "generate_pydantic", return_value=output_task_payload + ) as mock_generate_pydantic: + worker_with_tools._invoke_llm(input_task) + call_args = mock_generate_pydantic.call_args + self.assertEqual(call_args.kwargs["max_tool_rounds"], 7) + + # Without tools, max_tool_rounds must not be forwarded even if set. + worker_without_tools = LLMTaskWorker( + llm=self.llm, + prompt="Test prompt", + output_types=[DummyOutputTask], + tools=None, + max_tool_rounds=7, + ) + with patch.object( + self.llm, "generate_pydantic", return_value=output_task_payload + ) as mock_generate_pydantic: + worker_without_tools._invoke_llm(input_task) + call_args = mock_generate_pydantic.call_args + self.assertNotIn("max_tool_rounds", call_args.kwargs) + + @patch("planai.llm_task.LLMTaskWorker.publish_work") + def test_cache_salt_forwarded_only_when_hook_returns_value(self, mock_publish_work): + mock_tool = MagicMock(spec=LLMToolInstance) + output_task_payload = DummyOutputTask(result="Tool test output") + input_task = DummyTask(content="Tool test input") + + worker_with_tools = LLMTaskWorker( + llm=self.llm, + prompt="Test prompt", + output_types=[DummyOutputTask], + tools=[mock_tool], + ) + + # default get_cache_salt returns None -> no cache_salt kwarg + with patch.object( + self.llm, "generate_pydantic", return_value=output_task_payload + ) as mock_generate_pydantic: + worker_with_tools._invoke_llm(input_task) + self.assertNotIn("cache_salt", mock_generate_pydantic.call_args.kwargs) + + # overridden get_cache_salt returning a value -> forwarded + with patch( + "planai.llm_task.LLMTaskWorker.get_cache_salt", return_value="some-salt" + ): + with patch.object( + self.llm, "generate_pydantic", return_value=output_task_payload + ) as mock_generate_pydantic: + worker_with_tools._invoke_llm(input_task) + self.assertEqual( + mock_generate_pydantic.call_args.kwargs["cache_salt"], + "some-salt", + ) + def test_invoke_llm(self): input_task = DummyTask(content="Test input") output_task = DummyOutputTask(result="Test output") diff --git a/tests/planai/test_utils.py b/tests/planai/test_utils.py index 83526e3..9a3b0ee 100644 --- a/tests/planai/test_utils.py +++ b/tests/planai/test_utils.py @@ -82,7 +82,7 @@ def _generate_random_string(self, length: int) -> str: chars = ( string.printable + "".join(chr(i) for i in range(0x80, 0x110000, 997)) # sparse unicode - + "\x00\x01\x02\x03\x1F" # control chars + + "\x00\x01\x02\x03\x1f" # control chars + "<>\"'&" # XML special chars + "🌟🔥🌈" # emojis ) diff --git a/tests/planai/test_workspace_task.py b/tests/planai/test_workspace_task.py new file mode 100644 index 0000000..b64500f --- /dev/null +++ b/tests/planai/test_workspace_task.py @@ -0,0 +1,218 @@ +# test_workspace_task.py + +import tempfile +import unittest +from pathlib import Path +from typing import List, Type +from unittest.mock import Mock, patch + +from llm_interface import LLMInterface +from llm_interface.llm_tool import Tool as LLMToolInstance +from planai.task import Task +from planai.testing.helpers import MockCache, add_input_provenance +from planai.workspace_task import WorkspaceLLMTaskWorker, WorkspaceTask + + +class DummyTask(Task): + content: str + + +class OutputTask(Task): + result: str + + +class DummyWorkspaceWorker(WorkspaceLLMTaskWorker): + output_types: List[Type[Task]] = [OutputTask] + + +class WorkspaceWorkerTestCase(unittest.TestCase): + """Shared setup for WorkspaceLLMTaskWorker tests that don't need a real cache.""" + + def setUp(self): + self.llm = LLMInterface() + self.llm.client = Mock() + self.workspace_dir = tempfile.TemporaryDirectory() + self.addCleanup(self.workspace_dir.cleanup) + self.cache_dir = tempfile.TemporaryDirectory() + self.addCleanup(self.cache_dir.cleanup) + self.worker = DummyWorkspaceWorker( + llm=self.llm, prompt="test prompt", cache_dir=self.cache_dir.name + ) + + +class TestWorkspaceLLMTaskWorkerDefaults(WorkspaceWorkerTestCase): + def test_max_tool_rounds_overridden_to_40(self): + self.assertEqual(self.worker.max_tool_rounds, 40) + + def test_read_only_defaults_false(self): + self.assertFalse(self.worker.read_only) + + def test_input_globs_defaults_empty(self): + self.assertEqual(self.worker.input_globs, []) + + def test_expected_output_files_defaults_empty(self): + task = WorkspaceTask(workspace=self.workspace_dir.name) + self.assertEqual(self.worker.expected_output_files(task), []) + + +class TestGetWorkspace(WorkspaceWorkerTestCase): + def test_finds_workspace_from_task_itself(self): + task = WorkspaceTask(workspace=self.workspace_dir.name) + ws = self.worker.get_workspace(task) + self.assertEqual(ws.root, Path(self.workspace_dir.name).resolve()) + + def test_finds_workspace_from_input_provenance(self): + task = DummyTask(content="hi") + add_input_provenance(task, WorkspaceTask(workspace=self.workspace_dir.name)) + ws = self.worker.get_workspace(task) + self.assertEqual(ws.root, Path(self.workspace_dir.name).resolve()) + + def test_finds_nearest_workspace_when_multiple_in_chain(self): + older_dir = tempfile.TemporaryDirectory() + self.addCleanup(older_dir.cleanup) + task = DummyTask(content="hi") + add_input_provenance(task, WorkspaceTask(workspace=older_dir.name)) + add_input_provenance(task, WorkspaceTask(workspace=self.workspace_dir.name)) + ws = self.worker.get_workspace(task) + self.assertEqual(ws.root, Path(self.workspace_dir.name).resolve()) + + def test_raises_value_error_when_not_found(self): + task = DummyTask(content="hi") + with self.assertRaises(ValueError): + self.worker.get_workspace(task) + + +class TestGetTools(WorkspaceWorkerTestCase): + def setUp(self): + super().setUp() + self.task = WorkspaceTask(workspace=self.workspace_dir.name) + + def test_tools_are_bound_to_the_task_workspace(self): + tools = {t.name: t for t in self.worker.get_tools(self.task)} + self.assertIn("write_file", tools) + result = tools["write_file"].execute(path="hello.txt", content="hi there") + self.assertNotIn("Error", result) + self.assertEqual( + (Path(self.workspace_dir.name) / "hello.txt").read_text(), "hi there" + ) + + def test_read_only_omits_write_and_edit_tools(self): + self.worker.read_only = True + tools = {t.name: t for t in self.worker.get_tools(self.task)} + self.assertNotIn("write_file", tools) + self.assertNotIn("edit_file", tools) + self.assertIn("read_file", tools) + + def test_static_tools_are_appended(self): + custom_tool = LLMToolInstance( + name="custom_tool", + description="A custom tool", + parameters={"type": "object", "properties": {}, "required": []}, + func=lambda: "ok", + ) + self.worker.tools = [custom_tool] + tools = self.worker.get_tools(self.task) + names = {t.name for t in tools} + self.assertIn("custom_tool", names) + # file tools should still be present alongside the static tool + self.assertIn("write_file", names) + + +class TestExtraCacheKeyAndCacheSalt(WorkspaceWorkerTestCase): + def setUp(self): + super().setUp() + self.worker = DummyWorkspaceWorker( + llm=self.llm, + prompt="test prompt", + cache_dir=self.cache_dir.name, + input_globs=["*.txt"], + ) + self.task = WorkspaceTask(workspace=self.workspace_dir.name) + + def test_cache_key_changes_when_input_file_content_changes(self): + data_file = Path(self.workspace_dir.name) / "data.txt" + data_file.write_text("v1") + key1 = self.worker._get_cache_key(self.task) + + data_file.write_text("v2") + key2 = self.worker._get_cache_key(self.task) + + self.assertNotEqual(key1, key2) + + def test_cache_key_stable_without_changes(self): + data_file = Path(self.workspace_dir.name) / "data.txt" + data_file.write_text("v1") + key1 = self.worker._get_cache_key(self.task) + key2 = self.worker._get_cache_key(self.task) + self.assertEqual(key1, key2) + + def test_get_cache_salt_matches_cache_key(self): + self.assertEqual( + self.worker.get_cache_salt(self.task), + self.worker._get_cache_key(self.task), + ) + + +class OutputTaskFile(Task): + result: str + + +class FileWritingWorker(WorkspaceLLMTaskWorker): + output_types: List[Type[Task]] = [OutputTaskFile] + + def expected_output_files(self, task: Task) -> List[str]: + return ["out.txt"] + + +class TestCacheHitBypass(unittest.TestCase): + def setUp(self): + self.llm = LLMInterface() + self.llm.client = Mock() + self.workspace_dir = tempfile.TemporaryDirectory() + self.addCleanup(self.workspace_dir.cleanup) + + self.mock_cache = MockCache() + self.cache_patcher = patch( + "planai.cached_task.Cache", return_value=self.mock_cache + ) + self.cache_patcher.start() + self.addCleanup(self.cache_patcher.stop) + + self.worker = FileWritingWorker( + llm=self.llm, prompt="test prompt", cache_dir="./unused-cache-dir" + ) + self.task = WorkspaceTask(workspace=self.workspace_dir.name) + + def _seed_cache(self): + cache_key = self.worker._get_cache_key(self.task) + cached_result = [(None, OutputTaskFile(result="cached"))] + self.mock_cache.set(cache_key, [cached_result, self.task]) + self.mock_cache.clear_stats() + + def test_cache_hit_bypassed_when_expected_file_missing(self): + self._seed_cache() + fresh_output = OutputTaskFile(result="fresh") + self.llm.generate_pydantic = Mock(return_value=fresh_output) + + with patch.object(self.worker, "_publish_cached_results") as mock_publish: + with patch("planai.llm_task.LLMTaskWorker.publish_work"): + self.worker._pre_consume_work(self.task) + mock_publish.assert_not_called() + + self.llm.generate_pydantic.assert_called_once() + + def test_cache_hit_honored_when_expected_file_present(self): + self._seed_cache() + (Path(self.workspace_dir.name) / "out.txt").write_text("done") + + self.llm.generate_pydantic = Mock() + + with patch.object(self.worker, "_publish_cached_results") as mock_publish: + self.worker._pre_consume_work(self.task) + mock_publish.assert_called_once() + + self.llm.generate_pydantic.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/planai/tools/test_filesystem.py b/tests/planai/tools/test_filesystem.py new file mode 100644 index 0000000..7d61c19 --- /dev/null +++ b/tests/planai/tools/test_filesystem.py @@ -0,0 +1,318 @@ +# test_filesystem.py + +import tempfile +import unittest +from pathlib import Path + +from planai.tools.filesystem import Workspace, hash_files, make_file_tools + + +class TestWorkspaceJail(unittest.TestCase): + def setUp(self): + self.tempdir = tempfile.TemporaryDirectory() + self.addCleanup(self.tempdir.cleanup) + self.workspace = Workspace(self.tempdir.name) + + def test_root_is_created_and_resolved(self): + self.assertTrue(self.workspace.root.exists()) + self.assertTrue(self.workspace.root.is_absolute()) + + def test_nested_path_accepted(self): + resolved = self.workspace.resolve("a/b/c.txt") + self.assertEqual(resolved, self.workspace.root / "a" / "b" / "c.txt") + + def test_absolute_path_rejected(self): + with self.assertRaises(ValueError): + self.workspace.resolve("/etc/passwd") + + def test_empty_path_rejected(self): + with self.assertRaises(ValueError): + self.workspace.resolve("") + + def test_dotdot_escape_rejected(self): + with self.assertRaises(ValueError): + self.workspace.resolve("../escape.txt") + with self.assertRaises(ValueError): + self.workspace.resolve("sub/../../escape.txt") + + def test_symlink_escape_rejected(self): + outside = tempfile.TemporaryDirectory() + self.addCleanup(outside.cleanup) + link = Path(self.tempdir.name) / "evil" + link.symlink_to(outside.name) + with self.assertRaises(ValueError): + self.workspace.resolve("evil/secret.txt") + + def test_symlink_within_root_accepted(self): + target_dir = Path(self.tempdir.name) / "real" + target_dir.mkdir() + link = Path(self.tempdir.name) / "link" + link.symlink_to(target_dir) + resolved = self.workspace.resolve("link/file.txt") + self.assertEqual(resolved, target_dir.resolve() / "file.txt") + + def test_parent_dirs_created_on_write(self): + tools = {t.name: t for t in make_file_tools(self.workspace)} + result = tools["write_file"].execute(path="a/b/c.txt", content="hi") + self.assertNotIn("Error", result) + self.assertTrue((Path(self.tempdir.name) / "a" / "b" / "c.txt").exists()) + + def test_error_messages_do_not_leak_host_path(self): + with self.assertRaises(ValueError) as ctx: + self.workspace.resolve("/etc/passwd") + self.assertNotIn(self.tempdir.name, str(ctx.exception)) + + +class FileToolsTestCase(unittest.TestCase): + def setUp(self): + self.tempdir = tempfile.TemporaryDirectory() + self.addCleanup(self.tempdir.cleanup) + self.workspace = Workspace(self.tempdir.name) + self.tools = {t.name: t for t in make_file_tools(self.workspace)} + + def write(self, rel_path: str, content: str): + full = Path(self.tempdir.name) / rel_path + full.parent.mkdir(parents=True, exist_ok=True) + full.write_text(content, encoding="utf-8") + return full + + +class TestReadFile(FileToolsTestCase): + def test_line_numbering(self): + self.write("file.txt", "one\ntwo\nthree\n") + result = self.tools["read_file"].execute(path="file.txt") + lines = result.splitlines() + self.assertEqual(lines[0], " 1\tone") + self.assertEqual(lines[1], " 2\ttwo") + self.assertEqual(lines[2], " 3\tthree") + + def test_offset_and_limit_window(self): + self.write("file.txt", "\n".join(f"line{i}" for i in range(1, 11)) + "\n") + result = self.tools["read_file"].execute(path="file.txt", offset=3, limit=2) + lines = result.splitlines() + self.assertEqual(len(lines), 2) + self.assertIn("line3", lines[0]) + self.assertIn("line4", lines[1]) + self.assertTrue(lines[0].startswith(" 3\t")) + + def test_truncation_note(self): + self.write("big.txt", "x" * 1000) + result = self.tools["read_file"].execute(path="big.txt", offset=0, limit=0) + # use a small max_read_chars via a fresh tool set + tools = {t.name: t for t in make_file_tools(self.workspace, max_read_chars=50)} + truncated = tools["read_file"].execute(path="big.txt") + self.assertLess(len(result), len(truncated) + 10000) # sanity: result exists + self.assertIn("truncated", truncated.lower()) + self.assertIn("offset", truncated.lower()) + + def test_missing_file_error(self): + result = self.tools["read_file"].execute(path="nope.txt") + self.assertTrue(result.startswith("Error:")) + + def test_binary_file_error(self): + full = Path(self.tempdir.name) / "binary.bin" + full.write_bytes(bytes([0xFF, 0xFE, 0x00, 0x80, 0x81])) + result = self.tools["read_file"].execute(path="binary.bin") + self.assertTrue(result.startswith("Error:")) + + def test_jail_error_propagated(self): + result = self.tools["read_file"].execute(path="../outside.txt") + self.assertTrue(result.startswith("Error:")) + + +class TestWriteFile(FileToolsTestCase): + def test_write_and_confirm(self): + result = self.tools["write_file"].execute(path="out.txt", content="a\nb\n") + self.assertIn("2", result) + full = Path(self.tempdir.name) / "out.txt" + self.assertEqual(full.read_text(), "a\nb\n") + + def test_write_overwrites(self): + self.write("out.txt", "old content") + self.tools["write_file"].execute(path="out.txt", content="new") + full = Path(self.tempdir.name) / "out.txt" + self.assertEqual(full.read_text(), "new") + + def test_omitted_when_read_only(self): + tools = {t.name: t for t in make_file_tools(self.workspace, read_only=True)} + self.assertNotIn("write_file", tools) + self.assertNotIn("edit_file", tools) + self.assertIn("read_file", tools) + self.assertIn("list_files", tools) + self.assertIn("grep_files", tools) + + +class TestEditFile(FileToolsTestCase): + def test_edit_unique_match(self): + self.write("file.txt", "hello world") + result = self.tools["edit_file"].execute( + path="file.txt", old_string="world", new_string="there" + ) + self.assertIn("1", result) + full = Path(self.tempdir.name) / "file.txt" + self.assertEqual(full.read_text(), "hello there") + + def test_edit_zero_matches_error(self): + self.write("file.txt", "hello world") + result = self.tools["edit_file"].execute( + path="file.txt", old_string="missing", new_string="x" + ) + self.assertTrue(result.startswith("Error:")) + self.assertIn("0", result) + + def test_edit_multiple_matches_error(self): + self.write("file.txt", "foo foo foo") + result = self.tools["edit_file"].execute( + path="file.txt", old_string="foo", new_string="bar" + ) + self.assertTrue(result.startswith("Error:")) + self.assertIn("3", result) + + def test_edit_replace_all(self): + self.write("file.txt", "foo foo foo") + result = self.tools["edit_file"].execute( + path="file.txt", + old_string="foo", + new_string="bar", + replace_all=True, + ) + self.assertIn("3", result) + full = Path(self.tempdir.name) / "file.txt" + self.assertEqual(full.read_text(), "bar bar bar") + + def test_omitted_when_read_only(self): + tools = {t.name: t for t in make_file_tools(self.workspace, read_only=True)} + self.assertNotIn("edit_file", tools) + + +class TestListFiles(FileToolsTestCase): + def test_format_and_sorting(self): + self.write("b.txt", "22") + self.write("a.txt", "1") + self.write("sub/c.txt", "333") + result = self.tools["list_files"].execute() + lines = result.splitlines() + rel_paths = [line.split(None, 1)[1] for line in lines] + self.assertEqual(rel_paths, sorted(rel_paths)) + self.assertIn("a.txt", rel_paths) + self.assertIn("sub/c.txt", rel_paths) + + def test_pattern_filters(self): + self.write("a.py", "x") + self.write("b.txt", "x") + result = self.tools["list_files"].execute(pattern="*.py") + self.assertIn("a.py", result) + self.assertNotIn("b.txt", result) + + def test_cap(self): + for i in range(10): + self.write(f"file{i}.txt", "x") + tools = {t.name: t for t in make_file_tools(self.workspace, max_list_entries=3)} + result = tools["list_files"].execute() + lines = [line for line in result.splitlines() if line.strip()] + self.assertIn("truncated", result.lower()) + # 3 file lines + note line(s) + file_lines = [line for line in lines if not line.startswith("[")] + self.assertEqual(len(file_lines), 3) + + def test_missing_directory_error(self): + result = self.tools["list_files"].execute(path="nope") + self.assertTrue(result.startswith("Error:")) + + def test_empty_directory(self): + result = self.tools["list_files"].execute() + self.assertEqual(result, "(no files found)") + + +class TestGrepFiles(FileToolsTestCase): + def test_format(self): + self.write("a.md", "hello world\nfoo bar\n") + result = self.tools["grep_files"].execute(pattern="foo") + self.assertEqual(result, "a.md:2: foo bar") + + def test_glob_filters_files(self): + self.write("a.md", "needle\n") + self.write("b.txt", "needle\n") + result = self.tools["grep_files"].execute(pattern="needle") + self.assertIn("a.md", result) + self.assertNotIn("b.txt", result) + + def test_invalid_regex(self): + result = self.tools["grep_files"].execute(pattern="(unclosed") + self.assertTrue(result.startswith("Error:")) + + def test_no_matches(self): + self.write("a.md", "hello\n") + result = self.tools["grep_files"].execute(pattern="zzz") + self.assertEqual(result, "(no matches found)") + + def test_cap(self): + self.write("a.md", "\n".join("match" for _ in range(20))) + tools = {t.name: t for t in make_file_tools(self.workspace, max_grep_matches=5)} + result = tools["grep_files"].execute(pattern="match") + self.assertIn("truncated", result.lower()) + match_lines = [line for line in result.splitlines() if line.startswith("a.md:")] + self.assertEqual(len(match_lines), 5) + + +class TestHashFiles(unittest.TestCase): + def setUp(self): + self.tempdir = tempfile.TemporaryDirectory() + self.addCleanup(self.tempdir.cleanup) + self.workspace = Workspace(self.tempdir.name) + + def write(self, rel_path, content): + full = Path(self.tempdir.name) / rel_path + full.parent.mkdir(parents=True, exist_ok=True) + full.write_text(content) + return full + + def test_empty_when_no_match(self): + self.assertEqual(hash_files(self.workspace, ["**/*.md"]), "") + + def test_stable_across_runs(self): + self.write("a.md", "content") + h1 = hash_files(self.workspace, ["**/*.md"]) + h2 = hash_files(self.workspace, ["**/*.md"]) + self.assertEqual(h1, h2) + self.assertNotEqual(h1, "") + + def test_changes_with_content(self): + self.write("a.md", "content") + h1 = hash_files(self.workspace, ["**/*.md"]) + self.write("a.md", "different content") + h2 = hash_files(self.workspace, ["**/*.md"]) + self.assertNotEqual(h1, h2) + + def test_changes_with_new_file(self): + self.write("a.md", "content") + h1 = hash_files(self.workspace, ["**/*.md"]) + self.write("b.md", "more content") + h2 = hash_files(self.workspace, ["**/*.md"]) + self.assertNotEqual(h1, h2) + + +class TestToolSchemas(unittest.TestCase): + def test_schemas_have_descriptions_for_all_params(self): + with tempfile.TemporaryDirectory() as tmpdir: + workspace = Workspace(tmpdir) + tools = make_file_tools(workspace) + self.assertGreaterEqual(len(tools), 5) + for t in tools: + schema = t.to_dict() + fn = schema["function"] + self.assertTrue(fn["name"]) + self.assertTrue(fn["description"]) + params = fn["parameters"] + self.assertEqual(params["type"], "object") + for pname, pschema in params["properties"].items(): + self.assertIn( + "description", + pschema, + f"Parameter {pname} of tool {fn['name']} missing description", + ) + self.assertTrue(pschema["description"]) + + +if __name__ == "__main__": + unittest.main() From fe98afda570168cf8d9f323b7f9332f4522d1c7b Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 13:44:40 -0700 Subject: [PATCH 2/8] chore: require llm-interface 0.2 for max_tool_rounds and cache_salt Co-Authored-By: Claude Fable 5.1 --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 31523b1..5e31ba6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,7 @@ html2text = "^2024.2.26" dicttoxml = "^1.7.16" psutil = "^6.1.1" waitress = "^3.0.2" -llm-interface = "^0.1.13" +llm-interface = ">=0.2.0,<0.3.0" [tool.poetry.group.dev.dependencies] From 7c46625d38ce96d9fe6dc18ae1c2a582cb445dc0 Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 17:49:35 -0700 Subject: [PATCH 3/8] ci: fix PR comment step (context.repo.repo) and grant it permissions The comment step passed context.repo.name (undefined) as the repo, so every pull-request job ended in a 404 after its tests had passed. Jobs that comment now declare pull-requests/issues write permission. Co-Authored-By: Claude Fable 5.1 --- .github/workflows/ci_cd.yml | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci_cd.yml b/.github/workflows/ci_cd.yml index a9fb434..1c479bc 100644 --- a/.github/workflows/ci_cd.yml +++ b/.github/workflows/ci_cd.yml @@ -11,6 +11,10 @@ on: jobs: test: runs-on: ubuntu-latest + permissions: + contents: read + pull-requests: write + issues: write strategy: matrix: python-version: [3.12, '3.10', '3.11'] @@ -37,7 +41,7 @@ jobs: github.rest.issues.createComment({ issue_number: context.issue.number, owner: context.repo.owner, - repo: context.repo.name, + repo: context.repo.repo, body: '✅ Tests passed for Python ${{ matrix.python-version }}!' }) - name: Upload coverage reports to Codecov @@ -48,6 +52,10 @@ jobs: test-examples: runs-on: ubuntu-latest + permissions: + contents: read + pull-requests: write + issues: write strategy: matrix: example: [deepsearch ] @@ -78,12 +86,16 @@ jobs: github.rest.issues.createComment({ issue_number: context.issue.number, owner: context.repo.owner, - repo: context.repo.name, + repo: context.repo.repo, body: '✅ Example tests passed for ${{ matrix.example }} (Python ${{ matrix.python-version }})!' }) lint: runs-on: ubuntu-latest + permissions: + contents: read + pull-requests: write + issues: write steps: - uses: actions/checkout@v4 - name: Set up Python @@ -107,7 +119,7 @@ jobs: github.rest.issues.createComment({ issue_number: context.issue.number, owner: context.repo.owner, - repo: context.repo.name, + repo: context.repo.repo, body: '✅ Linting passed!' }) From 9caa0d01d2cb03d5b64549c1751d559d01b64ab0 Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 17:54:22 -0700 Subject: [PATCH 4/8] fix: keep list_files, grep_files and hash_files inside the jail Glob follows symlinks, so a link under the workspace could expose files outside it. Every match now goes through Workspace.resolve() and entries that escape are skipped. Co-Authored-By: Claude Fable 5.1 --- src/planai/tools/filesystem.py | 27 ++++++++++++---- tests/planai/tools/test_filesystem.py | 45 +++++++++++++++++++++++++++ 2 files changed, 66 insertions(+), 6 deletions(-) diff --git a/src/planai/tools/filesystem.py b/src/planai/tools/filesystem.py index 518f48e..b774082 100644 --- a/src/planai/tools/filesystem.py +++ b/src/planai/tools/filesystem.py @@ -85,6 +85,22 @@ def resolve(self, rel_path: str) -> Path: return resolved +def jailed_files(ws: Workspace, base: Path, pattern: str) -> List[Path]: + """Files under ``base`` matching ``pattern`` whose real location is inside the + workspace. Glob follows symlinks, so a link pointing outside the root would + otherwise be readable; such entries are skipped.""" + files = [] + for path in base.glob(pattern): + if not path.is_file(): + continue + try: + ws.resolve(path.relative_to(ws.root).as_posix()) + except ValueError: + continue + files.append(path) + return sorted(files) + + def hash_files(workspace: Union[Workspace, str, Path], globs: List[str]) -> str: """Compute a stable hash over the content of files matching the given globs. @@ -104,10 +120,9 @@ def hash_files(workspace: Union[Workspace, str, Path], globs: List[str]) -> str: matched = {} for pattern in globs: - for path in ws.root.glob(pattern): - if path.is_file(): - rel = path.relative_to(ws.root).as_posix() - matched[rel] = path + for path in jailed_files(ws, ws.root, pattern): + rel = path.relative_to(ws.root).as_posix() + matched[rel] = path if not matched: return "" @@ -296,7 +311,7 @@ def list_files(path: str = ".", pattern: str = "**/*") -> str: return f"Error: Not a directory: {path}" try: - matches = sorted(p for p in target.glob(pattern) if p.is_file()) + matches = jailed_files(ws, target, pattern) except (re.error, ValueError) as e: return f"Error: Invalid pattern: {e}" @@ -345,7 +360,7 @@ def grep_files(pattern: str, path: str = ".", glob: str = "**/*.md") -> str: return f"Error: Invalid regular expression: {e}" try: - candidates = sorted(p for p in target.glob(glob) if p.is_file()) + candidates = jailed_files(ws, target, glob) except ValueError as e: return f"Error: Invalid pattern: {e}" diff --git a/tests/planai/tools/test_filesystem.py b/tests/planai/tools/test_filesystem.py index 7d61c19..854294d 100644 --- a/tests/planai/tools/test_filesystem.py +++ b/tests/planai/tools/test_filesystem.py @@ -316,3 +316,48 @@ def test_schemas_have_descriptions_for_all_params(self): if __name__ == "__main__": unittest.main() + + +class TestSymlinkEscapeInWalks(unittest.TestCase): + """list_files, grep_files and hash_files must not follow symlinks out of the jail.""" + + def setUp(self): + import tempfile + + self.outside_dir = tempfile.TemporaryDirectory() + self.root_dir = tempfile.TemporaryDirectory() + outside = Path(self.outside_dir.name) / "secret.md" + outside.write_text("top secret needle") + root = Path(self.root_dir.name) + (root / "inside.md").write_text("inside needle") + (root / "link.md").symlink_to(outside) + (root / "linkdir").symlink_to(Path(self.outside_dir.name)) + self.ws = Workspace(root) + + def tearDown(self): + self.outside_dir.cleanup() + self.root_dir.cleanup() + + def _tool(self, name): + return next(t for t in make_file_tools(self.ws) if t.name == name) + + def test_list_skips_symlinked_files_and_directories(self): + listing = self._tool("list_files").execute(path=".", pattern="**/*") + self.assertIn("inside.md", listing) + self.assertNotIn("link.md", listing) + self.assertNotIn("secret.md", listing) + + def test_grep_skips_symlinked_content(self): + hits = self._tool("grep_files").execute( + pattern="needle", path=".", glob="**/*.md" + ) + self.assertIn("inside.md", hits) + self.assertNotIn("top secret", hits) + self.assertNotIn("link.md", hits) + + def test_hash_ignores_symlinked_content(self): + before = hash_files(self.ws, ["**/*.md"]) + (Path(self.outside_dir.name) / "secret.md").write_text("changed outside") + self.assertEqual(before, hash_files(self.ws, ["**/*.md"])) + (Path(self.root_dir.name) / "inside.md").write_text("changed inside") + self.assertNotEqual(before, hash_files(self.ws, ["**/*.md"])) From 9e4cdb7429fae89c8b60016362b7d0d731bb2ae7 Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 18:17:09 -0700 Subject: [PATCH 5/8] chore: lock llm-interface 0.2.0 Co-Authored-By: Claude Fable 5.1 --- poetry.lock | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/poetry.lock b/poetry.lock index 7712c12..cc37272 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.1.2 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.4.3 and should not be changed by hand. [[package]] name = "alabaster" @@ -451,6 +451,7 @@ files = [ {file = "colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6"}, {file = "colorama-0.4.6.tar.gz", hash = "sha256:08695f5cb7ed6e0531a20572697297273c47b8cae5a63ffc6d6ed5c201be6e44"}, ] +markers = {dev = "platform_system == \"Windows\" or sys_platform == \"win32\"", docs = "sys_platform == \"win32\""} [[package]] name = "coverage" @@ -1089,14 +1090,14 @@ files = [ [[package]] name = "llm-interface" -version = "0.1.13" +version = "0.2.0" description = "A flexible interface for working with various LLM providers" optional = false python-versions = "<4.0,>=3.10" groups = ["main"] files = [ - {file = "llm_interface-0.1.13-py3-none-any.whl", hash = "sha256:7e6cad32e9677f2abaeb31cd88f7cb9423d08e510158684f02b99676e7859318"}, - {file = "llm_interface-0.1.13.tar.gz", hash = "sha256:f1f89efce3b1e3ba9d86dac482477dd13d667488021cbf86be209dc5335618fa"}, + {file = "llm_interface-0.2.0-py3-none-any.whl", hash = "sha256:ce5b04fcc89829e2d5520992d03806152270dfffdd611c21a3b4d26b8c205679"}, + {file = "llm_interface-0.2.0.tar.gz", hash = "sha256:001a5db4d116aba761b0e668d9819c0dbc93c847f0f416f1e9079fecc4a5a74e"}, ] [package.dependencies] @@ -2352,4 +2353,4 @@ docs = [] [metadata] lock-version = "2.1" python-versions = "^3.10" -content-hash = "93aec2a7d1a661e225e53910d62d791d5d05cf6d6fbc36ca28fd6d38330d8023" +content-hash = "25d7f1374756ba517eaa8955d2c328b0c5825b41a77334b1772ab7946bd8618b" From a9f49216a35b11d4ca8c78f380e1d3a305ccd199 Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 20:38:36 -0700 Subject: [PATCH 6/8] docs: document workspaces and file tools in the Astro site Add a Workspaces and File Tools feature page and cover the new WorkspaceLLMTaskWorker, WorkspaceTask, Workspace, make_file_tools, hash_files, the get_tools()/max_tool_rounds hooks and the _cache_hit_is_valid cache hook on the LLM integration, caching, task worker, usage and API reference pages. Co-Authored-By: Claude Fable 5.1 --- docs-astro/astro.config.mjs | 1 + docs-astro/src/content/docs/api/index.md | 14 ++ docs-astro/src/content/docs/api/taskworker.md | 54 +++++- .../src/content/docs/features/caching.md | 27 +++ .../content/docs/features/llm-integration.md | 36 +++- .../src/content/docs/features/taskworkers.md | 19 ++ .../src/content/docs/features/workspaces.md | 175 ++++++++++++++++++ docs-astro/src/content/docs/guide/usage.md | 30 +++ 8 files changed, 352 insertions(+), 4 deletions(-) create mode 100644 docs-astro/src/content/docs/features/workspaces.md diff --git a/docs-astro/astro.config.mjs b/docs-astro/astro.config.mjs index 02a4039..685b530 100644 --- a/docs-astro/astro.config.mjs +++ b/docs-astro/astro.config.mjs @@ -35,6 +35,7 @@ export default defineConfig({ items: [ { label: 'Task Workers', slug: 'features/taskworkers' }, { label: 'LLM Integration', slug: 'features/llm-integration' }, + { label: 'Workspaces & File Tools', slug: 'features/workspaces' }, { label: 'Caching', slug: 'features/caching' }, { label: 'Subgraphs', slug: 'features/subgraphs' }, ], diff --git a/docs-astro/src/content/docs/api/index.md b/docs-astro/src/content/docs/api/index.md index 9ed2ab3..528bd6c 100644 --- a/docs-astro/src/content/docs/api/index.md +++ b/docs-astro/src/content/docs/api/index.md @@ -29,6 +29,13 @@ Specialized workers for integrating Large Language Models into workflows, includ - **CachedTaskWorker**: Base class for workers with caching - **CachedLLMTaskWorker**: LLM worker with response caching +### Workspaces and File Tools +- **Workspace**: A sandboxed directory that every file path is resolved against +- **make_file_tools**: Builds read/write/edit/list/grep tools jailed to a Workspace +- **hash_files**: Content hash of the files matching glob patterns, for cache keys +- **WorkspaceTask**: Task that carries the workspace path through provenance +- **WorkspaceLLMTaskWorker**: CachedLLMTaskWorker with file tools and file-aware caching + ### Advanced Workers - **InitialTaskWorker**: Entry point for workflows - **JoinedTaskWorker**: Aggregates multiple task results @@ -68,6 +75,13 @@ from planai import ( CachedTaskWorker, SubGraphWorker, + # Workspaces and file tools + Workspace, + make_file_tools, + hash_files, + WorkspaceTask, + WorkspaceLLMTaskWorker, + # Utilities Dispatcher, InputProvenance, diff --git a/docs-astro/src/content/docs/api/taskworker.md b/docs-astro/src/content/docs/api/taskworker.md index 5a3d05b..b4e45b5 100644 --- a/docs-astro/src/content/docs/api/taskworker.md +++ b/docs-astro/src/content/docs/api/taskworker.md @@ -264,6 +264,15 @@ def extra_validation(self, response: Task, input_task: Task) -> Optional[str]: return None ``` +##### get_tools +```python +def get_tools(self, task: Task) -> Optional[List[Tool]]: + """Tools for this task; defaults to the static ``tools`` field""" + return make_file_tools(Workspace(task.job_dir), read_only=True) +``` + +The `max_tool_rounds` field bounds the number of tool-calling rounds per request. When it is reached the model is asked for a final answer with tools disabled. See [Workspaces and File Tools](/features/workspaces/). + ### Real-World Example ```python @@ -314,6 +323,33 @@ class ExpensiveAnalysis(CachedLLMTaskWorker): pass ``` +### WorkspaceLLMTaskWorker + +A `CachedLLMTaskWorker` that hands the LLM file tools jailed to a per-job directory found in the task's provenance: + +```python +from planai import WorkspaceLLMTaskWorker, WorkspaceTask + +class Editor(WorkspaceLLMTaskWorker): + prompt: str = "Fix the inconsistencies between the sections in draft/report.md" + llm_input_type: Type[Task] = WorkspaceTask + output_types: List[Type[Task]] = [EditSummary] + input_globs: List[str] = ["draft/*.md"] + max_tool_rounds: int = 40 + + def expected_output_files(self, task: WorkspaceTask) -> List[str]: + return ["draft/report.md"] +``` + +Fields: `read_only` (default `False`), `max_tool_rounds` (default `40`), `input_globs` (default empty), and `max_read_chars` (default `100000`). + +Hooks: +- `get_workspace(task)` returns the `Workspace`, taken from the nearest task with a string `workspace` attribute; override it to source the directory elsewhere. +- `expected_output_files(task)` lists workspace-relative files whose absence invalidates a cache hit. +- `get_cache_salt(task)` forwards the cache key to the LLM response cache. + +See [Workspaces and File Tools](/features/workspaces/) for the full guide. + ## CachedTaskWorker Specialized worker for expensive operations that benefit from persistent caching: @@ -381,6 +417,17 @@ def extra_cache_key(self, task: Task) -> str: return f"{self.custom_setting}_{task.priority}" ``` +#### _cache_hit_is_valid +```python +def _cache_hit_is_valid( + self, task: Task, cached_results: List[Tuple[str, Task]] +) -> bool: + """Reject a cache hit based on state the key does not capture""" + return (Path(task.output_dir) / "index.json").exists() +``` + +Called on every cache hit before the cached results are published. Returning `False` logs the hit as no longer valid and runs `consume_work()` as on a miss. + ### Real-World Example ```python @@ -417,9 +464,10 @@ class DocumentAnalyzer(CachedTaskWorker): #### Cache Hit When input matches cached data: 1. `pre_consume_work()` is called -2. Cached results are published directly -3. `consume_work()` is **skipped** -4. `post_consume_work()` is called +2. `_cache_hit_is_valid()` may reject the hit, in which case the miss path runs +3. Cached results are published directly +4. `consume_work()` is **skipped** +5. `post_consume_work()` is called #### Cache Miss When no cached data exists: diff --git a/docs-astro/src/content/docs/features/caching.md b/docs-astro/src/content/docs/features/caching.md index cf9f383..66d3000 100644 --- a/docs-astro/src/content/docs/features/caching.md +++ b/docs-astro/src/content/docs/features/caching.md @@ -118,6 +118,33 @@ class VersionedCacheWorker(CachedTaskWorker): return algorithm_version ``` +## Caching Work That Touches Files + +The cache key covers the input task, the prompt, and `extra_cache_key()`, but not the files a worker reads or writes. Two hooks close that gap: + +- `hash_files(workspace, globs)` returns a content hash for the files matching a set of glob patterns. Return it from `extra_cache_key()` so a changed input file produces a new key. +- `_cache_hit_is_valid(task, cached_results)` runs on every cache hit and can reject it based on state the key does not capture. Return `False` when an output file from the cached run no longer exists, and the worker re-executes instead of replaying stale results. + +```python +from pathlib import Path +from planai import CachedTaskWorker, hash_files + +class Indexer(CachedTaskWorker): + output_types: List[Type[Task]] = [IndexBuilt] + + def extra_cache_key(self, task: JobTask) -> str: + return hash_files(task.job_dir, ["docs/**/*.md"]) + + def _cache_hit_is_valid(self, task: JobTask, cached_results) -> bool: + return (Path(task.job_dir) / "index.json").exists() + + def consume_work(self, task: JobTask): + # build docs/index.json from the markdown files + ... +``` + +`WorkspaceLLMTaskWorker` implements both through its `input_globs` field and `expected_output_files()` hook, and additionally salts the LLM response cache with the same key. See [Workspaces and File Tools](/features/workspaces/). + ## Next Steps - Learn about [Task Workers](/features/taskworkers/) that can be cached diff --git a/docs-astro/src/content/docs/features/llm-integration.md b/docs-astro/src/content/docs/features/llm-integration.md index 13d2c7e..9e33da0 100644 --- a/docs-astro/src/content/docs/features/llm-integration.md +++ b/docs-astro/src/content/docs/features/llm-integration.md @@ -22,6 +22,20 @@ llm = llm_from_config( ) ``` +### Anthropic + +```python +# Set ANTHROPIC_API_KEY in your environment +llm = llm_from_config( + provider="anthropic", + model_name="claude-sonnet-5", + thinking={"type": "adaptive"}, # optional extended thinking + effort="high", # optional: low, medium, high, xhigh, max +) +``` + +Anthropic requests use structured outputs and prompt caching by default. + ### Ollama (Local Models) ```python @@ -176,6 +190,25 @@ class AssistantWorker(LLMTaskWorker): # Tools are automatically registered with the LLM ``` +### Per-Task Tools + +The `tools` field binds the same tools to every task. Override `get_tools()` when the tools depend on the task, for example file tools jailed to a directory that arrives with the task: + +```python +from planai import LLMTaskWorker, Workspace, make_file_tools + +class RepoReviewer(LLMTaskWorker): + prompt = "Review the repository and report the three most important issues" + llm_input_type = RepoTask + output_types: List[Type[Task]] = [Review] + max_tool_rounds: int = 30 + + def get_tools(self, task: RepoTask): + return make_file_tools(Workspace(task.checkout_dir), read_only=True) +``` + +`max_tool_rounds` bounds how many rounds of tool calls the model may make within one request. When the limit is reached, the model is asked for its final structured answer with tools disabled, so a worker always produces an output task. See [Workspaces and File Tools](/features/workspaces/) for `WorkspaceLLMTaskWorker`, which wires up the file tools, workspace discovery, and file-aware caching for you. + ## Streaming Responses For real-time applications, enable response streaming: @@ -402,4 +435,5 @@ class TokenTrackingWorker(LLMTaskWorker): - Learn about [Prompt Engineering](/guide/prompts/) for better results - Explore [Caching Strategies](/features/caching/) for cost optimization - See [Examples](https://github.com/provos/planai/tree/main/examples) using LLMs -- Read about [Prompt Optimization](/cli/prompt-optimization/) tools \ No newline at end of file +- Read about [Prompt Optimization](/cli/prompt-optimization/) tools +- Give the LLM files to work on with [Workspaces and File Tools](/features/workspaces/) \ No newline at end of file diff --git a/docs-astro/src/content/docs/features/taskworkers.md b/docs-astro/src/content/docs/features/taskworkers.md index eedfd80..b692782 100644 --- a/docs-astro/src/content/docs/features/taskworkers.md +++ b/docs-astro/src/content/docs/features/taskworkers.md @@ -143,6 +143,25 @@ class CachedAnalyzer(CachedLLMTaskWorker): This helps during development to save model costs or avoiding repeat processing if a graph fails to run. Any changes to the prompt, model or input_data will lead to a new cache key. +### WorkspaceLLMTaskWorker + +A `CachedLLMTaskWorker` whose LLM can read, write, edit, list, and search files inside a per-job directory: + +```python +from planai import WorkspaceLLMTaskWorker, WorkspaceTask + +class ReportWriter(WorkspaceLLMTaskWorker): + prompt = "Read notes/*.md and write the report to report.md" + llm_input_type = WorkspaceTask + output_types: List[Type[Task]] = [ReportSummary] + input_globs: List[str] = ["notes/*.md"] + + def expected_output_files(self, task: WorkspaceTask) -> List[str]: + return ["report.md"] +``` + +The directory comes from a `WorkspaceTask` upstream in the provenance chain, every path the model uses is jailed to it, and the cache is invalidated when the input files change or the output files are missing. See [Workspaces and File Tools](/features/workspaces/) for details. + ### JoinedTaskWorker Aggregates results from multiple tasks: diff --git a/docs-astro/src/content/docs/features/workspaces.md b/docs-astro/src/content/docs/features/workspaces.md new file mode 100644 index 0000000..32dfc3a --- /dev/null +++ b/docs-astro/src/content/docs/features/workspaces.md @@ -0,0 +1,175 @@ +--- +title: Workspaces and File Tools +description: Let an LLM read, write, edit, and search files inside a sandboxed per-job directory +--- + +Some work does not fit through a prompt and a structured response: reviewing a repository, drafting a long report one section at a time, or editing a document that another worker produced. `WorkspaceLLMTaskWorker` gives the LLM a set of file tools that are jailed to a per-job working directory. The worker returns a small structured result while the large text lives on disk, and many jobs can run concurrently without any of them reaching outside the directory assigned to it. + +## Overview + +PlanAI exports four pieces that work together: + +- **`Workspace`** is the sandbox: a root directory that every path is resolved against. +- **`make_file_tools()`** builds `read_file`, `write_file`, `edit_file`, `list_files`, and `grep_files` tools bound to one workspace. +- **`WorkspaceTask`** carries the workspace path through the graph. +- **`WorkspaceLLMTaskWorker`** is a `CachedLLMTaskWorker` that finds the workspace in the task's provenance, hands the tools to the LLM, and keeps the cache honest about files on disk. + +## A Minimal Example + +```python +from typing import List, Type +from planai import Graph, Task, WorkspaceLLMTaskWorker, WorkspaceTask, llm_from_config + +class ReviewResult(Task): + summary: str + issues_found: int + +class CodeReviewer(WorkspaceLLMTaskWorker): + prompt = "Review the code in this workspace and write your findings to review.md." + llm_input_type: Type[Task] = WorkspaceTask + output_types: List[Type[Task]] = [ReviewResult] + + def expected_output_files(self, task: WorkspaceTask) -> List[str]: + # a cache hit whose files are missing (e.g. a fresh checkout) is re-executed + return ["review.md"] + +llm = llm_from_config(provider="anthropic", model_name="claude-sonnet-5") +reviewer = CodeReviewer(llm=llm, max_tool_rounds=40) + +graph = Graph(name="Review Workflow") +graph.add_workers(reviewer) +graph.set_entry(reviewer) +graph.set_exit(reviewer) +graph.run(initial_tasks=[(reviewer, WorkspaceTask(workspace="/jobs/job-123"))]) +``` + +The model reads and writes files with the tools during the request, and the structured `ReviewResult` is what flows to downstream workers. + +## How the Workspace Is Found + +`WorkspaceLLMTaskWorker.get_workspace()` walks the task and then its input provenance, nearest first, and returns the first task with a non-empty string `workspace` attribute. Publish a `WorkspaceTask` (or your own `Task` subclass with a `workspace: str` field) once, and every workspace worker downstream operates on the same directory: + +```python +from pathlib import Path +from planai import TaskWorker, WorkspaceTask + +class JobSetup(TaskWorker): + output_types: List[Type[Task]] = [WorkspaceTask] + + def consume_work(self, task: JobRequest): + job_dir = Path("work") / task.job_id + job_dir.mkdir(parents=True, exist_ok=True) + (job_dir / "input.md").write_text(task.text) + self.publish_work( + WorkspaceTask(workspace=str(job_dir.resolve())), input_task=task + ) +``` + +Override `get_workspace()` when the directory should come from worker configuration instead of the provenance chain. + +## The File Tools + +| Tool | What it does | +| --- | --- | +| `read_file(path, offset=0, limit=0)` | Returns the file with 1-based line numbers, like `cat -n`. `offset` and `limit` page through large files; output is truncated at `max_read_chars`. | +| `write_file(path, content)` | Creates or overwrites a file. Parent directories are created as needed. | +| `edit_file(path, old_string, new_string, replace_all=False)` | Replaces an exact snippet. `old_string` must occur exactly once unless `replace_all` is true. | +| `list_files(path=".", pattern="**/*")` | Lists files with their sizes, sorted by relative path. | +| `grep_files(pattern, path=".", glob="**/*.md")` | Searches each line of the matching files with a Python regular expression. | + +Every path is workspace-relative. Absolute paths, `..` components, and symlinks that resolve outside the root are rejected. The tools return an `Error: ...` string instead of raising, so the model sees what went wrong and can correct itself. Set `read_only=True` on the worker to omit `write_file` and `edit_file`. + +## Caching and Files on Disk + +`WorkspaceLLMTaskWorker` extends `CachedLLMTaskWorker`, and files introduce two problems the normal cache key does not cover: + +1. **Changed inputs.** The cache key is built from the input task and the prompt, not from the files the model will read. Set `input_globs` to fold the content of the matching files into the key through `hash_files()`. +2. **Missing outputs.** A cache hit replays the published output tasks, but not the files the tools wrote on the earlier run. Override `expected_output_files()` to list the workspace-relative paths the worker must produce. When any of them is missing, the hit is treated as a miss and the worker runs again. + +```python +class SectionWriter(WorkspaceLLMTaskWorker): + prompt = "Read the notes and write the requested section to the given file." + llm_input_type: Type[Task] = SectionRequest + output_types: List[Type[Task]] = [SectionDraft] + input_globs: List[str] = ["notes/*.md"] + + def expected_output_files(self, task: SectionRequest) -> List[str]: + return [task.output_file] +``` + +The worker also forwards its cache key to `llm-interface` as `cache_salt`. The library's own response cache is keyed on the initial prompt, so without the salt a second run could receive a stale answer after the files changed. + +## Validating What the Model Wrote + +Use `extra_validation()` to check the files and not only the structured response. Returning a string sends it back to the model as feedback and the request is retried: + +```python +from typing import Optional + +class SectionWriter(WorkspaceLLMTaskWorker): + ... + + def extra_validation( + self, response: SectionDraft, input_task: SectionRequest + ) -> Optional[str]: + path = self.get_workspace(input_task).resolve(input_task.output_file) + if not path.exists(): + return f"Write the section to {input_task.output_file} with write_file." + if len(path.read_text()) < 800: + return "The section is too short; expand it to at least 800 characters." + return None +``` + +This keeps the expensive judgement with the model and the cheap, deterministic checks in Python. + +## Configuration + +| Field | Default | Purpose | +| --- | --- | --- | +| `read_only` | `False` | Expose only `read_file`, `list_files`, and `grep_files`. | +| `max_tool_rounds` | `40` | Maximum tool-calling rounds per request. When the limit is reached the model is asked for its final answer with tools disabled. | +| `input_globs` | `[]` | Glob patterns whose file content is folded into the cache key. | +| `max_read_chars` | `100000` | Truncation limit for `read_file` output. | +| `tools` | `None` | Additional tools, appended after the file tools. | + +## Using the Pieces Directly + +`Workspace`, `make_file_tools()`, and `hash_files()` are exported from `planai` and work with any `LLMTaskWorker` through the `get_tools()` hook: + +```python +from planai import LLMTaskWorker, Workspace, make_file_tools + +class Summarizer(LLMTaskWorker): + prompt = "Summarize every file in the workspace." + llm_input_type: Type[Task] = JobTask + output_types: List[Type[Task]] = [Summary] + max_tool_rounds: int = 20 + + def get_tools(self, task: JobTask): + return make_file_tools(Workspace(task.job_dir), read_only=True) +``` + +`make_file_tools()` also accepts `max_read_chars`, `max_list_entries`, and `max_grep_matches` to bound the size of tool results. `Workspace.resolve(rel_path)` is the same check the tools use, so Python code can validate a model-supplied path before touching it. + +## Testing + +The tools are plain functions wrapped as `Tool` objects, so they can be exercised without an LLM: + +```python +def test_edit_requires_unique_match(tmp_path): + (tmp_path / "a.md").write_text("one two one") + tools = {t.name: t for t in make_file_tools(tmp_path)} + + result = tools["edit_file"].execute(path="a.md", old_string="one", new_string="1") + + assert result.startswith("Error: old_string is not unique") + assert tools["read_file"].execute(path="a.md") == " 1\tone two one" +``` + +For the worker itself, point a `WorkspaceTask` at `tmp_path`, drive it with [`InvokeTaskWorker`](/guide/testing/), and assert on both the published task and the files left in the directory. + +## Next Steps + +- See [LLM Integration](/features/llm-integration/) for the `get_tools()` hook and `max_tool_rounds` +- Read about [Caching](/features/caching/) to understand how file hashes and validity checks extend the cache key +- Review the [TaskWorker API](/api/taskworker/) for the full list of hooks diff --git a/docs-astro/src/content/docs/guide/usage.md b/docs-astro/src/content/docs/guide/usage.md index d50025a..e1d0f45 100644 --- a/docs-astro/src/content/docs/guide/usage.md +++ b/docs-astro/src/content/docs/guide/usage.md @@ -203,6 +203,36 @@ initial_input = Task1WorkItem(data="start") main_graph.run(initial_tasks=[(subgraph_worker, initial_input)]) ``` +### Letting an LLM Work with Files + +`WorkspaceLLMTaskWorker` is a `CachedLLMTaskWorker` that gives the LLM `read_file`, `write_file`, `edit_file`, `list_files`, and `grep_files` tools jailed to a per-job working directory. Use it when the text is too large to pass through a prompt and a structured response, or when many concurrent jobs each need their own directory that the model cannot escape. + +The directory is discovered from the task's provenance: publish a `WorkspaceTask` (or any `Task` with a string `workspace` attribute) upstream and `get_workspace()` finds the nearest one. All paths the model uses are relative to that directory; absolute paths, `..` components, and symlinks that leave the directory are rejected. + +```python +from planai import Graph, WorkspaceLLMTaskWorker, WorkspaceTask, llm_from_config + +class CodeReviewer(WorkspaceLLMTaskWorker): + prompt = "Review the code in this repository and write your findings to review.md." + llm_input_type: Type[Task] = WorkspaceTask + output_types: List[Type[Task]] = [ReviewResult] + + def expected_output_files(self, task: WorkspaceTask) -> List[str]: + # a cache hit whose files are missing (e.g. a fresh checkout) is re-executed + return ["review.md"] + +llm = llm_from_config(provider="anthropic", model_name="claude-sonnet-5") +reviewer = CodeReviewer(llm=llm, max_tool_rounds=40) + +graph = Graph(name="Review Workflow") +graph.add_workers(reviewer) +graph.set_entry(reviewer) +graph.set_exit(reviewer) +graph.run(initial_tasks=[(reviewer, WorkspaceTask(workspace="/jobs/job-123"))]) +``` + +Because a cache hit replays only the published output tasks and not the files the tools wrote, `expected_output_files()` lists the paths the worker must produce; if any are missing the entry is treated as a miss. Set `input_globs` to fold the content of matching workspace files into the cache key, and pass `read_only=True` to omit the writing tools. The full guide is [Workspaces and File Tools](/features/workspaces/). + ## Best Practices 1. **Modular Design**: Break down complex tasks into smaller, reusable TaskWorkers. From ecaee180e971fd82a915dd181822551ff04955aa Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 22:37:31 -0700 Subject: [PATCH 7/8] fix: harden workspace file tools and their cache interaction Address review findings on the file-tools branch: - Workers whose tools write files now pass a fresh cache_salt on every execution so llm-interface's response cache cannot replay an answer whose tool calls (the file writes) would be skipped, e.g. after a PlanAI cache hit was rejected for missing output files. Read-only workers keep using the cache key as the salt. - The workspace root is always part of the cache key, so identical payloads in different workspaces no longer share an entry. - extra_cache_key() no longer raises when no workspace is in the provenance chain; get_tools() reports that instead. - The lookup cache key is memoized per task on the worker thread and reused by get_cache_salt(); the store still recomputes it after the run on purpose (documented) so a worker that edits files it also hashes is found by a later run over the edited files. - Static tools that would shadow a file tool are rejected. - Workspace() no longer creates the directory as a side effect. - Absolute or empty glob patterns are rejected with ValueError instead of surfacing pathlib's NotImplementedError. - grep_files searches every file by default (was Markdown only) and truncates long matching lines (max_grep_line_chars). - edit_file preserves CRLF line endings; write_file writes verbatim. - hash_files raises a clear OSError naming the unreadable file. - Fix the mypy error on the cached-results Optional and annotate get_tools(). Co-Authored-By: Claude Fable 5.1 --- docs-astro/src/content/docs/api/taskworker.md | 2 +- .../src/content/docs/features/caching.md | 2 +- .../src/content/docs/features/workspaces.md | 10 +- src/planai/cached_task.py | 93 ++++++++++++------- src/planai/tools/filesystem.py | 78 +++++++++++----- src/planai/workspace_task.py | 48 ++++++++-- tests/planai/test_workspace_task.py | 67 ++++++++++++- tests/planai/tools/test_filesystem.py | 93 ++++++++++++++++++- 8 files changed, 322 insertions(+), 71 deletions(-) diff --git a/docs-astro/src/content/docs/api/taskworker.md b/docs-astro/src/content/docs/api/taskworker.md index b4e45b5..609509e 100644 --- a/docs-astro/src/content/docs/api/taskworker.md +++ b/docs-astro/src/content/docs/api/taskworker.md @@ -346,7 +346,7 @@ Fields: `read_only` (default `False`), `max_tool_rounds` (default `40`), `input_ Hooks: - `get_workspace(task)` returns the `Workspace`, taken from the nearest task with a string `workspace` attribute; override it to source the directory elsewhere. - `expected_output_files(task)` lists workspace-relative files whose absence invalidates a cache hit. -- `get_cache_salt(task)` forwards the cache key to the LLM response cache. +- `get_cache_salt(task)` salts the LLM response cache: the cache key for read-only workers, a fresh value per execution for workers whose tools write files. See [Workspaces and File Tools](/features/workspaces/) for the full guide. diff --git a/docs-astro/src/content/docs/features/caching.md b/docs-astro/src/content/docs/features/caching.md index 66d3000..33ddfc0 100644 --- a/docs-astro/src/content/docs/features/caching.md +++ b/docs-astro/src/content/docs/features/caching.md @@ -143,7 +143,7 @@ class Indexer(CachedTaskWorker): ... ``` -`WorkspaceLLMTaskWorker` implements both through its `input_globs` field and `expected_output_files()` hook, and additionally salts the LLM response cache with the same key. See [Workspaces and File Tools](/features/workspaces/). +`WorkspaceLLMTaskWorker` implements both through its `input_globs` field and `expected_output_files()` hook. Note that the key is computed again when results are stored, after the worker ran, so a worker that changes files it also hashes is found by a later run over the changed files. See [Workspaces and File Tools](/features/workspaces/). ## Next Steps diff --git a/docs-astro/src/content/docs/features/workspaces.md b/docs-astro/src/content/docs/features/workspaces.md index 32dfc3a..7cd436b 100644 --- a/docs-astro/src/content/docs/features/workspaces.md +++ b/docs-astro/src/content/docs/features/workspaces.md @@ -75,13 +75,13 @@ Override `get_workspace()` when the directory should come from worker configurat | `write_file(path, content)` | Creates or overwrites a file. Parent directories are created as needed. | | `edit_file(path, old_string, new_string, replace_all=False)` | Replaces an exact snippet. `old_string` must occur exactly once unless `replace_all` is true. | | `list_files(path=".", pattern="**/*")` | Lists files with their sizes, sorted by relative path. | -| `grep_files(pattern, path=".", glob="**/*.md")` | Searches each line of the matching files with a Python regular expression. | +| `grep_files(pattern, path=".", glob="**/*")` | Searches each line of the matching files with a Python regular expression. Long matching lines are truncated. | Every path is workspace-relative. Absolute paths, `..` components, and symlinks that resolve outside the root are rejected. The tools return an `Error: ...` string instead of raising, so the model sees what went wrong and can correct itself. Set `read_only=True` on the worker to omit `write_file` and `edit_file`. ## Caching and Files on Disk -`WorkspaceLLMTaskWorker` extends `CachedLLMTaskWorker`, and files introduce two problems the normal cache key does not cover: +`WorkspaceLLMTaskWorker` extends `CachedLLMTaskWorker`. Its cache key always includes the workspace root, so two jobs with the same payload in different directories never share an entry. Files introduce two further problems the normal cache key does not cover: 1. **Changed inputs.** The cache key is built from the input task and the prompt, not from the files the model will read. Set `input_globs` to fold the content of the matching files into the key through `hash_files()`. 2. **Missing outputs.** A cache hit replays the published output tasks, but not the files the tools wrote on the earlier run. Override `expected_output_files()` to list the workspace-relative paths the worker must produce. When any of them is missing, the hit is treated as a miss and the worker runs again. @@ -97,7 +97,9 @@ class SectionWriter(WorkspaceLLMTaskWorker): return [task.output_file] ``` -The worker also forwards its cache key to `llm-interface` as `cache_salt`. The library's own response cache is keyed on the initial prompt, so without the salt a second run could receive a stale answer after the files changed. +The key is computed again when the results are stored, after the worker ran. A worker that edits files it also lists in `input_globs` is therefore found by a later run over the edited files, and is not replayed over the unedited ones, which is when replaying without running would be wrong. + +`llm-interface` keeps its own response cache keyed on the initial prompt. A replayed response skips the tool calls, so for a worker that can write files the files would never be written. Such workers pass a fresh `cache_salt` on every execution and only ever hit the PlanAI cache. Read-only workers pass their cache key as the salt, so an identical request can still be served from the response cache. ## Validating What the Model Wrote @@ -149,7 +151,7 @@ class Summarizer(LLMTaskWorker): return make_file_tools(Workspace(task.job_dir), read_only=True) ``` -`make_file_tools()` also accepts `max_read_chars`, `max_list_entries`, and `max_grep_matches` to bound the size of tool results. `Workspace.resolve(rel_path)` is the same check the tools use, so Python code can validate a model-supplied path before touching it. +`make_file_tools()` also accepts `max_read_chars`, `max_list_entries`, `max_grep_matches`, and `max_grep_line_chars` to bound the size of tool results. Constructing a `Workspace` has no side effects: the directory is created by the first `write_file`. `Workspace.resolve(rel_path)` is the same check the tools use, so Python code can validate a model-supplied path before touching it. ## Testing diff --git a/src/planai/cached_task.py b/src/planai/cached_task.py index 159b469..b1774f0 100644 --- a/src/planai/cached_task.py +++ b/src/planai/cached_task.py @@ -16,7 +16,7 @@ import logging import sys import threading -from typing import List, Tuple +from typing import List, Optional, Tuple from diskcache import Cache from pydantic import Field, PrivateAttr @@ -44,38 +44,67 @@ def _pre_consume_work(self, task: Task): self.pre_consume_work(task) cache_key = self._get_cache_key(task) + # remember the lookup key for the duration of this task so that hooks + # running inside consume_work (e.g. get_cache_salt) do not recompute it + self._local.lookup_cache_key = (task, cache_key) try: - result = self._cache.get(cache_key) - except Exception as e: - logging.error("Error getting data from cache %s: %s", cache_key, str(e)) - result = None - - cache_hit = False - cached_results = None - if result is not None: - cached_results, _ = result - cache_hit = self._cache_hit_is_valid(task, cached_results) - if not cache_hit: - logging.info( - "Cache hit for %s with key: %s is no longer valid; re-executing", - self.name, - cache_key, - ) - - if cache_hit: - logging.info("Cache hit for %s with key: %s", self.name, cache_key) - self._publish_cached_results(cached_results, task) - else: - logging.info("Cache miss for %s with key: %s", self.name, cache_key) - self.consume_work(task) - input_task, outputs = self._local.ctx.get_input_and_outputs() - # strip private fields from outputs - outputs = [ - [consumer.name, task.copy_public()] for consumer, task in outputs - ] - self._set_cache(input_task, outputs) - - self.post_consume_work(task) + self._consume_with_cache(task, cache_key) + finally: + self._local.lookup_cache_key = None + + def _consume_with_cache(self, task: Task, cache_key: str): + try: + result = self._cache.get(cache_key) + except Exception as e: + logging.error("Error getting data from cache %s: %s", cache_key, str(e)) + result = None + + cached_results: Optional[List[Tuple[str, Task]]] = None + if result is not None: + cached_results, _ = result + if not self._cache_hit_is_valid(task, cached_results): + logging.info( + "Cache hit for %s with key: %s is no longer valid; re-executing", + self.name, + cache_key, + ) + cached_results = None + + if cached_results is not None: + logging.info("Cache hit for %s with key: %s", self.name, cache_key) + self._publish_cached_results(cached_results, task) + else: + logging.info("Cache miss for %s with key: %s", self.name, cache_key) + self.consume_work(task) + input_task, outputs = self._local.ctx.get_input_and_outputs() + # strip private fields from outputs + outputs = [ + [consumer.name, task.copy_public()] for consumer, task in outputs + ] + # The key is deliberately computed again inside _set_cache, after + # consume_work: extra_cache_key() may depend on state the worker just + # changed (e.g. files it edits in place), and the entry must be found + # by a later run that sees that state, which is exactly when replaying + # the outputs without re-running is valid. + self._set_cache(input_task, outputs) + + self.post_consume_work(task) + + def _lookup_cache_key(self, task: Task) -> str: + """ + The cache key that was used to look up ``task`` on this thread, computed once + per task. Falls back to computing it when called outside of task consumption. + + Args: + task (Task): The task being consumed. + + Returns: + str: The cache key. + """ + memo = getattr(self._local, "lookup_cache_key", None) + if memo is not None and memo[0] is task: + return memo[1] + return self._get_cache_key(task) def _cache_hit_is_valid( self, task: Task, cached_results: List[Tuple[str, Task]] diff --git a/src/planai/tools/filesystem.py b/src/planai/tools/filesystem.py index b774082..32db0ad 100644 --- a/src/planai/tools/filesystem.py +++ b/src/planai/tools/filesystem.py @@ -38,14 +38,15 @@ class Workspace: workspace root and are rejected with a :class:`ValueError` if they would escape that root, whether through an absolute path, a ``..`` component, or a symlink that points outside of the root. + + Constructing a Workspace has no side effects: the root directory is not + created until ``write_file`` first writes into it. """ def __init__(self, root: Union[str, "Path"]): - root_path = Path(root).expanduser() - root_path.mkdir(parents=True, exist_ok=True) # resolve() follows symlinks and normalizes the path so that every # subsequent comparison against self.root is done against the real path. - self.root: Path = root_path.resolve() + self.root: Path = Path(root).expanduser().resolve() def resolve(self, rel_path: str) -> Path: """Resolve a workspace-relative path to an absolute path inside the root. @@ -88,16 +89,29 @@ def resolve(self, rel_path: str) -> Path: def jailed_files(ws: Workspace, base: Path, pattern: str) -> List[Path]: """Files under ``base`` matching ``pattern`` whose real location is inside the workspace. Glob follows symlinks, so a link pointing outside the root would - otherwise be readable; such entries are skipped.""" + otherwise be readable; such entries are skipped. + + Raises: + ValueError: If the pattern is empty, absolute, or not supported by + :meth:`pathlib.Path.glob`. + """ + if not pattern: + raise ValueError("Glob pattern must not be empty") + if pattern.startswith(("/", "\\")) or Path(pattern).is_absolute(): + raise ValueError(f"Glob pattern must be relative: {pattern!r}") + files = [] - for path in base.glob(pattern): - if not path.is_file(): - continue - try: - ws.resolve(path.relative_to(ws.root).as_posix()) - except ValueError: - continue - files.append(path) + try: + for path in base.glob(pattern): + if not path.is_file(): + continue + try: + ws.resolve(path.relative_to(ws.root).as_posix()) + except ValueError: + continue + files.append(path) + except NotImplementedError as e: + raise ValueError(f"Unsupported glob pattern: {pattern!r}") from e return sorted(files) @@ -129,9 +143,15 @@ def hash_files(workspace: Union[Workspace, str, Path], globs: List[str]) -> str: digest = hashlib.sha1() for rel in sorted(matched.keys()): + try: + content = matched[rel].read_bytes() + except OSError as e: + raise OSError( + f"Could not read {rel} while hashing workspace files: {e.strerror or e}" + ) from e digest.update(rel.encode("utf-8")) digest.update(b"\x00") - digest.update(matched[rel].read_bytes()) + digest.update(content) digest.update(b"\x00") return digest.hexdigest() @@ -143,16 +163,18 @@ def make_file_tools( max_read_chars: int = 100_000, max_list_entries: int = 500, max_grep_matches: int = 200, + max_grep_line_chars: int = 500, ) -> List[Tool]: """Create llm_interface Tool objects bound to a single sandboxed workspace. Args: - workspace: The Workspace to jail all file operations to (or a path to one, - which will be created if missing). + workspace: The Workspace to jail all file operations to (or a path to one). read_only: If True, omit the write_file and edit_file tools. max_read_chars: Maximum number of characters read_file returns before truncating. max_list_entries: Maximum number of entries list_files returns before truncating. max_grep_matches: Maximum number of matches grep_files returns before truncating. + max_grep_line_chars: Maximum number of characters of a matching line that + grep_files includes before truncating it. Returns: List[Tool]: Tool objects usable as the ``tools`` argument to an LLMTaskWorker. @@ -223,7 +245,9 @@ def write_file(path: str, content: str) -> str: try: target.parent.mkdir(parents=True, exist_ok=True) - target.write_text(content, encoding="utf-8") + # write_bytes keeps the content verbatim; write_text would translate + # line endings on some platforms + target.write_bytes(content.encode("utf-8")) except OSError as e: return f"Error: Could not write file {path}: {e.strerror or e}" @@ -237,7 +261,8 @@ def edit_file( """Replace an exact snippet of text within an existing file. By default, old_string must occur exactly once in the file; use replace_all - to replace every occurrence instead. + to replace every occurrence instead. Line endings are matched as "\\n" and + a file that uses CRLF keeps CRLF. Args: path: Workspace-relative path to the file to edit. @@ -260,12 +285,17 @@ def edit_file( return "Error: old_string must not be empty" try: - text = target.read_text(encoding="utf-8") + raw = target.read_bytes().decode("utf-8") except UnicodeDecodeError: return f"Error: File is not valid UTF-8 text (binary?): {path}" except OSError as e: return f"Error: Could not read file {path}: {e.strerror or e}" + # match on "\n" so the model can quote the file as read_file showed it, + # but preserve the file's CRLF line endings when writing it back + uses_crlf = "\r\n" in raw + text = raw.replace("\r\n", "\n") if uses_crlf else raw + count = text.count(old_string) if count == 0: return f"Error: old_string not found in {path} (0 matches)" @@ -282,8 +312,11 @@ def edit_file( new_text = text.replace(old_string, new_string, 1) replacements = 1 + if uses_crlf: + new_text = new_text.replace("\n", "\r\n") + try: - target.write_text(new_text, encoding="utf-8") + target.write_bytes(new_text.encode("utf-8")) except OSError as e: return f"Error: Could not write file {path}: {e.strerror or e}" @@ -333,7 +366,7 @@ def list_files(path: str = ".", pattern: str = "**/*") -> str: output += f"\n\n[Output truncated at {max_list_entries} entries.]" return output - def grep_files(pattern: str, path: str = ".", glob: str = "**/*.md") -> str: + def grep_files(pattern: str, path: str = ".", glob: str = "**/*") -> str: """Search for a regular expression across text files in the workspace. Args: @@ -342,7 +375,8 @@ def grep_files(pattern: str, path: str = ".", glob: str = "**/*.md") -> str: path: Workspace-relative directory to search within. Use "." for the workspace root. glob: Glob pattern, relative to path, selecting which files to search, - e.g. "**/*.md" for every Markdown file recursively. + e.g. "**/*" for every file recursively or "**/*.md" for Markdown + files only. """ try: target = ws.resolve(path) @@ -375,6 +409,8 @@ def grep_files(pattern: str, path: str = ".", glob: str = "**/*.md") -> str: rel = file_path.relative_to(ws.root).as_posix() for lineno, line in enumerate(text.splitlines(), start=1): if regex.search(line): + if len(line) > max_grep_line_chars: + line = line[:max_grep_line_chars] + " [line truncated]" results.append(f"{rel}:{lineno}: {line}") if len(results) >= max_grep_matches: truncated = True diff --git a/src/planai/workspace_task.py b/src/planai/workspace_task.py index 7baf289..152cd9d 100644 --- a/src/planai/workspace_task.py +++ b/src/planai/workspace_task.py @@ -14,8 +14,10 @@ """A CachedLLMTaskWorker that lets an LLM read, write, and search files inside a per-job working directory carried through task provenance.""" +import uuid from typing import List, Optional, Tuple +from llm_interface import Tool from pydantic import Field from .llm_task import CachedLLMTaskWorker @@ -45,6 +47,12 @@ class WorkspaceLLMTaskWorker(CachedLLMTaskWorker): published output tasks and not any files the tools wrote on a prior run, subclasses can declare expected_output_files() so that a cache hit whose files are missing is treated as a cache miss and re-executed. + + The cache key includes the workspace root and, through input_globs, the + content of the declared input files. Workers that can write files pass a + fresh cache_salt to the LLM on every execution so that llm_interface's + response cache never replays an answer whose tool calls (the file writes) + would be skipped; read-only workers pass the cache key instead. """ read_only: bool = Field( @@ -97,22 +105,48 @@ def get_workspace(self, task: Task) -> Workspace: "string 'workspace' attribute) upstream" ) - def get_tools(self, task: Task): + def get_tools(self, task: Task) -> List[Tool]: file_tools = make_file_tools( self.get_workspace(task), read_only=self.read_only, max_read_chars=self.max_read_chars, ) - if self.tools: - return file_tools + list(self.tools) - return file_tools + if not self.tools: + return file_tools + + file_tool_names = {t.name for t in file_tools} + clashes = sorted(t.name for t in self.tools if t.name in file_tool_names) + if clashes: + raise ValueError( + f"{self.name}: static tools {clashes} would shadow the workspace " + "file tools of the same name" + ) + return file_tools + list(self.tools) def get_cache_salt(self, task: Task) -> Optional[str]: - return self._get_cache_key(task) + """ + Salt for llm_interface's response cache. Read-only workers use the cache + key, so an identical request can be served from the response cache. Workers + that can write files get a fresh value on every execution: a replayed + response skips the tool calls, so the files it describes would never be + written. + """ + cache_key = self._lookup_cache_key(task) + if self.read_only: + return cache_key + return f"{cache_key}:{uuid.uuid4().hex}" def extra_cache_key(self, task: Task) -> str: - workspace = self.get_workspace(task) - return hash_files(workspace, self.input_globs) + try: + workspace = self.get_workspace(task) + except ValueError: + # nothing to add; get_tools() raises a descriptive error when the + # worker actually runs + return "" + parts = [str(workspace.root)] + if self.input_globs: + parts.append(hash_files(workspace, self.input_globs)) + return ":".join(parts) def expected_output_files(self, task: Task) -> List[str]: """ diff --git a/tests/planai/test_workspace_task.py b/tests/planai/test_workspace_task.py index b64500f..0072e3f 100644 --- a/tests/planai/test_workspace_task.py +++ b/tests/planai/test_workspace_task.py @@ -103,6 +103,18 @@ def test_read_only_omits_write_and_edit_tools(self): self.assertNotIn("edit_file", tools) self.assertIn("read_file", tools) + def test_static_tool_shadowing_a_file_tool_is_rejected(self): + shadow = LLMToolInstance( + name="read_file", + description="not the jailed one", + parameters={"type": "object", "properties": {}, "required": []}, + func=lambda: "escaped", + ) + self.worker.tools = [shadow] + with self.assertRaises(ValueError) as ctx: + self.worker.get_tools(self.task) + self.assertIn("read_file", str(ctx.exception)) + def test_static_tools_are_appended(self): custom_tool = LLMToolInstance( name="custom_tool", @@ -146,12 +158,39 @@ def test_cache_key_stable_without_changes(self): key2 = self.worker._get_cache_key(self.task) self.assertEqual(key1, key2) - def test_get_cache_salt_matches_cache_key(self): + def test_cache_salt_is_the_cache_key_for_read_only_workers(self): + self.worker.read_only = True self.assertEqual( self.worker.get_cache_salt(self.task), self.worker._get_cache_key(self.task), ) + def test_cache_salt_is_fresh_per_execution_for_writers(self): + key = self.worker._get_cache_key(self.task) + salt1 = self.worker.get_cache_salt(self.task) + salt2 = self.worker.get_cache_salt(self.task) + self.assertTrue(salt1.startswith(key + ":")) + self.assertNotEqual(salt1, salt2) + + +class TestCacheKeyIdentity(WorkspaceWorkerTestCase): + def test_cache_key_differs_between_workspaces(self): + other_dir = tempfile.TemporaryDirectory() + self.addCleanup(other_dir.cleanup) + task_a = DummyTask(content="same payload") + add_input_provenance(task_a, WorkspaceTask(workspace=self.workspace_dir.name)) + task_b = DummyTask(content="same payload") + add_input_provenance(task_b, WorkspaceTask(workspace=other_dir.name)) + + self.assertNotEqual( + self.worker._get_cache_key(task_a), self.worker._get_cache_key(task_b) + ) + + def test_cache_key_without_workspace_does_not_raise(self): + task = DummyTask(content="orphan") + self.assertEqual(self.worker.extra_cache_key(task), "") + self.worker._get_cache_key(task) + class OutputTaskFile(Task): result: str @@ -201,6 +240,32 @@ def test_cache_hit_bypassed_when_expected_file_missing(self): self.llm.generate_pydantic.assert_called_once() + def test_rerun_after_invalid_hit_does_not_reuse_llm_response_cache(self): + self._seed_cache() + self.llm.generate_pydantic = Mock(return_value=OutputTaskFile(result="fresh")) + + with patch.object(self.worker, "_publish_cached_results"): + with patch("planai.llm_task.LLMTaskWorker.publish_work"): + self.worker._pre_consume_work(self.task) + + salt = self.llm.generate_pydantic.call_args.kwargs["cache_salt"] + key = self.worker._get_cache_key(self.task) + self.assertNotEqual(salt, key) + self.assertTrue(salt.startswith(key + ":")) + + def test_lookup_key_is_computed_once_per_execution(self): + self.llm.generate_pydantic = Mock(return_value=OutputTaskFile(result="fresh")) + + with patch.object( + self.worker, "_get_cache_key", wraps=self.worker._get_cache_key + ) as spy: + with patch("planai.llm_task.LLMTaskWorker.publish_work"): + self.worker._pre_consume_work(self.task) + + # once for the lookup and once, after consume_work, for the store; + # get_cache_salt reuses the lookup key instead of computing a third + self.assertEqual(spy.call_count, 2) + def test_cache_hit_honored_when_expected_file_present(self): self._seed_cache() (Path(self.workspace_dir.name) / "out.txt").write_text("done") diff --git a/tests/planai/tools/test_filesystem.py b/tests/planai/tools/test_filesystem.py index 854294d..9ca4564 100644 --- a/tests/planai/tools/test_filesystem.py +++ b/tests/planai/tools/test_filesystem.py @@ -1,5 +1,6 @@ # test_filesystem.py +import os import tempfile import unittest from pathlib import Path @@ -13,9 +14,16 @@ def setUp(self): self.addCleanup(self.tempdir.cleanup) self.workspace = Workspace(self.tempdir.name) - def test_root_is_created_and_resolved(self): - self.assertTrue(self.workspace.root.exists()) - self.assertTrue(self.workspace.root.is_absolute()) + def test_root_is_resolved_but_not_created(self): + missing = Path(self.tempdir.name) / "not-yet" + ws = Workspace(missing) + self.assertTrue(ws.root.is_absolute()) + self.assertFalse(ws.root.exists()) + + tools = {t.name: t for t in make_file_tools(ws)} + result = tools["write_file"].execute(path="a.txt", content="hi") + self.assertNotIn("Error", result) + self.assertEqual((missing / "a.txt").read_text(), "hi") def test_nested_path_accepted(self): resolved = self.workspace.resolve("a/b/c.txt") @@ -141,6 +149,12 @@ def test_omitted_when_read_only(self): self.assertIn("list_files", tools) self.assertIn("grep_files", tools) + def test_content_written_verbatim(self): + self.tools["write_file"].execute(path="win.txt", content="a\r\nb\n") + self.assertEqual( + (Path(self.tempdir.name) / "win.txt").read_bytes(), b"a\r\nb\n" + ) + class TestEditFile(FileToolsTestCase): def test_edit_unique_match(self): @@ -185,6 +199,31 @@ def test_omitted_when_read_only(self): self.assertNotIn("edit_file", tools) +class TestEditFileLineEndings(FileToolsTestCase): + def test_crlf_preserved(self): + full = Path(self.tempdir.name) / "win.txt" + full.write_bytes(b"hello\r\nworld\r\n") + result = self.tools["edit_file"].execute( + path="win.txt", old_string="world", new_string="there" + ) + self.assertNotIn("Error", result) + self.assertEqual(full.read_bytes(), b"hello\r\nthere\r\n") + + def test_multiline_old_string_matches_crlf_file(self): + full = Path(self.tempdir.name) / "win.txt" + full.write_bytes(b"a\r\nb\r\nc\r\n") + result = self.tools["edit_file"].execute( + path="win.txt", old_string="a\nb", new_string="ab" + ) + self.assertNotIn("Error", result) + self.assertEqual(full.read_bytes(), b"ab\r\nc\r\n") + + def test_lf_file_stays_lf(self): + full = self.write("unix.txt", "a\nb\n") + self.tools["edit_file"].execute(path="unix.txt", old_string="b", new_string="c") + self.assertEqual(full.read_bytes(), b"a\nc\n") + + class TestListFiles(FileToolsTestCase): def test_format_and_sorting(self): self.write("b.txt", "22") @@ -230,13 +269,28 @@ def test_format(self): result = self.tools["grep_files"].execute(pattern="foo") self.assertEqual(result, "a.md:2: foo bar") + def test_default_glob_searches_every_file(self): + self.write("a.md", "needle\n") + self.write("src/b.py", "needle\n") + result = self.tools["grep_files"].execute(pattern="needle") + self.assertIn("a.md", result) + self.assertIn("src/b.py", result) + def test_glob_filters_files(self): self.write("a.md", "needle\n") self.write("b.txt", "needle\n") - result = self.tools["grep_files"].execute(pattern="needle") + result = self.tools["grep_files"].execute(pattern="needle", glob="*.md") self.assertIn("a.md", result) self.assertNotIn("b.txt", result) + def test_long_lines_are_truncated(self): + self.write("a.md", "x" * 30 + "needle" + "y" * 30 + "\n") + tools = { + t.name: t for t in make_file_tools(self.workspace, max_grep_line_chars=20) + } + result = tools["grep_files"].execute(pattern="needle") + self.assertEqual(result, "a.md:1: " + "x" * 20 + " [line truncated]") + def test_invalid_regex(self): result = self.tools["grep_files"].execute(pattern="(unclosed") self.assertTrue(result.startswith("Error:")) @@ -255,6 +309,24 @@ def test_cap(self): self.assertEqual(len(match_lines), 5) +class TestAbsolutePatternsRejected(FileToolsTestCase): + def test_list_files(self): + result = self.tools["list_files"].execute(pattern="/etc/*") + self.assertTrue(result.startswith("Error:"), result) + + def test_grep_files(self): + result = self.tools["grep_files"].execute(pattern="root", glob="/etc/*") + self.assertTrue(result.startswith("Error:"), result) + + def test_hash_files(self): + with self.assertRaises(ValueError): + hash_files(self.workspace, ["/etc/*"]) + + def test_empty_pattern(self): + result = self.tools["list_files"].execute(pattern="") + self.assertTrue(result.startswith("Error:"), result) + + class TestHashFiles(unittest.TestCase): def setUp(self): self.tempdir = tempfile.TemporaryDirectory() @@ -267,6 +339,19 @@ def write(self, rel_path, content): full.write_text(content) return full + def test_unreadable_file_raises_clear_error(self): + if os.geteuid() == 0: + self.skipTest("root can read files regardless of permissions") + secret = self.write("secret.txt", "x") + secret.chmod(0) + self.addCleanup(secret.chmod, 0o600) + with self.assertRaises(OSError) as ctx: + hash_files(self.workspace, ["*.txt"]) + self.assertIn("secret.txt", str(ctx.exception)) + + def test_missing_root_hashes_to_empty(self): + self.assertEqual(hash_files(Path(self.tempdir.name) / "nope", ["*"]), "") + def test_empty_when_no_match(self): self.assertEqual(hash_files(self.workspace, ["**/*.md"]), "") From 3f447df02f5870ca3d0e556aa071fe57301ad33f Mon Sep 17 00:00:00 2001 From: provos Date: Sat, 5 Sep 2026 22:40:16 -0700 Subject: [PATCH 8/8] fix: avoid a race in Dispatcher.stop() when graphs share a dispatcher Two graphs sharing one dispatcher stop it concurrently; one caller set _dispatch_thread to None between the other's None check and its is_alive() call, which made test_shared_dispatcher_shutdown flaky. Hold the thread in a local before joining it. Co-Authored-By: Claude Fable 5.1 --- src/planai/dispatcher.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/src/planai/dispatcher.py b/src/planai/dispatcher.py index 7793388..7eca968 100644 --- a/src/planai/dispatcher.py +++ b/src/planai/dispatcher.py @@ -702,9 +702,12 @@ def stop(self, timeout: float = None): timeout (float, optional): Maximum time to wait for thread completion in seconds """ self.stop_event.set() - if self._dispatch_thread: - self._dispatch_thread.join(timeout=timeout) - if self._dispatch_thread.is_alive(): + # graphs sharing a dispatcher may call stop() concurrently; hold the + # thread locally so another caller clearing the attribute cannot race us + thread = self._dispatch_thread + if thread: + thread.join(timeout=timeout) + if thread.is_alive(): logging.warning("Dispatcher thread did not stop within timeout") else: logging.info("Dispatcher thread stopped")