Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion amplifier_foundation/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,7 @@ def __init__(
*,
strict: bool = False,
include_source_resolver: Callable[[str], str | None] | None = None,
persist: bool = True,
) -> None:
"""Initialize registry.

Expand All @@ -198,9 +199,15 @@ def __init__(
source URIs. When provided, called before default resolution
logic. Returns a resolved URI string or None to fall back to
default behavior.
persist: If False, read the shared registrations but keep all changes
in this registry instance. Loads, include tracking, stale-path
cleanup and explicit save() calls cannot write registry.json.
Source downloads still use the shared content cache. Use a
fresh instance per session with scoped source overrides.
"""
self._home = self._resolve_home(home)
self._strict = strict
self._persist = persist
self._include_source_resolver = include_source_resolver
self._registry: dict[str, BundleState] = {}
self._source_resolver = SimpleSourceResolver(
Expand Down Expand Up @@ -1399,7 +1406,9 @@ def get_state(
# =========================================================================

def save(self) -> None:
"""Persist registry state to home/registry.json."""
"""Persist state, unless this is an isolated in-memory registry view."""
if not self._persist:
return
self._home.mkdir(parents=True, exist_ok=True)
registry_path = self._home / "registry.json"

Expand Down
107 changes: 107 additions & 0 deletions tests/test_registry_isolation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
"""Session source selections must never become shared registrations."""

import asyncio

import pytest
import yaml

from amplifier_foundation.exceptions import BundleDependencyError
from amplifier_foundation.registry import BundleRegistry


def bundle(root, name="shared", *, broken=False):
root.mkdir()
child = root / "child.yaml"
child.write_text(yaml.safe_dump({
"bundle": {"name": "shared-child"},
"tools": [{"module": "tool-marker", "config": {"value": root.name}}],
}))
(root / "bundle.yaml").write_text(yaml.safe_dump({
"bundle": {"name": name},
"includes": ["shared:child.yaml", *([str(root / "missing.yaml")] if broken else [])],
}))
return root.as_uri()


@pytest.mark.asyncio
@pytest.mark.parametrize("concurrent", [False, True])
async def test_scoped_aliases_and_discovered_state_never_escape(tmp_path, concurrent):
original = bundle(tmp_path / "original")
variants = [bundle(tmp_path / label) for label in ("first", "second")]
home = tmp_path / "registry"
global_registry = BundleRegistry(home, strict=True)
global_registry.register({"shared": original})
await global_registry.load("shared")
before = (home / "registry.json").read_bytes()

async def load(uri):
scoped = BundleRegistry(home, strict=True, persist=False)
scoped.register({"shared": uri, original: uri})
await asyncio.sleep(0)
loaded = await scoped.load(original)
scoped.save() # A host's final save must not bypass isolation.
return loaded

loaded = (await asyncio.gather(*(load(uri) for uri in variants)) if concurrent
else [await load(uri) for uri in variants])
assert [row.tools[0]["config"]["value"] for row in loaded] == ["first", "second"]
assert (home / "registry.json").read_bytes() == before
fresh = BundleRegistry(home, strict=True, persist=False)
assert fresh.find(original) is None
assert fresh.find("shared") == original
assert (await fresh.load(original)).tools[0]["config"]["value"] == "original"
assert (home / "registry.json").read_bytes() == before
assert global_registry.find("shared") == original


@pytest.mark.asyncio
async def test_failed_composition_keeps_registry_bytes(tmp_path):
original = bundle(tmp_path / "original")
broken = bundle(tmp_path / "broken", broken=True)
home = tmp_path / "registry"
registry = BundleRegistry(home)
registry.register({"shared": original})
registry.save()
before = (home / "registry.json").read_bytes()
scoped = BundleRegistry(home, persist=False, strict=True)
scoped.register({original: broken})
with pytest.raises(BundleDependencyError):
await scoped.load(original)
scoped.save()
assert (home / "registry.json").read_bytes() == before


def test_stale_path_cleanup_is_local_and_no_file_is_created(tmp_path):
home = tmp_path / "registry"
registry = BundleRegistry(home)
registry.register({"missing": "file:///not-present"})
registry.get_state("missing").local_path = str(tmp_path / "absent")
registry.save()
before = (home / "registry.json").read_bytes()
scoped = BundleRegistry(home, persist=False)
assert scoped.get_state("missing").local_path is None
scoped.unregister("missing")
scoped.save()
assert (home / "registry.json").read_bytes() == before
empty = tmp_path / "empty"
scoped = BundleRegistry(empty, persist=False)
scoped.register({"temporary": "file:///temporary"})
scoped.save()
assert not (empty / "registry.json").exists()


@pytest.mark.asyncio
async def test_explicit_global_registration_and_update_still_persist(tmp_path):
old = bundle(tmp_path / "old")
new = bundle(tmp_path / "new")
home = tmp_path / "registry"
registry = BundleRegistry(home)
registry.register({"shared": old})
registry.save()
assert BundleRegistry(home).find("shared") == old
registry.register({"shared": new})
await registry.load("shared")
restored = BundleRegistry(home)
assert restored.find("shared") == new
assert restored.get_state("shared").includes == ["shared-child"]
assert restored.get_state("shared-child").uri == (tmp_path / "new/child.yaml").as_uri()
Loading