diff --git a/amplifier_foundation/bundle/_prepared.py b/amplifier_foundation/bundle/_prepared.py index 04fee88a..24e32d32 100644 --- a/amplifier_foundation/bundle/_prepared.py +++ b/amplifier_foundation/bundle/_prepared.py @@ -424,6 +424,7 @@ def _create_system_prompt_factory( from amplifier_foundation.mentions import ContentDeduplicator from amplifier_foundation.mentions import format_context_block from amplifier_foundation.mentions import load_mentions + from amplifier_foundation.mentions import load_mentions_from_file # Capture state for the closure captured_bundle = bundle @@ -486,9 +487,7 @@ async def factory() -> str: # Add to deduplicator and mention_to_path for unified formatting for context_name, context_path in captured_bundle.context.items(): if context_path.exists(): - content = context_path.read_text(encoding="utf-8") - # Add to deduplicator for content-based deduplication - deduplicator.add_file(context_path, content) + await load_mentions_from_file(context_path, resolver, deduplicator) # Add to mention_to_path for attribution (context_name → path) mention_to_path[context_name] = context_path diff --git a/amplifier_foundation/mentions/__init__.py b/amplifier_foundation/mentions/__init__.py index 5119fe86..a7ef6561 100644 --- a/amplifier_foundation/mentions/__init__.py +++ b/amplifier_foundation/mentions/__init__.py @@ -1,25 +1,29 @@ """@mention parsing and loading utilities.""" from .deduplicator import ContentDeduplicator -from .loader import expand_mentions_in_instruction -from .loader import format_context_block -from .loader import load_mentions -from .models import ContextFile -from .models import MentionResult +from .loader import ( + expand_mentions_in_instruction, + format_context_block, + load_mentions, + load_mentions_from_file, +) +from .models import ContextFile, MentionResult from .parser import parse_mentions -from .protocol import MentionResolverProtocol +from .protocol import MentionResolverProtocol, RelativeMentionResolverProtocol from .resolver import BaseMentionResolver from .utils import format_directory_listing __all__ = [ - "parse_mentions", - "load_mentions", - "expand_mentions_in_instruction", - "format_context_block", - "format_directory_listing", + "BaseMentionResolver", "ContentDeduplicator", "ContextFile", - "MentionResult", "MentionResolverProtocol", - "BaseMentionResolver", + "MentionResult", + "RelativeMentionResolverProtocol", + "expand_mentions_in_instruction", + "format_context_block", + "format_directory_listing", + "load_mentions", + "load_mentions_from_file", + "parse_mentions", ] diff --git a/amplifier_foundation/mentions/loader.py b/amplifier_foundation/mentions/loader.py index 157dfcbd..8cffd35a 100644 --- a/amplifier_foundation/mentions/loader.py +++ b/amplifier_foundation/mentions/loader.py @@ -9,7 +9,7 @@ from .deduplicator import ContentDeduplicator from .models import MentionResult from .parser import parse_mentions -from .protocol import MentionResolverProtocol +from .protocol import MentionResolverProtocol, RelativeMentionResolverProtocol from .utils import format_directory_listing @@ -91,7 +91,10 @@ async def load_mentions( text: Text containing @mentions. resolver: Resolver to convert mentions to paths. deduplicator: Optional deduplicator for content. If None, creates one. - relative_to: Base path for relative mentions (defaults to cwd). + relative_to: Base path for local mentions (defaults to the resolver's + base). Nested explicit ./ and ../ mentions use the referring file's + directory with RelativeMentionResolverProtocol; bare local mentions + retain the resolver's base. max_depth: Maximum recursion depth to prevent infinite loops (default 3). Returns: @@ -101,6 +104,7 @@ async def load_mentions( deduplicator = ContentDeduplicator() results: list[MentionResult] = [] + visited_paths: set[Path] = set() mentions = parse_mentions(text) for mention in mentions: @@ -111,6 +115,7 @@ async def load_mentions( relative_to=relative_to, max_depth=max_depth, current_depth=0, + visited_paths=visited_paths, ) results.append(result) @@ -138,7 +143,10 @@ async def expand_mentions_in_instruction( instruction: The text to expand. May contain @mention tokens. resolver: Resolver to convert @mentions to file paths. deduplicator: Optional deduplicator for content. If None, creates a fresh one. - relative_to: Base path for relative mentions (defaults to cwd). + relative_to: Base path for local mentions (defaults to the resolver's + base). Nested explicit ./ and ../ mentions use the referring file's + directory with RelativeMentionResolverProtocol; bare local mentions + retain the resolver's base. Returns: Instruction with blocks prepended, or the original instruction @@ -172,6 +180,31 @@ async def expand_mentions_in_instruction( return f"{block}\n\n{instruction}" +async def load_mentions_from_file( + path: Path, + resolver: MentionResolverProtocol, + deduplicator: ContentDeduplicator | None = None, + max_depth: int = 3, +) -> MentionResult: + """Load a declared context file and its nested mentions using the same rules. + + Taking a Path avoids reparsing a known filename as mention syntax (including + spaces or Windows drive letters). Explicit relative references use this + file's directory; bare references retain the resolver's workspace root. + """ + return await _load_file( + mention=str(path), + path=path, + resolver=resolver, + deduplicator=deduplicator + if deduplicator is not None + else ContentDeduplicator(), + max_depth=max_depth, + current_depth=0, + visited_paths=set(), + ) + + async def _resolve_mention( mention: str, resolver: MentionResolverProtocol, @@ -179,10 +212,18 @@ async def _resolve_mention( relative_to: Path | None, max_depth: int, current_depth: int, + visited_paths: set[Path], ) -> MentionResult: """Resolve a single mention and recursively load its mentions.""" # Resolve mention to path - path = resolver.resolve(mention) + if ( + relative_to is not None + and (current_depth == 0 or mention.startswith(("@./", "@../"))) + and isinstance(resolver, RelativeMentionResolverProtocol) + ): + path = resolver.resolve_relative(mention, relative_to) + else: + path = resolver.resolve(mention) if path is None: return MentionResult( mention=mention, @@ -192,6 +233,22 @@ async def _resolve_mention( failure_reason="not_found", ) + return await _load_file( + mention, path, resolver, deduplicator, max_depth, current_depth, visited_paths + ) + + +async def _load_file( + mention: str, + path: Path, + resolver: MentionResolverProtocol, + deduplicator: ContentDeduplicator, + max_depth: int, + current_depth: int, + visited_paths: set[Path], +) -> MentionResult: + """Shared read/recursion path for resolved mentions and declared context.""" + # Handle directories: generate listing as content if path.is_dir(): try: @@ -223,6 +280,15 @@ async def _resolve_mention( failure_reason="not_found", ) + # Path identity stops cycles, while content identity deduplicates output. + # Identical files in different directories may reference different siblings. + canonical_path = path.resolve() + if canonical_path in visited_paths: + return MentionResult( + mention=mention, resolved_path=path, content=None, error=None + ) + visited_paths.add(canonical_path) + # Read file try: content = await read_with_retry(path) @@ -244,13 +310,7 @@ async def _resolve_mention( ) # Check for duplicate content - if not deduplicator.add_file(path, content): - return MentionResult( - mention=mention, - resolved_path=path, - content=None, # Already seen, don't include again - error=None, - ) + new_content = deduplicator.add_file(path, content) # Recursively load mentions from this file (if not at max depth) if current_depth < max_depth: @@ -263,11 +323,12 @@ async def _resolve_mention( relative_to=path.parent, max_depth=max_depth, current_depth=current_depth + 1, + visited_paths=visited_paths, ) return MentionResult( mention=mention, resolved_path=path, - content=content, + content=content if new_content else None, error=None, ) diff --git a/amplifier_foundation/mentions/protocol.py b/amplifier_foundation/mentions/protocol.py index 42ef90e4..e833a8e0 100644 --- a/amplifier_foundation/mentions/protocol.py +++ b/amplifier_foundation/mentions/protocol.py @@ -3,7 +3,7 @@ from __future__ import annotations from pathlib import Path -from typing import Protocol +from typing import Protocol, runtime_checkable class MentionResolverProtocol(Protocol): @@ -23,3 +23,17 @@ def resolve(self, mention: str) -> Path | None: Path to the resolved file, or None if not found. """ ... + + +@runtime_checkable +class RelativeMentionResolverProtocol(Protocol): + """Optional per-call context for resolvers used by the recursive loader. + + Legacy ``resolve(mention)`` implementations remain supported. Implement this + extension to anchor local paths to the referring file without changing the + resolver's workspace, namespace roots, or state between calls. + """ + + def resolve_relative(self, mention: str, relative_to: Path) -> Path | None: + """Resolve local paths relative to ``relative_to``; retain shortcut roots.""" + ... diff --git a/amplifier_foundation/mentions/resolver.py b/amplifier_foundation/mentions/resolver.py index d412b1c4..ac205dcf 100644 --- a/amplifier_foundation/mentions/resolver.py +++ b/amplifier_foundation/mentions/resolver.py @@ -2,6 +2,7 @@ from __future__ import annotations +from copy import copy from pathlib import Path from typing import TYPE_CHECKING @@ -51,7 +52,7 @@ def resolve(self, mention: str) -> Path | None: mention_body = mention[1:] # Remove @ prefix # Pattern 1: @bundle-name:context-name - if ":" in mention_body: + if ":" in mention_body and not Path(mention_body).is_absolute(): namespace, name = mention_body.split(":", 1) if bundle := self.bundles.get(namespace): return bundle.resolve_context_path(name) @@ -84,3 +85,13 @@ def register_bundle(self, name: str, bundle: Bundle) -> None: bundle: Bundle instance. """ self.bundles[name] = bundle + + def resolve_relative(self, mention: str, relative_to: Path) -> Path | None: + """Resolve local mentions from a referring file, without shared mutation. + + Use ``resolve`` on a scoped copy so subclasses retain their resolution + policy. Home paths and bundle namespaces keep their explicit roots. + """ + scoped = copy(self) + scoped.base_path = relative_to + return scoped.resolve(mention) diff --git a/docs/API_REFERENCE.md b/docs/API_REFERENCE.md index 1ee94b80..4367d870 100644 --- a/docs/API_REFERENCE.md +++ b/docs/API_REFERENCE.md @@ -198,3 +198,26 @@ from amplifier_foundation import load_mentions, BaseMentionResolver resolver = BaseMentionResolver(bundles={"foundation": foundation_bundle}) results = await load_mentions("See @foundation:context/guidelines.md", resolver) ``` + +Local mentions in the initial text use the resolver's base directory, or the +explicit `relative_to` passed to `load_mentions`. Inside an included file, explicit +relative mentions (`@./journal.md` or `@../rules.md`) use that file's directory. +Bare local mentions such as `@AGENTS.md` retain the resolver's workspace root; +this preserves bundles that intentionally include the current project's rules. +Bundle namespaces and explicit home/absolute paths retain their own roots. +The resolver is never mutated while loading a nested file. Content is included +once, but identical instruction files in different directories still load their +own relative references. Missing files remain optional; recursion depth and +canonical-path cycle detection bound traversal. + +Bundle-declared `context:` files use the same recursive loading rules. +`load_mentions_from_file(path, resolver, deduplicator)` exposes that path-based +entry point without reparsing filenames as mention syntax. This affects context +assembly, not ordinary file-tool results or attachments, whose content remains +literal. + +Custom resolvers keep the existing `resolve(mention)` contract. To support +per-file relative resolution, also implement the optional +`RelativeMentionResolverProtocol.resolve_relative(mention, relative_to)` method. +App shortcuts such as `@user:` and `@project:` should retain their configured +roots. Legacy resolvers without that method continue to resolve exactly as before. diff --git a/tests/test_relative_mentions.py b/tests/test_relative_mentions.py new file mode 100644 index 00000000..d50afc46 --- /dev/null +++ b/tests/test_relative_mentions.py @@ -0,0 +1,248 @@ +"""Nested instructions resolve beside their source, without changing roots.""" + +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from amplifier_foundation.bundle import Bundle, BundleModuleResolver, PreparedBundle +from amplifier_foundation.mentions import ( + BaseMentionResolver, + ContentDeduplicator, + load_mentions, + load_mentions_from_file, +) + + +def write(path: Path, text: str) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(text, encoding="utf-8") + return path + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mention", ["@./tasks.md", "@./tasks"]) +async def test_nested_reference_uses_referring_file_not_workspace(tmp_path, mention): + write(tmp_path / "rules/AGENTS.md", f"Follow {mention} and @workspace.md") + write(tmp_path / "workspace.md", "WORKSPACE RULE") + write(tmp_path / "rules/tasks.md", "NESTED TASKS RULE") + write(tmp_path / "tasks.md", "WRONG WORKSPACE RULE") + resolver = BaseMentionResolver(base_path=tmp_path) + dedup = ContentDeduplicator() + await load_mentions("@rules/AGENTS.md", resolver, dedup) + contents = [entry.content for entry in dedup.get_unique_files()] + assert "NESTED TASKS RULE" in contents + assert "WORKSPACE RULE" in contents + assert "WRONG WORKSPACE RULE" not in contents + assert resolver.base_path == tmp_path + assert resolver.resolve("@tasks.md") == tmp_path / "tasks.md" + + +@pytest.mark.asyncio +async def test_same_content_different_roots_preserves_both_nested_files(tmp_path): + for name in ("global", "project"): + write(tmp_path / name / "AGENTS.md", "Read @./tasks.md") + write(tmp_path / name / "tasks.md", f"{name} tasks rule @./AGENTS.md") + dedup = ContentDeduplicator() + await load_mentions( + "@global/AGENTS.md @project/AGENTS.md", + BaseMentionResolver(base_path=tmp_path), + dedup, + ) + entries = dedup.get_unique_files() + assert len(entries) == 3 + assert len(entries[0].paths) == 2 + assert {entry.content for entry in entries} == { + "Read @./tasks.md", + "global tasks rule @./AGENTS.md", + "project tasks rule @./AGENTS.md", + } + + +@pytest.mark.asyncio +async def test_explicit_base_and_nested_namespace_home_absolute_paths( + tmp_path, monkeypatch +): + home = tmp_path / "home" + original_expanduser = Path.expanduser + monkeypatch.setattr( + Path, + "expanduser", + lambda self: ( + home.joinpath(*self.parts[1:]) + if self.parts and self.parts[0] == "~" + else original_expanduser(self) + ), + ) + write(home / "rules.md", "HOME RULE") + absolute = write(tmp_path / "absolute.md", "ABSOLUTE RULE") + write(tmp_path / "bundle/rules.md", "BUNDLE RULE @./child.md") + write(tmp_path / "bundle/child.md", "BUNDLE CHILD") + write( + tmp_path / "workspace/AGENTS.md", + f"@~/rules.md @{absolute.as_posix()} @bundle:rules.md @missing:rules.md @missing.md", + ) + bundle = Bundle(name="bundle", base_path=tmp_path / "bundle") + resolver = BaseMentionResolver( + bundles={"bundle": bundle}, base_path=tmp_path / "wrong" + ) + dedup = ContentDeduplicator() + await load_mentions( + "@AGENTS.md", resolver, dedup, relative_to=tmp_path / "workspace" + ) + contents = [entry.content for entry in dedup.get_unique_files()] + assert all( + value in contents + for value in ( + "HOME RULE", + "ABSOLUTE RULE", + "BUNDLE RULE @./child.md", + "BUNDLE CHILD", + ) + ) + assert resolver.base_path == tmp_path / "wrong" + + +@pytest.mark.asyncio +async def test_recursion_depth_and_legacy_resolver_contract(tmp_path): + first = write(tmp_path / "first.md", "@second.md") + second = write(tmp_path / "second.md", "@third.md") + third = write(tmp_path / "third.md", "THIRD") + + class LegacyResolver: + def resolve(self, mention): + return {"@first.md": first, "@second.md": second, "@third.md": third}.get( + mention + ) + + dedup = ContentDeduplicator() + await load_mentions( + "@first.md", + LegacyResolver(), + dedup, + relative_to=tmp_path / "elsewhere", + max_depth=1, + ) + assert [entry.content for entry in dedup.get_unique_files()] == [ + "@second.md", + "@third.md", + ] + + +def test_scoped_resolution_keeps_subclass_policy(tmp_path): + write(tmp_path / "rules/private.md", "PRIVATE") + + class RestrictedResolver(BaseMentionResolver): + def resolve(self, mention): + return None if "private" in mention else super().resolve(mention) + + resolver = RestrictedResolver(base_path=tmp_path) + assert resolver.resolve_relative("@private.md", tmp_path / "rules") is None + + +@pytest.mark.asyncio +async def test_explicit_parent_reference_and_cycle_are_bounded(tmp_path): + write(tmp_path / "rules/sub/AGENTS.md", "@../tasks.md") + write(tmp_path / "rules/tasks.md", "RIGHT RULE @./sub/AGENTS.md") + write(tmp_path / "tasks.md", "WRONG RULE") + dedup = ContentDeduplicator() + await load_mentions( + "@rules/sub/AGENTS.md", + BaseMentionResolver(base_path=tmp_path), + dedup, + max_depth=100, + ) + assert [entry.content for entry in dedup.get_unique_files()] == [ + "@../tasks.md", + "RIGHT RULE @./sub/AGENTS.md", + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bundle_name", ["anchors-fixture", "work-fixture"]) +async def test_tasks_rules_refresh_in_real_prompt_factory( + tmp_path, monkeypatch, bundle_name +): + home, workspace, cache = ( + tmp_path / name for name in ("home", "workspace", "cache") + ) + original_expanduser = Path.expanduser + monkeypatch.setattr( + Path, + "expanduser", + lambda self: ( + home.joinpath(*self.parts[1:]) + if self.parts and self.parts[0] == "~" + else original_expanduser(self) + ), + ) + for root, label in ( + (home / ".amplifier", "GLOBAL"), + (workspace / ".amplifier", "PROJECT"), + ): + write(root / "AGENTS.md", "Read @./rules/tasks.md") + write(root / "rules/tasks.md", f"{label} TASKS RULE") + write(workspace / "rules/tasks.md", "WRONG WORKSPACE RULE") + write(workspace / "AGENTS.md", "BARE WORKSPACE RULE") + write(cache / "context/system.md", "BUNDLE RULE @AGENTS.md") + bundle = Bundle( + name=bundle_name, + base_path=cache, + instruction=f"@{bundle_name}:context/system.md\n@~/.amplifier/AGENTS.md\n@.amplifier/AGENTS.md", + ) + prepared = PreparedBundle({}, BundleModuleResolver({}), bundle) + session = SimpleNamespace( + coordinator=SimpleNamespace(hooks=SimpleNamespace(emit=AsyncMock())) + ) + render = prepared.create_system_prompt_factory(session, session_cwd=workspace) + for factory in ( + render, + render, + prepared.create_system_prompt_factory(session, session_cwd=workspace), + ): + prompt = await factory() + for label in ( + "GLOBAL TASKS RULE", + "PROJECT TASKS RULE", + "BARE WORKSPACE RULE", + ): + assert prompt.count(label) == 1 + assert "WRONG WORKSPACE RULE" not in prompt + write(home / ".amplifier/rules/tasks.md", "UPDATED TASKS RULE") + assert "UPDATED TASKS RULE" in await render() + assert "GLOBAL TASKS RULE" not in await render() + (workspace / ".amplifier/rules/tasks.md").unlink() + assert "PROJECT TASKS RULE" not in await render() + + +@pytest.mark.asyncio +async def test_declared_context_files_share_recursive_loading_and_roots(tmp_path): + workspace, cache = tmp_path / "workspace", tmp_path / "cache with spaces" + write(workspace / "AGENTS.md", "WORKSPACE RULE") + parent = write(cache / "context/rules.md", "@./tasks.md @AGENTS.md") + write(cache / "context/tasks.md", "INCLUDED TASK RULE @./rules.md") + bundle = Bundle( + name="fixture", base_path=cache, instruction="Root", context={"rules": parent} + ) + prepared = PreparedBundle({}, BundleModuleResolver({}), bundle) + session = SimpleNamespace( + coordinator=SimpleNamespace(hooks=SimpleNamespace(emit=AsyncMock())) + ) + render = prepared.create_system_prompt_factory(session, session_cwd=workspace) + for factory in ( + render, + prepared.create_system_prompt_factory(session, session_cwd=workspace), + ): + prompt = await factory() + assert prompt.count("INCLUDED TASK RULE") == 1 + assert prompt.count("WORKSPACE RULE") == 1 + assert str(parent) in prompt + write(cache / "context/tasks.md", "EDITED INCLUDED RULE") + assert "EDITED INCLUDED RULE" in await render() + assert "INCLUDED TASK RULE" not in await render() + dedup = ContentDeduplicator() + await load_mentions_from_file( + parent, BaseMentionResolver(base_path=workspace), dedup + ) + assert len(dedup.get_unique_files()) == 3