diff --git a/memu/neighborhood_lock.py b/memu/neighborhood_lock.py new file mode 100644 index 0000000..c8a0f83 --- /dev/null +++ b/memu/neighborhood_lock.py @@ -0,0 +1,586 @@ +"""Neighborhood locks for wiki-worker agents. + +When a ``wiki-worker`` starts editing a slug, it needs exclusive access to +that slug (and optionally a 1-hop neighborhood) so that two agents do not +clobber each other's writes. This module provides the lock registry shared +by both storage tiers plus the context manager agents use at runtime. + +Two registries ship in-tree: + +- :class:`SqliteLockRegistry` reuses the shared SQLite connection that the + wiki backend already opens (tables land alongside the nodes/slug tables + introduced by PR #21). It's the production path for anything that talks + to :class:`memu.storage.sqlite_backend.SqliteBackend`. +- :class:`MarkdownLockRegistry` keeps Tier 0 fully filesystem-native: locks + live in ``/.memu/locks/.json`` and a monotonic fencing + counter lives in ``/.memu/locks/_fence/``. + +Both expose the same :class:`LockRegistry` protocol so +:class:`NeighborhoodLock` can remain backend-agnostic. + +``NeighborhoodConflict`` subclasses the existing +:class:`memu.lane_lock.LaneContestedError`, so callers that already trap the +lane-lock hierarchy will catch wiki contention without new except clauses. +""" +from __future__ import annotations + +import asyncio +import json +import logging +import os +import sqlite3 +import uuid +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Iterable, Optional, Protocol, runtime_checkable + +from memu.lane_lock import LaneContestedError + +logger = logging.getLogger(__name__) + +DEFAULT_TTL_S: float = 30.0 +DEFAULT_RENEW_INTERVAL_S: float = 10.0 + + +class NeighborhoodConflict(LaneContestedError): + """Raised when a neighborhood lock cannot be acquired or kept.""" + + +@dataclass(frozen=True) +class LockRecord: + """Snapshot of a single slug's lock state.""" + + slug: str + holder_id: str + fencing_token: int + acquired_at: datetime + expires_at: datetime + + +@runtime_checkable +class LockRegistry(Protocol): + """Backend-agnostic interface for neighborhood locks.""" + + async def init(self) -> None: ... + + async def try_acquire( + self, slug: str, holder_id: str, ttl_s: float + ) -> LockRecord: + """Claim ``slug`` for ``holder_id``; raise :class:`NeighborhoodConflict` + if a live peer already holds it.""" + + async def renew( + self, slug: str, holder_id: str, fencing_token: int, ttl_s: float + ) -> LockRecord: + """Extend the lease. Must fail if someone stole the lock.""" + + async def release( + self, slug: str, holder_id: str, fencing_token: int + ) -> bool: + """Drop the lock if still owned; return True on success.""" + + async def peek(self, slug: str) -> Optional[LockRecord]: + """Return the current record without mutating state.""" + + async def close(self) -> None: ... + + +# --------------------------------------------------------------------------- +# SQLite registry +# --------------------------------------------------------------------------- + +_SQLITE_SCHEMA = """ +CREATE TABLE IF NOT EXISTS neighborhood_locks ( + slug TEXT PRIMARY KEY, + holder_id TEXT NOT NULL, + fencing_token INTEGER NOT NULL, + acquired_at TEXT NOT NULL, + expires_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS neighborhood_fence_tokens ( + slug TEXT PRIMARY KEY, + current_token INTEGER NOT NULL DEFAULT 0 +); +""" + + +class SqliteLockRegistry: + """Lock registry that piggybacks on an existing SQLite connection. + + The backend already holds the connection in WAL mode; we only mutate two + additional tables. All registry methods run under a single asyncio lock + to keep read-check-write sequences atomic inside the process. + """ + + def __init__(self, conn: sqlite3.Connection): + self._conn = conn + self._mu = asyncio.Lock() + + async def init(self) -> None: + self._conn.executescript(_SQLITE_SCHEMA) + self._conn.commit() + + async def try_acquire( + self, slug: str, holder_id: str, ttl_s: float + ) -> LockRecord: + async with self._mu: + now = datetime.now(timezone.utc) + row = self._conn.execute( + "SELECT holder_id, fencing_token, expires_at" + " FROM neighborhood_locks WHERE slug = ?", + (slug,), + ).fetchone() + if row is not None: + expires_at = _parse_iso(row[2]) + if expires_at > now: + raise NeighborhoodConflict( + f"slug {slug!r} held by {row[0]!r} until {row[2]}" + ) + # Lease expired -> reclaim by deleting; fence counter keeps + # advancing so any resumed writer is detectably stale. + self._conn.execute( + "DELETE FROM neighborhood_locks WHERE slug = ?", (slug,) + ) + + new_token = self._bump_fence(slug) + expires = now + timedelta(seconds=ttl_s) + self._conn.execute( + "INSERT INTO neighborhood_locks" + " (slug, holder_id, fencing_token, acquired_at, expires_at)" + " VALUES (?, ?, ?, ?, ?)", + ( + slug, + holder_id, + new_token, + now.isoformat(), + expires.isoformat(), + ), + ) + self._conn.commit() + return LockRecord(slug, holder_id, new_token, now, expires) + + async def renew( + self, slug: str, holder_id: str, fencing_token: int, ttl_s: float + ) -> LockRecord: + async with self._mu: + row = self._conn.execute( + "SELECT holder_id, fencing_token FROM neighborhood_locks" + " WHERE slug = ?", + (slug,), + ).fetchone() + if row is None: + raise NeighborhoodConflict( + f"slug {slug!r}: lock no longer present" + ) + if row[0] != holder_id or int(row[1]) != fencing_token: + raise NeighborhoodConflict( + f"slug {slug!r}: fencing token mismatch" + f" (held {fencing_token}, current {row[1]})" + ) + now = datetime.now(timezone.utc) + expires = now + timedelta(seconds=ttl_s) + self._conn.execute( + "UPDATE neighborhood_locks SET expires_at = ? WHERE slug = ?", + (expires.isoformat(), slug), + ) + self._conn.commit() + return LockRecord(slug, holder_id, fencing_token, now, expires) + + async def release( + self, slug: str, holder_id: str, fencing_token: int + ) -> bool: + async with self._mu: + row = self._conn.execute( + "SELECT holder_id, fencing_token FROM neighborhood_locks" + " WHERE slug = ?", + (slug,), + ).fetchone() + if row is None: + return False + if row[0] != holder_id or int(row[1]) != fencing_token: + return False + self._conn.execute( + "DELETE FROM neighborhood_locks WHERE slug = ?", (slug,) + ) + self._conn.commit() + return True + + async def peek(self, slug: str) -> Optional[LockRecord]: + row = self._conn.execute( + "SELECT slug, holder_id, fencing_token, acquired_at, expires_at" + " FROM neighborhood_locks WHERE slug = ?", + (slug,), + ).fetchone() + if row is None: + return None + return LockRecord( + slug=row[0], + holder_id=row[1], + fencing_token=int(row[2]), + acquired_at=_parse_iso(row[3]), + expires_at=_parse_iso(row[4]), + ) + + async def close(self) -> None: + # Connection is owned by the backend; nothing to release here. + return None + + def _bump_fence(self, slug: str) -> int: + row = self._conn.execute( + "SELECT current_token FROM neighborhood_fence_tokens" + " WHERE slug = ?", + (slug,), + ).fetchone() + new_token = int(row[0]) + 1 if row else 1 + self._conn.execute( + "INSERT INTO neighborhood_fence_tokens(slug, current_token)" + " VALUES (?, ?)" + " ON CONFLICT(slug) DO UPDATE SET current_token = excluded.current_token", + (slug, new_token), + ) + return new_token + + +# --------------------------------------------------------------------------- +# Markdown (sidecar-file) registry +# --------------------------------------------------------------------------- + + +class MarkdownLockRegistry: + """Filesystem-only registry. One sidecar file per lock. + + Layout:: + + /.memu/locks/.json lease payload + /.memu/locks/_fence/ monotonic fencing counter + + Atomicity hinges on two primitives: + + - ``os.O_CREAT | os.O_EXCL`` for lease creation so two concurrent acquire + attempts on the same filesystem can't both succeed. + - Write-to-tmp-then-``rename`` for fence/lease mutations so readers never + see torn payloads. + """ + + def __init__(self, root: os.PathLike[str] | str): + self.root = Path(root) + self._locks_dir = self.root / ".memu" / "locks" + self._fence_dir = self._locks_dir / "_fence" + self._mu = asyncio.Lock() + + async def init(self) -> None: + self._locks_dir.mkdir(parents=True, exist_ok=True) + self._fence_dir.mkdir(parents=True, exist_ok=True) + + async def try_acquire( + self, slug: str, holder_id: str, ttl_s: float + ) -> LockRecord: + async with self._mu: + now = datetime.now(timezone.utc) + expires = now + timedelta(seconds=ttl_s) + lock_path = self._lock_path(slug) + + # Clear an expired lease first so the O_EXCL create below can + # succeed on reclaim. + existing = self._read_lock_file(lock_path) + if existing is not None: + if _parse_iso(existing["expires_at"]) > now: + raise NeighborhoodConflict( + f"slug {slug!r} held by {existing['holder_id']!r}" + f" until {existing['expires_at']}" + ) + lock_path.unlink(missing_ok=True) + + new_token = self._bump_fence(slug) + payload = { + "slug": slug, + "holder_id": holder_id, + "fencing_token": new_token, + "acquired_at": now.isoformat(), + "expires_at": expires.isoformat(), + } + try: + fd = os.open( + str(lock_path), + os.O_CREAT | os.O_EXCL | os.O_WRONLY, + 0o644, + ) + except FileExistsError as e: + raise NeighborhoodConflict( + f"slug {slug!r}: lost race to acquire" + ) from e + try: + os.write(fd, json.dumps(payload).encode("utf-8")) + finally: + os.close(fd) + return LockRecord(slug, holder_id, new_token, now, expires) + + async def renew( + self, slug: str, holder_id: str, fencing_token: int, ttl_s: float + ) -> LockRecord: + async with self._mu: + self._assert_fence_matches(slug, fencing_token) + lock_path = self._lock_path(slug) + payload = self._read_lock_file(lock_path) + if payload is None: + raise NeighborhoodConflict( + f"slug {slug!r}: lock no longer present" + ) + if ( + payload.get("holder_id") != holder_id + or int(payload.get("fencing_token", -1)) != fencing_token + ): + raise NeighborhoodConflict( + f"slug {slug!r}: owner changed" + ) + now = datetime.now(timezone.utc) + expires = now + timedelta(seconds=ttl_s) + payload["expires_at"] = expires.isoformat() + self._atomic_write(lock_path, json.dumps(payload)) + return LockRecord(slug, holder_id, fencing_token, now, expires) + + async def release( + self, slug: str, holder_id: str, fencing_token: int + ) -> bool: + async with self._mu: + lock_path = self._lock_path(slug) + payload = self._read_lock_file(lock_path) + if payload is None: + return False + if ( + payload.get("holder_id") != holder_id + or int(payload.get("fencing_token", -1)) != fencing_token + ): + return False + # Honor fencing_token: the on-disk counter must still match the + # caller's token. If a reclaim already bumped it, the lease they + # hold is stale and release is a no-op. + if self._read_fence(slug) != fencing_token: + return False + lock_path.unlink(missing_ok=True) + return True + + async def peek(self, slug: str) -> Optional[LockRecord]: + lock_path = self._lock_path(slug) + payload = self._read_lock_file(lock_path) + if payload is None: + return None + return LockRecord( + slug=str(payload["slug"]), + holder_id=str(payload["holder_id"]), + fencing_token=int(payload["fencing_token"]), + acquired_at=_parse_iso(payload["acquired_at"]), + expires_at=_parse_iso(payload["expires_at"]), + ) + + async def close(self) -> None: + return None + + # --- internal helpers ------------------------------------------------- + + def _safe_name(self, slug: str) -> str: + # Slugs may include '/' (e.g., "paper/attention"). Collapse to a flat + # filename so the locks dir stays one level deep. + return slug.replace("/", "__") + + def _lock_path(self, slug: str) -> Path: + return self._locks_dir / f"{self._safe_name(slug)}.json" + + def _fence_path(self, slug: str) -> Path: + return self._fence_dir / self._safe_name(slug) + + def _read_lock_file(self, path: Path) -> Optional[dict]: + try: + raw = path.read_text(encoding="utf-8") + except FileNotFoundError: + return None + except OSError: + return None + try: + return json.loads(raw) + except json.JSONDecodeError: + return None + + def _read_fence(self, slug: str) -> int: + try: + return int(self._fence_path(slug).read_text().strip() or "0") + except FileNotFoundError: + return 0 + except (OSError, ValueError): + return 0 + + def _bump_fence(self, slug: str) -> int: + new_token = self._read_fence(slug) + 1 + self._atomic_write(self._fence_path(slug), str(new_token)) + return new_token + + def _assert_fence_matches(self, slug: str, expected: int) -> None: + current = self._read_fence(slug) + if current != expected: + raise NeighborhoodConflict( + f"slug {slug!r}: fencing token mismatch" + f" (held {expected}, current {current})" + ) + + def _atomic_write(self, path: Path, content: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(path.suffix + ".tmp") + tmp.write_text(content, encoding="utf-8") + os.replace(tmp, path) + + +# --------------------------------------------------------------------------- +# Context manager +# --------------------------------------------------------------------------- + + +class NeighborhoodLock: + """Async context manager that acquires a slug (plus optional 1-hop + neighbors) and keeps the lease alive via an async renewer task. + + Usage:: + + registry = backend.get_lock_registry() + await registry.init() + async with NeighborhoodLock(registry, "foo-design", neighbors=["bar"]): + # exclusive on foo-design + bar + ... + + Multiple slugs are acquired in a sorted order; if any one fails, every + lock already grabbed is released before the :class:`NeighborhoodConflict` + propagates, so callers never leak partial leases. + """ + + def __init__( + self, + registry: LockRegistry, + slug: str, + *, + neighbors: Iterable[str] = (), + holder_id: Optional[str] = None, + ttl_s: float = DEFAULT_TTL_S, + renew_interval_s: float = DEFAULT_RENEW_INTERVAL_S, + ): + self.registry = registry + self.slug = slug + self.neighbors = tuple(neighbors) + self.holder_id = holder_id or uuid.uuid4().hex + self.ttl_s = float(ttl_s) + self.renew_interval_s = float(renew_interval_s) + if self.renew_interval_s <= 0: + raise ValueError("renew_interval_s must be positive") + if self.renew_interval_s >= self.ttl_s: + raise ValueError("renew_interval_s must be < ttl_s") + self.records: list[LockRecord] = [] + self._renewer: Optional[asyncio.Task] = None + + @property + def record(self) -> LockRecord: + if not self.records: + raise RuntimeError("NeighborhoodLock is not currently held") + return self.records[0] + + @property + def fencing_token(self) -> int: + return self.record.fencing_token + + async def __aenter__(self) -> "NeighborhoodLock": + targets = [self.slug] + sorted( + s for s in self.neighbors if s != self.slug + ) + acquired: list[LockRecord] = [] + try: + for slug in targets: + rec = await self.registry.try_acquire( + slug, self.holder_id, self.ttl_s + ) + acquired.append(rec) + except BaseException: + for rec in acquired: + try: + await self.registry.release( + rec.slug, rec.holder_id, rec.fencing_token + ) + except Exception: + logger.exception( + "rollback release failed for %s", rec.slug + ) + raise + self.records = acquired + self._renewer = asyncio.create_task( + self._renew_loop(), name=f"nlock-renew:{self.slug}" + ) + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: + if self._renewer is not None: + self._renewer.cancel() + try: + await self._renewer + except asyncio.CancelledError: + pass + except Exception: + logger.exception("renewer task exited with error") + self._renewer = None + for rec in self.records: + try: + await self.registry.release( + rec.slug, rec.holder_id, rec.fencing_token + ) + except Exception: + logger.exception("release failed for %s", rec.slug) + self.records = [] + + async def _renew_loop(self) -> None: + try: + while True: + await asyncio.sleep(self.renew_interval_s) + renewed: list[LockRecord] = [] + for rec in self.records: + renewed.append( + await self.registry.renew( + rec.slug, + rec.holder_id, + rec.fencing_token, + self.ttl_s, + ) + ) + self.records = renewed + except asyncio.CancelledError: + raise + except NeighborhoodConflict: + # The lock was stolen; surface by letting the task die. Callers + # inside the ``async with`` body will notice on their next + # interaction with the registry (or when the body exits). + logger.warning( + "NeighborhoodLock renew lost ownership for %s", self.slug + ) + except Exception: + logger.exception("renew loop crashed for %s", self.slug) + + +# --------------------------------------------------------------------------- +# helpers +# --------------------------------------------------------------------------- + + +def _parse_iso(value: str) -> datetime: + """Parse ISO timestamps produced by datetime.isoformat(), tolerating + trailing ``Z`` notation.""" + if value.endswith("Z"): + value = value[:-1] + "+00:00" + return datetime.fromisoformat(value) + + +__all__ = [ + "DEFAULT_TTL_S", + "DEFAULT_RENEW_INTERVAL_S", + "LockRecord", + "LockRegistry", + "MarkdownLockRegistry", + "NeighborhoodConflict", + "NeighborhoodLock", + "SqliteLockRegistry", +] diff --git a/memu/storage/base.py b/memu/storage/base.py index c5cca65..dfe551f 100644 --- a/memu/storage/base.py +++ b/memu/storage/base.py @@ -286,4 +286,13 @@ async def put_embedding( # --- iteration (for sync/migration) --- async def iter_changed(self, since: Optional[datetime] = None) -> AsyncIterator[WikiNode]: ... + # --- concurrency --- + def get_lock_registry(self) -> "Any": + """Return the neighborhood lock registry bound to this backend. + + See :mod:`memu.neighborhood_lock` for the exact protocol. Declared as + ``Any`` here to avoid a circular import; implementations return + something that satisfies :class:`memu.neighborhood_lock.LockRegistry`. + """ + async def close(self) -> None: ... diff --git a/memu/storage/markdown_backend.py b/memu/storage/markdown_backend.py index d97b26b..ae59a41 100644 --- a/memu/storage/markdown_backend.py +++ b/memu/storage/markdown_backend.py @@ -245,6 +245,13 @@ async def iter_changed( async def close(self) -> None: # nothing to close return None + def get_lock_registry(self): + """Return a :class:`MarkdownLockRegistry` rooted at this vault. + Callers must ``await registry.init()`` before use.""" + from ..neighborhood_lock import MarkdownLockRegistry + + return MarkdownLockRegistry(root=self.layout.root) + # ---------------------------------------------------------------- helpers def _read(self, path: Path) -> Optional[WikiNode]: try: diff --git a/memu/storage/postgres_backend.py b/memu/storage/postgres_backend.py index 0c097e0..d3caf9a 100644 --- a/memu/storage/postgres_backend.py +++ b/memu/storage/postgres_backend.py @@ -72,5 +72,8 @@ async def iter_changed(self, since: Optional[datetime] = None) -> AsyncIterator[ if False: # pragma: no cover yield # type: ignore[misc] + def get_lock_registry(self): # pragma: no cover - stub + raise NotImplementedError + async def close(self) -> None: # pragma: no cover return None diff --git a/memu/storage/sqlite_backend.py b/memu/storage/sqlite_backend.py index 372bd10..98ecb68 100644 --- a/memu/storage/sqlite_backend.py +++ b/memu/storage/sqlite_backend.py @@ -150,6 +150,13 @@ def conn(self) -> sqlite3.Connection: raise RuntimeError("SqliteBackend.init() must be called before use") return self._conn + def get_lock_registry(self): + """Return a :class:`SqliteLockRegistry` sharing this backend's + connection. Callers must still ``await registry.init()`` before use.""" + from ..neighborhood_lock import SqliteLockRegistry + + return SqliteLockRegistry(conn=self.conn) + # ---- node CRUD async def put_node(self, node: WikiNode) -> WikiNode: now = datetime.now(timezone.utc) diff --git a/scripts/search_smoke.py b/scripts/search_smoke.py new file mode 100644 index 0000000..6ddcfa8 --- /dev/null +++ b/scripts/search_smoke.py @@ -0,0 +1,21 @@ +"""Manual smoke script: POST /search against a locally running memU API. + +Run with a live server on http://localhost:8000: + + python3 scripts/search_smoke.py +""" +import requests + + +def main() -> None: + res = requests.post( + "http://localhost:8000/search", + headers={"X-API-Key": "memu-dev-key"}, + json={"query": "test query", "limit": 2}, + ) + print("Status:", res.status_code) + print("Response:", res.text[:200]) + + +if __name__ == "__main__": + main() diff --git a/tests/storage/test_neighborhood_lock.py b/tests/storage/test_neighborhood_lock.py new file mode 100644 index 0000000..3ff07fe --- /dev/null +++ b/tests/storage/test_neighborhood_lock.py @@ -0,0 +1,311 @@ +"""Parametrized tests for the neighborhood lock registries. + +Covers the seven scenarios from the design doc. The multi-process SQLite +case is intentionally skipped here — the test harness is flaky across OSes +because SQLite's lock reporting depends on filesystem semantics. A +dedicated follow-up PR will bring that coverage in with an OS matrix. +""" +from __future__ import annotations + +import asyncio +import time +from pathlib import Path + +import pytest + +from memu.lane_lock import LaneContestedError +from memu.neighborhood_lock import ( + MarkdownLockRegistry, + NeighborhoodConflict, + NeighborhoodLock, + SqliteLockRegistry, +) +from memu.storage import get_backend + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +async def _make_sqlite_registry(tmp_path: Path): + backend = get_backend(f"sqlite:///{tmp_path / 'index.db'}") + await backend.init() + registry = backend.get_lock_registry() + await registry.init() + return registry, backend.close + + +async def _make_markdown_registry(tmp_path: Path): + backend = get_backend(f"file://{tmp_path}") + await backend.init() + registry = backend.get_lock_registry() + await registry.init() + return registry, backend.close + + +REGISTRY_FACTORIES = { + "sqlite": _make_sqlite_registry, + "markdown": _make_markdown_registry, +} + + +@pytest.fixture(params=sorted(REGISTRY_FACTORIES)) +def registry_factory(request): + return REGISTRY_FACTORIES[request.param] + + +def _run(coro): + return asyncio.run(coro) + + +# --------------------------------------------------------------------------- +# Scenarios +# --------------------------------------------------------------------------- + + +def test_scenario_1_basic_acquire_and_release(registry_factory, tmp_path): + """Single holder can acquire and release cleanly; fence starts at 1.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + rec = await registry.try_acquire("foo", "holder-a", ttl_s=30) + assert rec.slug == "foo" + assert rec.holder_id == "holder-a" + assert rec.fencing_token == 1 + + released = await registry.release("foo", "holder-a", rec.fencing_token) + assert released is True + assert await registry.peek("foo") is None + finally: + await close() + + _run(go()) + + +def test_scenario_2_concurrent_acquire_conflicts(registry_factory, tmp_path): + """Second holder gets NeighborhoodConflict while the lease is live.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + await registry.try_acquire("foo", "holder-a", ttl_s=30) + with pytest.raises(NeighborhoodConflict): + await registry.try_acquire("foo", "holder-b", ttl_s=30) + finally: + await close() + + _run(go()) + + +def test_scenario_3_conflict_subclasses_lane_contested( + registry_factory, tmp_path +): + """Existing lane-lock callers catch neighborhood conflicts too.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + await registry.try_acquire("foo", "holder-a", ttl_s=30) + with pytest.raises(LaneContestedError): + await registry.try_acquire("foo", "holder-b", ttl_s=30) + finally: + await close() + + _run(go()) + + +def test_scenario_4_release_enables_reacquire_with_monotonic_fence( + registry_factory, tmp_path +): + """After release, re-acquire succeeds and fence advances strictly.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + first = await registry.try_acquire("foo", "holder-a", ttl_s=30) + await registry.release("foo", "holder-a", first.fencing_token) + + second = await registry.try_acquire("foo", "holder-b", ttl_s=30) + assert second.fencing_token > first.fencing_token + assert second.fencing_token == first.fencing_token + 1 + finally: + await close() + + _run(go()) + + +def test_scenario_5_expired_lease_is_reclaimable(registry_factory, tmp_path): + """A lease past its TTL can be reclaimed by a new holder.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + # Very short TTL, then wait it out. + rec = await registry.try_acquire("foo", "holder-a", ttl_s=0.05) + await asyncio.sleep(0.08) + reclaimed = await registry.try_acquire( + "foo", "holder-b", ttl_s=30 + ) + assert reclaimed.holder_id == "holder-b" + assert reclaimed.fencing_token > rec.fencing_token + # Original holder's stale release must now fail. + assert ( + await registry.release("foo", "holder-a", rec.fencing_token) + is False + ) + finally: + await close() + + _run(go()) + + +def test_scenario_6_wrong_owner_cannot_release_or_renew( + registry_factory, tmp_path +): + """Release/renew are owner-scoped; wrong holder is a no-op / error.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + rec = await registry.try_acquire("foo", "holder-a", ttl_s=30) + # Wrong holder id. + assert ( + await registry.release("foo", "holder-b", rec.fencing_token) + is False + ) + # Wrong fencing token. + assert ( + await registry.release("foo", "holder-a", rec.fencing_token + 99) + is False + ) + # Renew by impostor should raise. + with pytest.raises(NeighborhoodConflict): + await registry.renew( + "foo", "holder-b", rec.fencing_token, ttl_s=30 + ) + # Real owner can still release. + assert ( + await registry.release("foo", "holder-a", rec.fencing_token) + is True + ) + finally: + await close() + + _run(go()) + + +def test_scenario_7_context_manager_renews_and_multi_lock( + registry_factory, tmp_path +): + """NeighborhoodLock acquires root+neighbors, renews in-flight, releases + everything on exit.""" + + async def go(): + registry, close = await registry_factory(tmp_path) + try: + holder_lock = NeighborhoodLock( + registry, + "root", + neighbors=["neighbor-a", "neighbor-b"], + holder_id="worker-1", + ttl_s=0.4, + renew_interval_s=0.1, + ) + async with holder_lock as nl: + assert nl.fencing_token >= 1 + # All three slugs should be live. + for slug in ("root", "neighbor-a", "neighbor-b"): + rec = await registry.peek(slug) + assert rec is not None + assert rec.holder_id == "worker-1" + + # A peer attempting any one of the three should be blocked. + with pytest.raises(NeighborhoodConflict): + await registry.try_acquire( + "neighbor-a", "worker-2", ttl_s=10 + ) + + # Sleep past the original TTL so renewal must have fired. + await asyncio.sleep(0.55) + rec_after = await registry.peek("root") + assert rec_after is not None + assert rec_after.holder_id == "worker-1" + + # Exit -> every slug released. + for slug in ("root", "neighbor-a", "neighbor-b"): + assert await registry.peek(slug) is None + finally: + await close() + + _run(go()) + + +# --------------------------------------------------------------------------- +# Markdown-specific: the fencing_token sidecar is actually honored. +# --------------------------------------------------------------------------- + + +def test_markdown_honors_fence_sidecar_on_release(tmp_path): + """Corrupting ``.memu/locks/_fence/`` should invalidate release.""" + + async def go(): + registry, close = await _make_markdown_registry(tmp_path) + try: + rec = await registry.try_acquire("paper/attention", "holder-a", ttl_s=30) + fence_file = tmp_path / ".memu" / "locks" / "_fence" / "paper__attention" + assert fence_file.exists() + # Fast-forward the fence counter as if another claim had bumped it. + fence_file.write_text(str(rec.fencing_token + 5)) + # Release now sees a mismatched fence and refuses to delete. + assert ( + await registry.release( + "paper/attention", "holder-a", rec.fencing_token + ) + is False + ) + # Renew with the stale token hits the same check. + with pytest.raises(NeighborhoodConflict): + await registry.renew( + "paper/attention", + "holder-a", + rec.fencing_token, + ttl_s=30, + ) + finally: + await close() + + _run(go()) + + +def test_sqlite_registry_shares_backend_connection(tmp_path): + """Sanity: the SQLite registry reuses the existing backend connection + (tables land in the same database file as PR #21's node tables).""" + + async def go(): + backend = get_backend(f"sqlite:///{tmp_path / 'index.db'}") + await backend.init() + registry = backend.get_lock_registry() + assert isinstance(registry, SqliteLockRegistry) + await registry.init() + await registry.try_acquire("foo", "holder-a", ttl_s=30) + row = backend.conn.execute( + "SELECT slug, holder_id FROM neighborhood_locks" + ).fetchone() + assert row["slug"] == "foo" + assert row["holder_id"] == "holder-a" + await backend.close() + + _run(go()) + + +def test_markdown_registry_type(tmp_path): + async def go(): + backend = get_backend(f"file://{tmp_path}") + await backend.init() + registry = backend.get_lock_registry() + assert isinstance(registry, MarkdownLockRegistry) + await backend.close() + + _run(go())