diff --git a/src/powercontext/builtin/artifacts/memory/service.py b/src/powercontext/builtin/artifacts/memory/service.py index ea22f251f..c1b422321 100644 --- a/src/powercontext/builtin/artifacts/memory/service.py +++ b/src/powercontext/builtin/artifacts/memory/service.py @@ -930,7 +930,10 @@ async def _validated_entries(self, memory: Memory) -> tuple[MemoryEntryVersion, if version.memory_artifact_id != memory.artifact_id or version.entry_id != item.entry_id: raise InvalidMemoryCitationError("cross-identity") material = self._material_from_version(version) - if material.content_hash != item.entry_content_hash: + if ( + material.content_hash != item.entry_content_hash + or version.entry_content_hash != item.entry_content_hash + ): raise InvalidMemoryCitationError("hash-mismatch") ordered.append(version) return tuple(ordered) diff --git a/src/powercontext/builtin/persistence/memory.py b/src/powercontext/builtin/persistence/memory.py index 97164a4c8..4d91d6ca7 100644 --- a/src/powercontext/builtin/persistence/memory.py +++ b/src/powercontext/builtin/persistence/memory.py @@ -16,6 +16,7 @@ from __future__ import annotations +import math from collections.abc import AsyncIterator, Mapping from contextlib import AbstractAsyncContextManager, asynccontextmanager from dataclasses import dataclass, replace @@ -23,6 +24,7 @@ from pydantic import RootModel from sqlalchemy import delete, insert, select +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncConnection from powercontext.artifacts import ArtifactDraft, ArtifactRef @@ -31,10 +33,12 @@ InvalidEmbeddingError, Memory, MemoryCapabilities, + MemoryChange, MemoryCommit, MemoryContent, MemoryEntryVersion, MemoryHit, + MemoryManifestEntry, MemoryProjection, MemoryRevisionChanges, MemorySearchChannels, @@ -44,8 +48,10 @@ from powercontext.builtin.artifacts.memory.canonical import ( canonical_embedding, embedding_content_hash, + entry_content_bytes, entry_content_hash, memory_content_hash, + validate_embedding, ) from powercontext.builtin.artifacts.memory.errors import ( InvalidMemoryCitationError, @@ -66,6 +72,8 @@ from powercontext.errors import ArtifactNotFoundError from powercontext.sources import SourceRef +_UNIT_VECTOR_ABS_TOLERANCE = 1e-12 + class _SourceRefs(RootModel[tuple[SourceRef, ...]]): pass @@ -89,13 +97,21 @@ class _InvalidMemoryCommitError(MemoryBackendConfigurationError): def __init__(self, code: str, actual: str | None = None) -> None: details = { "artifact-result": "generic Artifact result differs from prepared revision", + "base": "base is not the authoritative stored Memory revision", "base-identity": "base and revision identities differ", + "changes": "revision changes are not unique, ordered, and complete", "complete": "unit of work is already complete", "content-hash": "content hash does not match canonical content", + "entry-hash": "entry body, declared hash, and manifest hash differ", + "entry-history": "entry version history is incomplete or has an invalid predecessor", + "entry-identity": "entry versions do not match the committed Memory revision", "family": "family is not memory", + "manifest": "manifest entries are not unique and canonically ordered", "memory-type": f"expected Memory, got {actual}", - "projection": "active manifest and projections differ", + "projection": "active manifest and projections differ or contain non-canonical content", "revision": "revision is not the next prepared revision", + "transition": "changes do not describe the exact base-to-revision transition", + "vector": "embedding projection does not match the configured profile and entry content", } super().__init__(f"invalid relational Memory commit: {details[code]}") @@ -155,11 +171,11 @@ async def latest(self, artifact_id: str, /) -> Memory: return _require_memory(artifact) async def entries(self, memory: ArtifactRef, /) -> tuple[MemoryEntryVersion, ...]: - canonical = await self.get(memory) - version_ids = tuple(item.entry_version_id for item in canonical.content.manifest.entries) - if not version_ids: - return () async with self._database.connection(self._bound_connection) as connection: + canonical = await self._get_memory(connection, memory) + version_ids = tuple(item.entry_version_id for item in canonical.content.manifest.entries) + if not version_ids: + return () rows = ( await connection.execute( select(MEMORY_ENTRY_VERSIONS_TABLE).where( @@ -170,13 +186,18 @@ async def entries(self, memory: ArtifactRef, /) -> tuple[MemoryEntryVersion, ... ) ).mappings() by_id = {str(row["entry_version_id"]): _decode_entry(row) for row in rows} - if set(by_id) != set(version_ids): + if len(by_id) != len(version_ids) or set(by_id) != set(version_ids): raise InvalidMemoryCitationError("missing-version") - return tuple(by_id[version_id] for version_id in version_ids) + ordered: list[MemoryEntryVersion] = [] + for item in canonical.content.manifest.entries: + version = by_id[item.entry_version_id] + _validate_manifest_entry(canonical.as_ref(), item, version) + ordered.append(version) + return tuple(ordered) async def projections(self, memory: ArtifactRef, /) -> tuple[MemoryProjection, ...]: - canonical = await self.get(memory) async with self._database.connection(self._bound_connection) as connection: + canonical = await self._get_memory(connection, memory) rows = ( await connection.execute( select(MEMORY_ENTRY_HEADS_TABLE, MEMORY_ENTRY_VERSIONS_TABLE) @@ -208,9 +229,14 @@ async def projections(self, memory: ArtifactRef, /) -> tuple[MemoryProjection, . for row in rows ) projections = await self._index.hydrate(connection, self._scope_id, projections) - active_ids = {item.entry_version_id for item in canonical.content.manifest.entries if item.state == "active"} - if {item.entry_version.entry_version_id for item in projections} != active_ids: + active = {item.entry_version_id: item for item in canonical.content.manifest.entries if item.state == "active"} + by_id = {item.entry_version.entry_version_id: item for item in projections} + if len(by_id) != len(projections) or set(by_id) != set(active): raise InvalidMemoryCitationError("projection-version") + for entry_version_id, projection in by_id.items(): + _validate_manifest_entry(canonical.as_ref(), active[entry_version_id], projection.entry_version) + if projection.searchable_text != analyze_text(projection.entry_version.text): + raise InvalidMemoryCitationError("projection-version") return projections async def rebuild_projections(self, embedding_model: EmbeddingModel | None = None, /) -> None: @@ -286,8 +312,9 @@ async def vector_complete(self, memories: tuple[ArtifactRef, ...], profile: Embe async def search(self, request: MemorySearchRequest, /) -> MemorySearchChannels: async with self._database.connection(self._bound_connection) as connection: - await self._validate_search_heads(connection, request.memories) + memories = await self._validate_search_heads(connection, request.memories) channels = await self._index.search(connection, self._scope_id, request) + await self._validate_search_channels(connection, memories, channels) await self._validate_search_heads(connection, request.memories) return channels @@ -295,9 +322,10 @@ async def _validate_search_heads( self, connection: AsyncConnection, memories: tuple[ArtifactRef, ...], - ) -> None: + ) -> dict[tuple[str, int], Memory]: """Reject a projection read when any requested head has advanced.""" + canonical: dict[tuple[str, int], Memory] = {} for memory in memories: try: exact = _require_memory(await self._artifacts.get(connection, self._scope_id, memory)) @@ -313,11 +341,93 @@ async def _validate_search_heads( raise ArtifactNotFoundError(memory) from None if exact.as_ref() != latest.as_ref(): raise InvalidMemoryCitationError("memory-mismatch") + canonical[(memory.artifact_id, memory.revision)] = exact + return canonical + + async def _validate_search_channels( + self, + connection: AsyncConnection, + memories: Mapping[tuple[str, int], Memory], + channels: MemorySearchChannels, + ) -> None: + candidates = (*channels.fts, *channels.vector) + if not candidates: + return + version_ids = tuple({hit.entry_version_id for hit in candidates}) + rows = ( + await connection.execute( + select( + MEMORY_ENTRY_VERSIONS_TABLE, + MEMORY_ENTRY_HEADS_TABLE.c.head_revision.label("_head_revision"), + MEMORY_ENTRY_HEADS_TABLE.c.entry_content_hash.label("_head_content_hash"), + MEMORY_ENTRY_HEADS_TABLE.c.searchable_text.label("_searchable_text"), + ) + .join( + MEMORY_ENTRY_HEADS_TABLE, + (MEMORY_ENTRY_HEADS_TABLE.c.scope_id == MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id) + & ( + MEMORY_ENTRY_HEADS_TABLE.c.memory_artifact_id + == MEMORY_ENTRY_VERSIONS_TABLE.c.memory_artifact_id + ) + & (MEMORY_ENTRY_HEADS_TABLE.c.entry_id == MEMORY_ENTRY_VERSIONS_TABLE.c.entry_id) + & (MEMORY_ENTRY_HEADS_TABLE.c.entry_version_id == MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id), + ) + .where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == self._scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id.in_(version_ids), + ) + ) + ).mappings() + authoritative = { + ( + str(row["memory_artifact_id"]), + int(row["_head_revision"]), + str(row["entry_id"]), + str(row["entry_version_id"]), + ): row + for row in rows + } + manifests = { + key: {item.entry_id: item for item in memory.content.manifest.entries} for key, memory in memories.items() + } + for hit in candidates: + memory_key = (hit.memory_ref.artifact_id, hit.memory_ref.revision) + memory = memories.get(memory_key) + item = manifests.get(memory_key, {}).get(hit.entry_id) + row = authoritative.get((*memory_key, hit.entry_id, hit.entry_version_id)) + if ( + memory is None + or hit.memory_ref.family != Memory.family + or item is None + or item.state != "active" + or item.entry_version_id != hit.entry_version_id + or row is None + or str(row["_head_content_hash"]) != item.entry_content_hash + ): + raise InvalidMemoryCitationError("search-anchor") + version = _decode_entry(row) + _validate_manifest_entry(memory.as_ref(), item, version) + if hit.text != version.text or str(row["_searchable_text"]) != analyze_text(version.text): + raise InvalidMemoryCitationError("hash-mismatch") async def expand(self, hits: tuple[MemoryHit, ...], /) -> tuple[MemoryEntryVersion, ...]: expanded: list[MemoryEntryVersion] = [] async with self._database.connection(self._bound_connection) as connection: + memories: dict[tuple[str, int], Memory] = {} for hit in hits: + memory_key = (hit.memory_ref.artifact_id, hit.memory_ref.revision) + memory = memories.get(memory_key) + if memory is None: + memory = await self._get_memory(connection, hit.memory_ref) + memories[memory_key] = memory + item = next( + ( + candidate + for candidate in memory.content.manifest.entries + if candidate.entry_id == hit.entry_id and candidate.entry_version_id == hit.entry_version_id + ), + None, + ) row = ( ( await connection.execute( @@ -332,11 +442,21 @@ async def expand(self, hits: tuple[MemoryHit, ...], /) -> tuple[MemoryEntryVersi .mappings() .one_or_none() ) - if row is None: + if item is None or row is None: raise InvalidMemoryCitationError("expand-anchor") - expanded.append(_decode_entry(row)) + version = _decode_entry(row) + _validate_manifest_entry(memory.as_ref(), item, version) + expanded.append(version) return tuple(expanded) + async def _get_memory(self, connection: AsyncConnection, memory: ArtifactRef) -> Memory: + if memory.family != Memory.family: + raise ArtifactNotFoundError(memory) + try: + return _require_memory(await self._artifacts.get(connection, self._scope_id, memory)) + except RepositoryNotFoundError: + raise ArtifactNotFoundError(memory) from None + async def _authoritative_projections( self, connection: AsyncConnection, @@ -382,7 +502,7 @@ async def _authoritative_projections( version = versions.get(item.entry_version_id) if version is None: raise InvalidMemoryCitationError("missing-version") - _validate_rebuild_entry(ref, item.entry_id, item.entry_content_hash, version) + _validate_manifest_entry(ref, item, version) projections.append( MemoryProjection( entry_version=version, @@ -441,6 +561,7 @@ async def _embed_rebuild( async def _commit(self, connection: AsyncConnection, value: MemoryCommit) -> Memory: _validate_commit(value) + await self._validate_commit_relations(connection, value) draft = _MemoryDraft( content=value.memory.content, sources=value.memory.lineage.sources, @@ -465,10 +586,13 @@ async def _commit(self, connection: AsyncConnection, value: MemoryCommit) -> Mem raise _InvalidMemoryCommitError("artifact-result") if value.entry_versions: - await connection.execute( - insert(MEMORY_ENTRY_VERSIONS_TABLE), - [_entry_values(self._scope_id, entry) for entry in value.entry_versions], - ) + try: + await connection.execute( + insert(MEMORY_ENTRY_VERSIONS_TABLE), + [_entry_values(self._scope_id, entry) for entry in value.entry_versions], + ) + except IntegrityError as error: + raise _InvalidMemoryCommitError("entry-identity") from error await connection.execute( delete(MEMORY_ENTRY_HEADS_TABLE).where( MEMORY_ENTRY_HEADS_TABLE.c.scope_id == self._scope_id, @@ -491,6 +615,204 @@ async def _commit(self, connection: AsyncConnection, value: MemoryCommit) -> Mem ) return committed + async def _validate_commit_relations(self, connection: AsyncConnection, value: MemoryCommit) -> None: + canonical_base = await self._canonical_commit_base(connection, value) + await self._validate_base_history(connection, canonical_base) + new_by_entry = _validate_revision_transition(value, canonical_base) + versions = await self._commit_versions(connection, value) + await self._validate_new_version_history(connection, value, canonical_base, new_by_entry) + self._validate_commit_projections(value, versions) + + async def _canonical_commit_base( + self, + connection: AsyncConnection, + value: MemoryCommit, + ) -> Memory | None: + if value.base is not None: + try: + canonical = _require_memory(await self._artifacts.get(connection, self._scope_id, value.base.as_ref())) + except RepositoryNotFoundError: + raise _InvalidMemoryCommitError("base") from None + if canonical != value.base: + raise _InvalidMemoryCommitError("base") + return canonical + return None + + async def _validate_base_history(self, connection: AsyncConnection, base: Memory | None) -> None: + if base is None: + return + rows = ( + await connection.execute( + select(MEMORY_ENTRY_VERSIONS_TABLE).where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == self._scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.memory_artifact_id == base.artifact_id, + ) + ) + ).mappings() + stored = tuple(rows) + rows_by_id = {str(row["entry_version_id"]): row for row in stored} + if len(rows_by_id) != len(stored): + raise _InvalidMemoryCommitError("entry-history") + by_id: dict[str, MemoryEntryVersion] = {} + pending = [item.entry_version_id for item in base.content.manifest.entries] + while pending: + entry_version_id = pending.pop() + if entry_version_id in by_id: + continue + row = rows_by_id.get(entry_version_id) + if row is None: + continue + version = _decode_entry(row) + by_id[entry_version_id] = version + if version.previous_version_id is not None: + pending.append(version.previous_version_id) + for item in base.content.manifest.entries: + version = by_id.get(item.entry_version_id) + if version is None or not _manifest_entry_matches(base.as_ref(), item, version): + raise _InvalidMemoryCommitError("entry-history") + if not _entry_history_matches(base.as_ref(), item, version, by_id): + raise _InvalidMemoryCommitError("entry-history") + + async def _commit_versions( + self, + connection: AsyncConnection, + value: MemoryCommit, + ) -> dict[str, MemoryEntryVersion]: + manifest = {item.entry_version_id: item for item in value.memory.content.manifest.entries} + new_by_id = {version.entry_version_id: version for version in value.entry_versions} + rows = ( + await connection.execute( + select(MEMORY_ENTRY_VERSIONS_TABLE).where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == self._scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id.in_(tuple(manifest)), + ) + ) + ).mappings() + stored_rows = tuple(rows) + if any(str(row["entry_version_id"]) in new_by_id for row in stored_rows): + raise _InvalidMemoryCommitError("entry-identity") + stored_by_id = { + str(row["entry_version_id"]): _decode_entry(row) + for row in stored_rows + if str(row["memory_artifact_id"]) == value.memory.artifact_id + } + versions = stored_by_id | new_by_id + if set(versions) != set(manifest): + raise _InvalidMemoryCommitError("entry-identity") + for entry_version_id, item in manifest.items(): + if not _manifest_entry_matches(value.memory.as_ref(), item, versions[entry_version_id]): + raise _InvalidMemoryCommitError("entry-hash") + return versions + + async def _validate_new_version_history( + self, + connection: AsyncConnection, + value: MemoryCommit, + canonical_base: Memory | None, + new_by_entry: Mapping[str, MemoryEntryVersion], + ) -> None: + previous_ids = tuple( + version.previous_version_id for version in value.entry_versions if version.previous_version_id is not None + ) + previous_rows = ( + () + if not previous_ids + else ( + await connection.execute( + select(MEMORY_ENTRY_VERSIONS_TABLE).where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == self._scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id.in_(previous_ids), + ) + ) + ) + .mappings() + .all() + ) + previous_by_id: dict[str, list[MemoryEntryVersion]] = {} + for row in previous_rows: + previous_by_id.setdefault(str(row["entry_version_id"]), []).append(_decode_entry(row)) + base_manifest = ( + {} if canonical_base is None else {item.entry_id: item for item in canonical_base.content.manifest.entries} + ) + changes = {change.entry_id: change for change in value.memory.content.changes} + for entry_id, version in new_by_entry.items(): + if version.created_in_revision != value.memory.revision: + raise _InvalidMemoryCommitError("entry-history") + change = changes[entry_id] + if change.op == "add": + if version.version != 1 or version.previous_version_id is not None: + raise _InvalidMemoryCommitError("entry-history") + continue + if canonical_base is None: + raise _InvalidMemoryCommitError("entry-history") + base_item = base_manifest[entry_id] + predecessors = previous_by_id.get(version.previous_version_id or "", []) + if len(predecessors) != 1: + raise _InvalidMemoryCommitError("entry-history") + predecessor = predecessors[0] + if ( + version.previous_version_id != base_item.entry_version_id + or version.version != predecessor.version + 1 + or not _manifest_entry_matches(canonical_base.as_ref(), base_item, predecessor) + ): + raise _InvalidMemoryCommitError("entry-history") + if base_item.state == "inactive" and ( + change.reason != "normalize" or not _canonical_entry_content_matches(version, predecessor) + ): + raise _InvalidMemoryCommitError("entry-history") + + def _validate_commit_projections( + self, + value: MemoryCommit, + versions: Mapping[str, MemoryEntryVersion], + ) -> None: + active = {item.entry_id: item for item in value.memory.content.manifest.entries if item.state == "active"} + projected = {projection.entry_version.entry_id: projection for projection in value.projections} + if len(projected) != len(value.projections) or set(projected) != set(active): + raise _InvalidMemoryCommitError("projection") + for entry_id, item in active.items(): + projection = projected[entry_id] + authoritative = versions[item.entry_version_id] + if projection.entry_version != authoritative or projection.searchable_text != analyze_text( + authoritative.text + ): + raise _InvalidMemoryCommitError("projection") + self._validate_embedding_projection(projection) + + def _validate_embedding_projection(self, projection: MemoryProjection) -> None: + if projection.embedding is None and projection.embedding_content_hash is None: + return + if projection.embedding is None or projection.embedding_content_hash is None: + raise _InvalidMemoryCommitError("vector") + profile = self._index.capabilities.embedding_profile + if not self._index.capabilities.vector or profile is None: + raise _InvalidMemoryCommitError("vector") + try: + vector = validate_embedding( + projection.embedding, + dimension=profile.dimension, + ) + expected_hash = embedding_content_hash( + profile_id=profile.profile_id, + model=profile.model, + dimension=profile.dimension, + distance=profile.distance, + normalization=profile.normalization, + entry_content_hash=projection.entry_version.entry_content_hash, + ) + except (TypeError, ValueError) as error: + raise _InvalidMemoryCommitError("vector") from error + if ( + profile.normalization == "unit" + and not math.isclose( + math.hypot(*vector), + 1.0, + rel_tol=0.0, + abs_tol=_UNIT_VECTOR_ABS_TOLERANCE, + ) + ) or expected_hash != projection.embedding_content_hash: + raise _InvalidMemoryCommitError("vector") + class _RelationalMemoryUnitOfWork: def __init__(self, backend: RelationalMemoryBackend, connection: AsyncConnection) -> None: @@ -515,9 +837,15 @@ async def _unit_of_work( def _validate_commit(value: MemoryCommit) -> None: - if value.memory.family != Memory.family: + if type(value.memory) is not Memory or value.memory.family != Memory.family: raise _InvalidMemoryCommitError("family") - if value.content_hash != memory_content_hash(value.memory.content): + if value.base is not None and type(value.base) is not Memory: + raise _InvalidMemoryCommitError("base") + try: + canonical_hash = memory_content_hash(value.memory.content) + except (TypeError, ValueError) as error: + raise _InvalidMemoryCommitError("content-hash") from error + if value.content_hash != canonical_hash: raise _InvalidMemoryCommitError("content-hash") expected_revision = 1 if value.base is None else value.base.revision + 1 if value.memory.revision != expected_revision: @@ -530,6 +858,112 @@ def _validate_commit(value: MemoryCommit) -> None: raise _InvalidMemoryCommitError("projection") +def _validate_revision_transition( + value: MemoryCommit, + base: Memory | None, +) -> dict[str, MemoryEntryVersion]: + manifest_entries = value.memory.content.manifest.entries + manifest_ids = tuple(item.entry_id for item in manifest_entries) + manifest_version_ids = tuple(item.entry_version_id for item in manifest_entries) + if ( + len(set(manifest_ids)) != len(manifest_ids) + or len(set(manifest_version_ids)) != len(manifest_version_ids) + or manifest_ids != tuple(sorted(manifest_ids, key=str.encode)) + ): + raise _InvalidMemoryCommitError("manifest") + + changes = value.memory.content.changes + change_ids = tuple(change.entry_id for change in changes) + if ( + not changes + or len(set(change_ids)) != len(change_ids) + or change_ids != tuple(sorted(change_ids, key=str.encode)) + ): + raise _InvalidMemoryCommitError("changes") + + base_entries = () if base is None else base.content.manifest.entries + base_ids = tuple(item.entry_id for item in base_entries) + if len(set(base_ids)) != len(base_ids): + raise _InvalidMemoryCommitError("base") + before = {item.entry_id: item for item in base_entries} + after = {item.entry_id: item for item in manifest_entries} + changed: set[str] = set() + new_targets: dict[str, str] = {} + + for change in changes: + previous = before.get(change.entry_id) + current = after.get(change.entry_id) + target = _validate_transition_change(change, previous, current) + if target is not None: + new_targets[change.entry_id] = target + changed.add(change.entry_id) + + if set(after) != set(before) | {entry_id for entry_id in changed if entry_id not in before}: + raise _InvalidMemoryCommitError("transition") + for entry_id, item in before.items(): + if entry_id not in changed and after.get(entry_id) != item: + raise _InvalidMemoryCommitError("transition") + + entry_versions = value.entry_versions + by_entry = {version.entry_id: version for version in entry_versions} + version_ids = {version.entry_version_id for version in entry_versions} + if ( + len(by_entry) != len(entry_versions) + or len(version_ids) != len(entry_versions) + or set(by_entry) != set(new_targets) + or any(by_entry[entry_id].entry_version_id != target for entry_id, target in new_targets.items()) + ): + raise _InvalidMemoryCommitError("entry-identity") + return by_entry + + +def _validate_transition_change( + change: MemoryChange, + previous: MemoryManifestEntry | None, + current: MemoryManifestEntry | None, +) -> str | None: + if current is None: + raise _InvalidMemoryCommitError("transition") + if change.op == "add": + valid = ( + previous is None + and change.from_entry_version_id is None + and change.to_entry_version_id == current.entry_version_id + and current.state == "active" + ) + target = current.entry_version_id + elif change.op == "revise": + valid = ( + previous is not None + and change.from_entry_version_id == previous.entry_version_id + and change.to_entry_version_id == current.entry_version_id + and current.entry_version_id != previous.entry_version_id + and current.state == previous.state + ) + target = current.entry_version_id + elif change.op == "deactivate": + valid = ( + previous is not None + and previous.state == "active" + and change.from_entry_version_id == previous.entry_version_id + and change.to_entry_version_id is None + and current == previous.model_copy(update={"state": "inactive"}) + ) + target = None + else: + valid = ( + previous is not None + and previous.state == "inactive" + and change.from_entry_version_id is None + and change.to_entry_version_id == previous.entry_version_id + and current == previous.model_copy(update={"state": "active"}) + ) + target = None + if not valid: + raise _InvalidMemoryCommitError("transition") + return target + + def _entry_values(scope_id: str, value: MemoryEntryVersion) -> dict[str, object]: return { "scope_id": scope_id, @@ -601,32 +1035,112 @@ def _decode_entry(row: Mapping[Any, Any]) -> MemoryEntryVersion: ) -def _validate_rebuild_entry( +def _manifest_entry_matches( memory_ref: ArtifactRef, - entry_id: str, - content_hash: str, + item: MemoryManifestEntry, version: MemoryEntryVersion, -) -> None: +) -> bool: if ( version.memory_artifact_id != memory_ref.artifact_id - or version.entry_id != entry_id - or version.entry_content_hash != content_hash + or version.entry_id != item.entry_id + or version.entry_version_id != item.entry_version_id + or version.entry_content_hash != item.entry_content_hash + or version.created_in_revision < 1 + or version.created_in_revision > memory_ref.revision ): - raise InvalidMemoryCitationError("cross-identity") - actual_hash = entry_content_hash( + return False + return _entry_declared_hash_matches(version) + + +def _entry_history_matches( + memory_ref: ArtifactRef, + item: MemoryManifestEntry, + head: MemoryEntryVersion, + versions: Mapping[str, MemoryEntryVersion], +) -> bool: + current = head + expected_version = head.version + visited: set[str] = set() + while True: + if current.entry_version_id in visited: + return False + visited.add(current.entry_version_id) + if ( + current.memory_artifact_id != memory_ref.artifact_id + or current.entry_id != item.entry_id + or current.version != expected_version + or current.created_in_revision < 1 + or current.created_in_revision > memory_ref.revision + or not _entry_declared_hash_matches(current) + ): + return False + if expected_version == 1: + return current.previous_version_id is None + if current.previous_version_id is None: + return False + predecessor = versions.get(current.previous_version_id) + if predecessor is None or predecessor.created_in_revision >= current.created_in_revision: + return False + current = predecessor + expected_version -= 1 + + +def _entry_declared_hash_matches(version: MemoryEntryVersion) -> bool: + try: + actual_hash = entry_content_hash( + kind=version.kind, + text=version.text, + source_refs=_entry_source_refs(version), + artifact_refs=_entry_artifact_refs(version), + ) + except (TypeError, ValueError): + return False + return actual_hash == version.entry_content_hash + + +def _canonical_entry_content_matches(left: MemoryEntryVersion, right: MemoryEntryVersion) -> bool: + try: + return _canonical_entry_bytes(left) == _canonical_entry_bytes(right) + except (TypeError, ValueError): + return False + + +def _canonical_entry_bytes(version: MemoryEntryVersion) -> bytes: + return entry_content_bytes( kind=version.kind, text=version.text, - source_refs=tuple({"source_type": ref.source_type, "source_id": ref.source_id} for ref in version.sources), - artifact_refs=tuple( - { - "family": ref.family, - "artifact_id": ref.artifact_id, - "revision": ref.revision, - } - for ref in version.artifacts - ), + source_refs=_entry_source_refs(version), + artifact_refs=_entry_artifact_refs(version), ) - if actual_hash != content_hash: + + +def _entry_source_refs(version: MemoryEntryVersion) -> tuple[dict[str, str], ...]: + return tuple({"source_type": ref.source_type, "source_id": ref.source_id} for ref in version.sources) + + +def _entry_artifact_refs(version: MemoryEntryVersion) -> tuple[dict[str, str | int], ...]: + return tuple( + { + "family": ref.family, + "artifact_id": ref.artifact_id, + "revision": ref.revision, + } + for ref in version.artifacts + ) + + +def _validate_manifest_entry( + memory_ref: ArtifactRef, + item: MemoryManifestEntry, + version: MemoryEntryVersion, +) -> None: + if ( + version.memory_artifact_id != memory_ref.artifact_id + or version.entry_id != item.entry_id + or version.entry_version_id != item.entry_version_id + ): + raise InvalidMemoryCitationError("cross-identity") + if not _manifest_entry_matches(memory_ref, item, version): raise InvalidMemoryCitationError("hash-mismatch") diff --git a/src/powercontext/builtin/persistence/memory_schema.py b/src/powercontext/builtin/persistence/memory_schema.py new file mode 100644 index 000000000..15e55624c --- /dev/null +++ b/src/powercontext/builtin/persistence/memory_schema.py @@ -0,0 +1,92 @@ +# Copyright (c) 2026 OceanBase. +# +# 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. + +"""Schema upgrades for authoritative Memory entry history.""" + +from __future__ import annotations + +from sqlalchemy import func, select, text +from sqlalchemy.exc import DBAPIError, IntegrityError +from sqlalchemy.ext.asyncio import AsyncConnection + +from powercontext.builtin.artifacts.memory.errors import MemoryBackendConfigurationError +from powercontext.builtin.persistence.tables import ( + MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME, + MEMORY_ENTRY_VERSIONS_TABLE, +) + +_SQLITE_INDEX_EXISTS = text( + "SELECT COUNT(*) FROM pragma_index_list('pc_memory_entry_versions') WHERE name = :index_name" +).bindparams(index_name=MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME) +_MYSQL_INDEX_EXISTS = text( + "SELECT COUNT(*) FROM information_schema.statistics " + "WHERE table_schema = DATABASE() " + "AND table_name = 'pc_memory_entry_versions' " + "AND index_name = :index_name" +).bindparams(index_name=MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME) +_CREATE_SCOPE_VERSION_INDEX = ( + f"CREATE UNIQUE INDEX {MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME} " + "ON pc_memory_entry_versions (scope_id, entry_version_id)" +) + + +async def ensure_memory_entry_version_scope_identity(connection: AsyncConnection, /) -> None: + """Make entry-version identities scope-global on new and existing databases.""" + + if await _scope_version_index_exists(connection): + return + duplicate = ( + await connection.execute( + select( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id, + ) + .group_by( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id, + ) + .having(func.count() > 1) + .limit(1) + ) + ).first() + if duplicate is not None: + raise _duplicate_identity_error() + try: + await connection.exec_driver_sql(_CREATE_SCOPE_VERSION_INDEX) + except DBAPIError as error: + if await _scope_version_index_exists(connection): + return + if isinstance(error, IntegrityError): + raise _duplicate_identity_error() from error + raise + + +async def _scope_version_index_exists(connection: AsyncConnection) -> bool: + dialect = connection.dialect.name + if dialect == "sqlite": + statement = _SQLITE_INDEX_EXISTS + elif dialect == "mysql": + statement = _MYSQL_INDEX_EXISTS + else: + raise ValueError(f"unsupported Memory schema migration dialect: {dialect}") # noqa: TRY003 + return int(await connection.scalar(statement) or 0) > 0 + + +def _duplicate_identity_error() -> MemoryBackendConfigurationError: + return MemoryBackendConfigurationError( + "pc_memory_entry_versions contains duplicate scope-global entry_version_id values" + ) + + +__all__ = ["ensure_memory_entry_version_scope_identity"] diff --git a/src/powercontext/builtin/persistence/oceanbase/profile.py b/src/powercontext/builtin/persistence/oceanbase/profile.py index 497e927f8..73c548408 100644 --- a/src/powercontext/builtin/persistence/oceanbase/profile.py +++ b/src/powercontext/builtin/persistence/oceanbase/profile.py @@ -29,7 +29,9 @@ from powercontext.builtin.persistence.database import AsyncDatabase from powercontext.builtin.persistence.errors import PersistenceError +from powercontext.builtin.persistence.memory_schema import ensure_memory_entry_version_scope_identity from powercontext.builtin.persistence.schema import create_tables +from powercontext.builtin.persistence.tables import MEMORY_ENTRY_VERSIONS_TABLE _DIALECT_DRIVER = "mysql+aoceanbase" _DIALECT_REGISTRY_NAME = "mysql.aoceanbase" @@ -114,6 +116,8 @@ async def _initialized_profile(profile: OceanBaseProfile) -> AsyncIterator[Ocean async with profile.database.transaction() as connection: await _require_mysql_tenant(connection) await create_tables(connection, profile.tables) + if any(table is MEMORY_ENTRY_VERSIONS_TABLE for table in profile.tables): + await ensure_memory_entry_version_scope_identity(connection) yield profile finally: await profile.database.close() diff --git a/src/powercontext/builtin/persistence/sqlite/profile.py b/src/powercontext/builtin/persistence/sqlite/profile.py index c9cdc6be8..a0967f396 100644 --- a/src/powercontext/builtin/persistence/sqlite/profile.py +++ b/src/powercontext/builtin/persistence/sqlite/profile.py @@ -33,7 +33,9 @@ from sqlalchemy.pool import StaticPool from powercontext.builtin.persistence.database import AsyncDatabase +from powercontext.builtin.persistence.memory_schema import ensure_memory_entry_version_scope_identity from powercontext.builtin.persistence.schema import create_tables +from powercontext.builtin.persistence.tables import MEMORY_ENTRY_VERSIONS_TABLE _WARMUP_RETRY_SECONDS = 0.05 _WARMUP_LOCKS: WeakKeyDictionary[asyncio.AbstractEventLoop, dict[str, asyncio.Lock]] = WeakKeyDictionary() @@ -98,6 +100,8 @@ async def open( await _warm_sqlite(engine, config) async with database.transaction() as connection: await create_tables(connection, tables) + if any(table is MEMORY_ENTRY_VERSIONS_TABLE for table in tables): + await ensure_memory_entry_version_scope_identity(connection) yield profile finally: await database.close() diff --git a/src/powercontext/builtin/persistence/tables.py b/src/powercontext/builtin/persistence/tables.py index 99e1013a6..3f32e7992 100644 --- a/src/powercontext/builtin/persistence/tables.py +++ b/src/powercontext/builtin/persistence/tables.py @@ -21,6 +21,7 @@ Column, Date, ForeignKeyConstraint, + Index, Integer, LargeBinary, MetaData, @@ -364,6 +365,7 @@ def _entry_text_type(): MAX_MEMORY_ENTRY_ID_LENGTH = 128 MAX_MEMORY_ENTRY_KIND_LENGTH = 128 MAX_MEMORY_HASH_LENGTH = 64 +MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME = "uq_pc_memory_entry_versions_scope_version" MEMORY_ENTRY_VERSIONS_TABLE = Table( @@ -413,6 +415,13 @@ def _entry_text_type(): ), ) +MEMORY_ENTRY_VERSION_SCOPE_INDEX = Index( + MEMORY_ENTRY_VERSION_SCOPE_INDEX_NAME, + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id, + unique=True, +) + MEMORY_ENTRY_HEADS_TABLE = Table( "pc_memory_entry_heads", SHARED_METADATA, diff --git a/tests/builtin/artifacts/memory/test_service.py b/tests/builtin/artifacts/memory/test_service.py index 0805318f4..6938a133a 100644 --- a/tests/builtin/artifacts/memory/test_service.py +++ b/tests/builtin/artifacts/memory/test_service.py @@ -16,8 +16,15 @@ import asyncio +import pytest +from sqlalchemy.ext.asyncio import AsyncConnection + +from powercontext.artifacts import ArtifactRef from powercontext.builtin.artifacts.memory import ( + EmbeddingProfile, Memory, + MemoryCapabilities, + MemoryChange, MemoryCommit, MemoryContent, MemoryEntryInput, @@ -29,8 +36,11 @@ MemoryService, ) from powercontext.builtin.artifacts.memory.canonical import entry_content_hash, memory_content_hash -from powercontext.builtin.inference import InferenceUsage +from powercontext.builtin.artifacts.memory.errors import MemoryBackendConfigurationError +from powercontext.builtin.artifacts.search import analyze_text +from powercontext.builtin.inference import EmbeddingResult, InferenceUsage from powercontext.builtin.persistence.memory import RelationalMemoryBackend +from powercontext.builtin.persistence.memory_index import NoMemoryIndex from powercontext.builtin.persistence.sqlite import SQLiteConfig from powercontext.builtin.runtime import BuiltinConfig, open_builtin_contexts from powercontext.builtin.runtime.config import RuntimeConfig @@ -52,6 +62,39 @@ async def rerank(self, query, candidates, limit, /) -> MemoryRerankDecision: ) +_DENSE_PROFILE = EmbeddingProfile( + profile_id="dense-test-v1", + model="test:dense", + dimension=2, + distance="l2", + normalization="unit", +) + + +class _DenseEmbeddingModel: + profile = _DENSE_PROFILE + + async def embed(self, texts: tuple[str, ...], /) -> EmbeddingResult: + return EmbeddingResult(vectors=((0.2407121489724894, -0.9705965492093231),) * len(texts)) + + +class _RecordingVectorIndex(NoMemoryIndex): + capabilities = MemoryCapabilities(fts=False, vector=True, embedding_profile=_DENSE_PROFILE) + + def __init__(self) -> None: + self.projections: tuple[MemoryProjection, ...] = () + + async def replace( + self, + _connection: AsyncConnection, + _scope_id: str, + _memory_ref: ArtifactRef, + projections: tuple[MemoryProjection, ...], + /, + ) -> None: + self.projections = projections + + def test_memory_search_applies_injected_reranker_after_coarse_fusion() -> None: async def scenario() -> None: reranker = _SelectingReranker() @@ -78,6 +121,61 @@ async def scenario() -> None: asyncio.run(scenario()) +def test_memory_remember_accepts_a_dense_unit_embedding_after_service_normalization() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + index = _RecordingVectorIndex() + backend = RelationalMemoryBackend( + database=contexts.database, + scope_id="dense-vector", + artifacts=contexts.repositories.artifacts, + index=index, + ) + service = MemoryService(backend=backend, embedding_model=_DenseEmbeddingModel()) + + memory = await service.remember( + memory=None, + entries=(MemoryEntryInput(kind="preference", text="User prefers dense vectors."),), + mode="append", + ) + + assert memory is not None + assert len(index.projections) == 1 + assert index.projections[0].embedding == ( + 0.24071214897248938, + -0.9705965492093231, + ) + + asyncio.run(scenario()) + + +def test_memory_commit_rejects_a_finite_vector_outside_the_unit_norm_tolerance() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + index = _RecordingVectorIndex() + backend = RelationalMemoryBackend( + database=contexts.database, + scope_id="non-unit-vector", + artifacts=contexts.repositories.artifacts, + index=index, + ) + service = MemoryService(backend=backend, embedding_model=_DenseEmbeddingModel()) + plan = await service.plan_remember( + memory=None, + entries=(MemoryEntryInput(kind="preference", text="User prefers valid unit vectors."),), + mode="append", + ) + assert plan.commit is not None + projection = plan.commit.projections[0].model_copy(update={"embedding": (0.5, 0.5)}) + candidate = plan.commit.model_copy(update={"projections": (projection,)}) + + with pytest.raises(MemoryBackendConfigurationError): + async with backend.begin() as unit_of_work: + await unit_of_work.commit(candidate) + + asyncio.run(scenario()) + + def test_memory_entry_can_be_deactivated_and_reactivated_without_rewriting_content() -> None: async def scenario() -> None: async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: @@ -149,11 +247,21 @@ async def scenario() -> None: ) for version in versions ) - ) + ), + changes=tuple( + MemoryChange( + op="add", + entry_id=version.entry_id, + from_entry_version_id=None, + to_entry_version_id=version.entry_version_id, + ) + for version in versions + ), ) memory = Memory(artifact_id="memory", revision=1, content=content) projections = tuple( - MemoryProjection(entry_version=version, searchable_text="duplicate.") for version in versions + MemoryProjection(entry_version=version, searchable_text=analyze_text(version.text)) + for version in versions ) async with backend.begin() as unit_of_work: await unit_of_work.commit( @@ -180,3 +288,79 @@ async def scenario() -> None: assert entries[1] == versions[1] asyncio.run(scenario()) + + +def test_memory_organize_normalizes_an_inactive_entry_without_semantic_revision() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + backend = RelationalMemoryBackend( + database=contexts.database, + scope_id="normalize-inactive", + artifacts=contexts.repositories.artifacts, + index=contexts.index, + ) + content_hash = entry_content_hash( + kind="preference", + text="User prefers black tea.", + source_refs=(), + artifact_refs=(), + ) + original = MemoryEntryVersion( + memory_artifact_id="memory", + entry_id="preference", + entry_version_id="preference-v1", + version=1, + previous_version_id=None, + kind=" preference ", + text=" User prefers black tea. ", + entry_content_hash=content_hash, + created_in_revision=1, + ) + content = MemoryContent( + manifest=MemoryManifest( + entries=( + MemoryManifestEntry( + entry_id=original.entry_id, + entry_version_id=original.entry_version_id, + entry_content_hash=original.entry_content_hash, + state="active", + ), + ) + ), + changes=( + MemoryChange( + op="add", + entry_id=original.entry_id, + from_entry_version_id=None, + to_entry_version_id=original.entry_version_id, + ), + ), + ) + memory = Memory(artifact_id="memory", revision=1, content=content) + async with backend.begin() as unit_of_work: + await unit_of_work.commit( + MemoryCommit( + base=None, + memory=memory, + content_hash=memory_content_hash(content), + entry_versions=(original,), + projections=( + MemoryProjection(entry_version=original, searchable_text=analyze_text(original.text)), + ), + ) + ) + service = MemoryService(backend=backend) + inactive = await service.forget(memory, entries=(original,)) + + normalized = await service.organize(inactive, mode="normalize") + entry = (await service.entries(normalized))[0] + + assert normalized.revision == 3 + assert normalized.content.manifest.entries[0].state == "inactive" + assert normalized.content.changes[0].op == "revise" + assert normalized.content.changes[0].reason == "normalize" + assert entry.kind == "preference" + assert entry.text == "User prefers black tea." + assert entry.previous_version_id == original.entry_version_id + + asyncio.run(scenario()) diff --git a/tests/builtin/persistence/test_memory.py b/tests/builtin/persistence/test_memory.py index 86a96b8b8..7644f2538 100644 --- a/tests/builtin/persistence/test_memory.py +++ b/tests/builtin/persistence/test_memory.py @@ -15,13 +15,21 @@ from __future__ import annotations import asyncio +import sqlite3 +from types import SimpleNamespace +from typing import cast +from unittest.mock import AsyncMock +import pytest from sqlalchemy import BigInteger, DateTime, Integer, String from sqlalchemy.dialects import mysql +from sqlalchemy.ext.asyncio import AsyncConnection from sqlalchemy.schema import CreateTable, ForeignKeyConstraint, PrimaryKeyConstraint, UniqueConstraint from powercontext.builtin.artifacts.memory import MemoryEntryInput -from powercontext.builtin.persistence.sqlite import SQLiteConfig +from powercontext.builtin.artifacts.memory.errors import MemoryBackendConfigurationError +from powercontext.builtin.persistence.memory_schema import ensure_memory_entry_version_scope_identity +from powercontext.builtin.persistence.sqlite import SQLiteConfig, SQLiteProfile from powercontext.builtin.persistence.tables import ( MEMORY_ENTRY_HEADS_TABLE, MEMORY_ENTRY_VERSIONS_TABLE, @@ -30,6 +38,7 @@ from powercontext.builtin.sources import ContentCapture, ContentSource _INNODB_MAX_INDEX_BYTES = 3072 +_SCOPE_VERSION_INDEX = "uq_pc_memory_entry_versions_scope_version" class _UnbudgetedColumnTypeError(TypeError): @@ -71,6 +80,74 @@ def test_memory_schema_is_mysql_compilable_and_respects_key_and_payload_limits() assert max(budgets) == 2560 assert all(budget < _INNODB_MAX_INDEX_BYTES for budget in budgets) + scope_version_indexes = { + tuple(column.name for column in index.columns) for index in MEMORY_ENTRY_VERSIONS_TABLE.indexes if index.unique + } + assert ("scope_id", "entry_version_id") in scope_version_indexes + + +def test_sqlite_startup_adds_scope_global_version_identity_to_an_existing_table(tmp_path) -> None: + async def scenario() -> None: + database = tmp_path / "legacy-memory.db" + with sqlite3.connect(database) as connection: + connection.execute( + "CREATE TABLE pc_memory_entry_versions (scope_id TEXT NOT NULL, entry_version_id TEXT NOT NULL)" + ) + + async with ( + SQLiteProfile.open( + SQLiteConfig(url=f"sqlite+aiosqlite:///{database}"), + tables=(MEMORY_ENTRY_VERSIONS_TABLE,), + ) as profile, + profile.database.transaction() as connection, + ): + indexes = (await connection.exec_driver_sql("PRAGMA index_list('pc_memory_entry_versions')")).all() + + assert _SCOPE_VERSION_INDEX in {str(row[1]) for row in indexes} + + asyncio.run(scenario()) + + +def test_sqlite_startup_rejects_duplicate_scope_global_version_identities(tmp_path) -> None: + async def scenario() -> None: + database = tmp_path / "duplicate-memory.db" + with sqlite3.connect(database) as connection: + connection.execute( + "CREATE TABLE pc_memory_entry_versions (scope_id TEXT NOT NULL, entry_version_id TEXT NOT NULL)" + ) + connection.executemany( + "INSERT INTO pc_memory_entry_versions (scope_id, entry_version_id) VALUES (?, ?)", + (("scope", "duplicate"), ("scope", "duplicate")), + ) + + with pytest.raises(MemoryBackendConfigurationError): + async with SQLiteProfile.open( + SQLiteConfig(url=f"sqlite+aiosqlite:///{database}"), + tables=(MEMORY_ENTRY_VERSIONS_TABLE,), + ): + pass + + asyncio.run(scenario()) + + +def test_oceanbase_startup_adds_scope_global_version_identity_to_an_existing_table() -> None: + async def scenario() -> None: + connection = SimpleNamespace( + dialect=SimpleNamespace(name="mysql"), + scalar=AsyncMock(return_value=0), + execute=AsyncMock(return_value=SimpleNamespace(first=lambda: None)), + exec_driver_sql=AsyncMock(), + ) + + await ensure_memory_entry_version_scope_identity(cast(AsyncConnection, connection)) + + connection.exec_driver_sql.assert_awaited_once_with( + "CREATE UNIQUE INDEX uq_pc_memory_entry_versions_scope_version " + "ON pc_memory_entry_versions (scope_id, entry_version_id)" + ) + + asyncio.run(scenario()) + def test_sqlite_memory_backend_commits_authoritative_history_and_fts() -> None: async def scenario() -> None: diff --git a/tests/builtin/persistence/test_memory_integrity.py b/tests/builtin/persistence/test_memory_integrity.py new file mode 100644 index 000000000..d73f5533d --- /dev/null +++ b/tests/builtin/persistence/test_memory_integrity.py @@ -0,0 +1,443 @@ +# Copyright (c) 2026 OceanBase. +# +# 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. + +from __future__ import annotations + +import asyncio + +import pytest + +from powercontext.builtin.artifacts.memory import ( + InvalidMemoryCitationError, + Memory, + MemoryChange, + MemoryCommit, + MemoryContent, + MemoryEntryInput, + MemoryEntryVersion, + MemoryHit, + MemoryManifest, + MemoryManifestEntry, + MemoryProjection, + MemoryService, +) +from powercontext.builtin.artifacts.memory.canonical import entry_content_hash, memory_content_hash +from powercontext.builtin.artifacts.memory.errors import MemoryBackendConfigurationError +from powercontext.builtin.artifacts.search import analyze_text +from powercontext.builtin.persistence.memory import RelationalMemoryBackend +from powercontext.builtin.persistence.sqlite import SQLiteConfig +from powercontext.builtin.persistence.tables import MEMORY_ENTRY_HEADS_TABLE, MEMORY_ENTRY_VERSIONS_TABLE +from powercontext.builtin.runtime import BuiltinConfig, open_builtin_contexts +from powercontext.builtin.runtime.relational import RelationalContexts +from powercontext.errors import ArtifactNotFoundError + + +def _entry( + *, + artifact_id: str = "memory", + entry_id: str = "preference", + entry_version_id: str = "preference-v1", + version: int = 1, + previous_version_id: str | None = None, + text: str = "User prefers black tea.", + content_hash: str | None = None, + created_in_revision: int = 1, +) -> MemoryEntryVersion: + return MemoryEntryVersion( + memory_artifact_id=artifact_id, + entry_id=entry_id, + entry_version_id=entry_version_id, + version=version, + previous_version_id=previous_version_id, + kind="preference", + text=text, + entry_content_hash=( + entry_content_hash(kind="preference", text=text, source_refs=(), artifact_refs=()) + if content_hash is None + else content_hash + ), + created_in_revision=created_in_revision, + ) + + +def _initial_commit(*entries: MemoryEntryVersion) -> MemoryCommit: + ordered = tuple(sorted(entries, key=lambda entry: entry.entry_id.encode("utf-8"))) + content = MemoryContent( + manifest=MemoryManifest( + entries=tuple( + MemoryManifestEntry( + entry_id=entry.entry_id, + entry_version_id=entry.entry_version_id, + entry_content_hash=entry.entry_content_hash, + state="active", + ) + for entry in ordered + ) + ), + changes=tuple( + MemoryChange( + op="add", + entry_id=entry.entry_id, + from_entry_version_id=None, + to_entry_version_id=entry.entry_version_id, + ) + for entry in ordered + ), + ) + memory = Memory(artifact_id=ordered[0].memory_artifact_id, revision=1, content=content) + return MemoryCommit( + base=None, + memory=memory, + content_hash=memory_content_hash(content), + entry_versions=ordered, + projections=tuple( + MemoryProjection(entry_version=entry, searchable_text=analyze_text(entry.text)) for entry in ordered + ), + ) + + +def _revision_commit( + base: Memory, + current: tuple[MemoryEntryVersion, ...], + replacement: MemoryEntryVersion, +) -> MemoryCommit: + versions = {entry.entry_id: entry for entry in current} + previous = versions[replacement.entry_id] + versions[replacement.entry_id] = replacement + manifest = {item.entry_id: item for item in base.content.manifest.entries} + manifest[replacement.entry_id] = MemoryManifestEntry( + entry_id=replacement.entry_id, + entry_version_id=replacement.entry_version_id, + entry_content_hash=replacement.entry_content_hash, + state=manifest[replacement.entry_id].state, + ) + content = MemoryContent( + manifest=MemoryManifest( + entries=tuple(sorted(manifest.values(), key=lambda item: item.entry_id.encode("utf-8"))) + ), + changes=( + MemoryChange( + op="revise", + entry_id=replacement.entry_id, + from_entry_version_id=previous.entry_version_id, + to_entry_version_id=replacement.entry_version_id, + ), + ), + ) + memory = Memory(artifact_id=base.artifact_id, revision=base.revision + 1, content=content) + return MemoryCommit( + base=base, + memory=memory, + content_hash=memory_content_hash(content), + entry_versions=(replacement,), + projections=tuple( + MemoryProjection(entry_version=entry, searchable_text=analyze_text(entry.text)) + for entry in sorted(versions.values(), key=lambda entry: entry.entry_id.encode("utf-8")) + ), + ) + + +def _backend(contexts: RelationalContexts, scope_id: str) -> RelationalMemoryBackend: + return RelationalMemoryBackend( + database=contexts.database, + scope_id=scope_id, + artifacts=contexts.repositories.artifacts, + index=contexts.index, + ) + + +async def _commit(backend: RelationalMemoryBackend, value: MemoryCommit) -> Memory: + async with backend.begin() as unit_of_work: + return await unit_of_work.commit(value) + + +def test_memory_commit_rejects_incomplete_or_mismatched_revision_atomically() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + valid = _initial_commit(_entry()) + original = valid.entry_versions[0] + fake_predecessor = original.model_copy(update={"previous_version_id": "does-not-exist"}) + tampered = original.model_copy(update={"text": "tampered searchable body."}) + extra = _entry(entry_id="extra", entry_version_id="extra-v1") + mismatched_item = valid.memory.content.manifest.entries[0].model_copy( + update={"entry_content_hash": "0" * 64} + ) + mismatched_content = valid.memory.content.model_copy( + update={"manifest": MemoryManifest(entries=(mismatched_item,))} + ) + invalid = ( + valid.model_copy(update={"entry_versions": ()}), + valid.model_copy(update={"entry_versions": (original, original)}), + valid.model_copy(update={"entry_versions": (original, extra)}), + valid.model_copy(update={"projections": (valid.projections[0], valid.projections[0])}), + valid.model_copy( + update={ + "entry_versions": (fake_predecessor,), + "projections": ( + MemoryProjection( + entry_version=fake_predecessor, + searchable_text=analyze_text(fake_predecessor.text), + ), + ), + } + ), + valid.model_copy( + update={ + "entry_versions": (tampered,), + "projections": ( + MemoryProjection(entry_version=tampered, searchable_text=analyze_text(tampered.text)), + ), + } + ), + valid.model_copy( + update={ + "projections": (valid.projections[0].model_copy(update={"searchable_text": "not canonical"}),) + } + ), + valid.model_copy( + update={ + "projections": ( + valid.projections[0].model_copy( + update={"embedding": (1.0,), "embedding_content_hash": "0" * 64} + ), + ) + } + ), + valid.model_copy( + update={ + "memory": valid.memory.model_copy(update={"content": mismatched_content}), + "content_hash": memory_content_hash(mismatched_content), + } + ), + ) + + for index, candidate in enumerate(invalid): + backend = _backend(contexts, f"invalid-{index}") + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, candidate) + with pytest.raises(ArtifactNotFoundError): + await backend.latest("memory") + + asyncio.run(scenario()) + + +def test_memory_commit_requires_the_direct_same_entry_predecessor() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + first_a = _entry(entry_id="entry-a", entry_version_id="entry-a-v1") + first_b = _entry(entry_id="entry-b", entry_version_id="entry-b-v1", text="User prefers green tea.") + initial = _initial_commit(first_a, first_b) + + for scope_id, predecessor in (("missing", "does-not-exist"), ("cross-entry", "entry-b-v1")): + backend = _backend(contexts, scope_id) + base = await _commit(backend, initial) + replacement = _entry( + entry_id="entry-a", + entry_version_id="entry-a-v2", + version=2, + previous_version_id=predecessor, + text="User prefers oolong tea.", + created_in_revision=2, + ) + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, _revision_commit(base, (first_a, first_b), replacement)) + assert await backend.latest("memory") == base + + backend = _backend(contexts, "cross-memory") + base = await _commit(backend, initial) + await _commit( + backend, + _initial_commit( + _entry( + artifact_id="other-memory", + entry_id="other-entry", + entry_version_id="other-entry-v1", + ) + ), + ) + cross_memory = _entry( + entry_id="entry-a", + entry_version_id="entry-a-v2", + version=2, + previous_version_id="other-entry-v1", + text="User prefers oolong tea.", + created_in_revision=2, + ) + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, _revision_commit(base, (first_a, first_b), cross_memory)) + assert await backend.latest("memory") == base + + backend = _backend(contexts, "duplicate-version-id") + await _commit(backend, initial) + collision = _initial_commit( + _entry( + artifact_id="other-memory", + entry_id="other-entry", + entry_version_id="entry-a-v1", + ) + ) + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, collision) + with pytest.raises(ArtifactNotFoundError): + await backend.latest("other-memory") + + backend = _backend(contexts, "skipped") + base = await _commit(backend, initial) + second_a = _entry( + entry_id="entry-a", + entry_version_id="entry-a-v2", + version=2, + previous_version_id="entry-a-v1", + text="User prefers oolong tea.", + created_in_revision=2, + ) + second = await _commit(backend, _revision_commit(base, (first_a, first_b), second_a)) + skipped = _entry( + entry_id="entry-a", + entry_version_id="entry-a-v3", + version=3, + previous_version_id="entry-a-v1", + text="User prefers white tea.", + created_in_revision=3, + ) + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, _revision_commit(second, (second_a, first_b), skipped)) + assert await backend.latest("memory") == second + + asyncio.run(scenario()) + + +def test_memory_commit_rejects_corrupted_history_already_referenced_by_base() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + backend = _backend(contexts, "corrupted-base-history") + original = _entry() + base = await _commit(backend, _initial_commit(original)) + second = _entry( + entry_version_id="preference-v2", + version=2, + previous_version_id=original.entry_version_id, + text="User prefers oolong tea.", + created_in_revision=2, + ) + head = await _commit(backend, _revision_commit(base, (original,), second)) + async with contexts.database.transaction() as connection: + await connection.execute( + MEMORY_ENTRY_VERSIONS_TABLE + .update() + .where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == "corrupted-base-history", + MEMORY_ENTRY_VERSIONS_TABLE.c.memory_artifact_id == base.artifact_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.entry_version_id == "preference-v1", + ) + .values(previous_version_id="does-not-exist") + ) + + service = MemoryService(backend=backend) + with pytest.raises(MemoryBackendConfigurationError): + await service.remember( + memory=head, + entries=(MemoryEntryInput(kind="fact", text="An unrelated fact."),), + mode="append", + ) + assert await backend.latest(base.artifact_id) == head + + asyncio.run(scenario()) + + +def test_memory_commit_rejects_semantic_revision_of_an_inactive_entry() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + backend = _backend(contexts, "inactive-revision") + original = _entry() + base = await _commit(backend, _initial_commit(original)) + service = MemoryService(backend=backend) + inactive = await service.forget(base, entries=(original,), reason="paused") + replacement = _entry( + entry_version_id="preference-v2", + version=2, + previous_version_id=original.entry_version_id, + text="User prefers green tea.", + created_in_revision=3, + ) + candidate = _revision_commit(inactive, (original,), replacement).model_copy(update={"projections": ()}) + + with pytest.raises(MemoryBackendConfigurationError): + await _commit(backend, candidate) + assert await backend.latest(base.artifact_id) == inactive + + asyncio.run(scenario()) + + +def test_corrupted_authoritative_entry_is_rejected_by_entries_search_and_expand() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + for scope_id, values in ( + ("body-corruption", {"text": "tampered searchable body."}), + ("declared-hash-corruption", {"entry_content_hash": "0" * 64}), + ): + backend = _backend(contexts, scope_id) + memory = await _commit(backend, _initial_commit(_entry())) + async with contexts.database.transaction() as connection: + await connection.execute( + MEMORY_ENTRY_VERSIONS_TABLE + .update() + .where( + MEMORY_ENTRY_VERSIONS_TABLE.c.scope_id == scope_id, + MEMORY_ENTRY_VERSIONS_TABLE.c.memory_artifact_id == memory.artifact_id, + ) + .values(**values) + ) + service = MemoryService(backend=backend) + hit = MemoryHit( + memory_ref=memory.as_ref(), + entry_id="preference", + entry_version_id="preference-v1", + text="User prefers black tea.", + score=1.0, + matched_by=("fts",), + ) + + with pytest.raises(InvalidMemoryCitationError): + await service.entries(memory) + with pytest.raises(InvalidMemoryCitationError): + await service.search("black tea", memories=(memory,), mode="fts") + with pytest.raises(InvalidMemoryCitationError): + await service.expand((hit,)) + + asyncio.run(scenario()) + + +def test_search_rejects_noncanonical_rebuildable_projection() -> None: + async def scenario() -> None: + async with open_builtin_contexts(BuiltinConfig(database=SQLiteConfig())) as contexts: + backend = _backend(contexts, "projection-corruption") + memory = await _commit(backend, _initial_commit(_entry())) + async with contexts.database.transaction() as connection: + await connection.execute( + MEMORY_ENTRY_HEADS_TABLE + .update() + .where( + MEMORY_ENTRY_HEADS_TABLE.c.scope_id == "projection-corruption", + MEMORY_ENTRY_HEADS_TABLE.c.memory_artifact_id == memory.artifact_id, + ) + .values(searchable_text="tampered projection") + ) + + service = MemoryService(backend=backend) + assert tuple(entry.text for entry in await service.entries(memory)) == ("User prefers black tea.",) + with pytest.raises(InvalidMemoryCitationError): + await service.search("black tea", memories=(memory,), mode="fts") + + asyncio.run(scenario()) diff --git a/tests/e2e/test_memory_search_concurrency.py b/tests/e2e/test_memory_search_concurrency.py index ba8328f30..a19c57f6d 100644 --- a/tests/e2e/test_memory_search_concurrency.py +++ b/tests/e2e/test_memory_search_concurrency.py @@ -27,21 +27,34 @@ from powercontext.builtin.artifacts.memory import ( EmbeddingProfile, + Memory, + MemoryChange, + MemoryCommit, + MemoryContent, MemoryEntryInput, + MemoryEntryVersion, MemoryHit, + MemoryManifest, + MemoryManifestEntry, + MemoryProjection, MemoryRerankDecision, MemorySearchMode, ) +from powercontext.builtin.artifacts.memory.canonical import entry_content_hash, memory_content_hash +from powercontext.builtin.artifacts.memory.errors import MemoryBackendConfigurationError +from powercontext.builtin.artifacts.search import analyze_text from powercontext.builtin.inference import EmbeddingResult, InferenceUsage +from powercontext.builtin.persistence.memory import RelationalMemoryBackend from powercontext.builtin.persistence.oceanbase import OceanBaseConfig from powercontext.builtin.persistence.sqlite import SQLiteConfig from powercontext.builtin.runtime import ( BuiltinConfig, RememberMemoryRequest, SearchMemoryRequest, + open_builtin_contexts, open_builtin_runtime, ) -from powercontext.errors import RevisionConflictError +from powercontext.errors import ArtifactNotFoundError, RevisionConflictError DatabaseKind = Literal["sqlite", "oceanbase"] TIMEOUT_SECONDS = 15 @@ -57,6 +70,7 @@ "vector": ("vector",), "hybrid": ("fts", "vector"), } +OCEANBASE_URL = os.environ.get("POWERCONTEXT_TEST_OCEANBASE_URL") class _KeywordEmbeddingModel: @@ -227,6 +241,128 @@ async def advance_head_before_search(_self: Any, *args: Any, **kwargs: Any) -> A asyncio.run(scenario()) +@pytest.mark.skipif( + not OCEANBASE_URL, + reason="set POWERCONTEXT_TEST_OCEANBASE_URL to a dedicated OceanBase MySQL-mode test database", +) +def test_oceanbase_memory_version_identity_race_is_database_enforced() -> None: + async def scenario() -> None: + assert OCEANBASE_URL is not None + scope_id = f"concurrent-version-identity-{uuid4()}" + async with open_builtin_contexts( + BuiltinConfig(database=OceanBaseConfig(url=SecretStr(OCEANBASE_URL))) + ) as contexts: + backend = RelationalMemoryBackend( + database=contexts.database, + scope_id=scope_id, + artifacts=contexts.repositories.artifacts, + index=contexts.index, + ) + mutable_backend: Any = backend + original = mutable_backend._commit_versions + both_checked = asyncio.Event() + release = asyncio.Event() + arrivals = 0 + arrival_lock = asyncio.Lock() + + async def pause_after_identity_read( + _self: RelationalMemoryBackend, + connection: Any, + value: MemoryCommit, + ) -> dict[str, MemoryEntryVersion]: + nonlocal arrivals + versions = await original(connection, value) + async with arrival_lock: + arrivals += 1 + if arrivals == 2: + both_checked.set() + await release.wait() + return versions + + mutable_backend._commit_versions = MethodType(pause_after_identity_read, backend) + commits = ( + _colliding_initial_commit("memory-a", "entry-a"), + _colliding_initial_commit("memory-b", "entry-b"), + ) + pending = tuple(asyncio.create_task(_commit_memory(backend, commit)) for commit in commits) + results: list[Memory | BaseException] = [] + try: + await asyncio.wait_for(both_checked.wait(), timeout=TIMEOUT_SECONDS) + release.set() + results = await asyncio.wait_for( + asyncio.gather(*pending, return_exceptions=True), + timeout=TIMEOUT_SECONDS, + ) + finally: + release.set() + for task in pending: + if not task.done(): + task.cancel() + with suppress(asyncio.CancelledError): + await task + + assert len(tuple(result for result in results if isinstance(result, Memory))) == 1 + failures = tuple(result for result in results if isinstance(result, BaseException)) + assert len(failures) == 1 + assert isinstance(failures[0], MemoryBackendConfigurationError) + stored = [] + for artifact_id in ("memory-a", "memory-b"): + with suppress(ArtifactNotFoundError): + stored.append(await backend.latest(artifact_id)) + assert len(stored) == 1 + + asyncio.run(scenario()) + + +def _colliding_initial_commit(memory_id: str, entry_id: str) -> MemoryCommit: + text = f"Concurrent identity for {memory_id}." + content_hash = entry_content_hash(kind="fact", text=text, source_refs=(), artifact_refs=()) + version = MemoryEntryVersion( + memory_artifact_id=memory_id, + entry_id=entry_id, + entry_version_id="shared-concurrent-version", + version=1, + previous_version_id=None, + kind="fact", + text=text, + entry_content_hash=content_hash, + created_in_revision=1, + ) + content = MemoryContent( + manifest=MemoryManifest( + entries=( + MemoryManifestEntry( + entry_id=entry_id, + entry_version_id=version.entry_version_id, + entry_content_hash=content_hash, + state="active", + ), + ) + ), + changes=( + MemoryChange( + op="add", + entry_id=entry_id, + from_entry_version_id=None, + to_entry_version_id=version.entry_version_id, + ), + ), + ) + memory = Memory(artifact_id=memory_id, revision=1, content=content) + return MemoryCommit( + base=None, + memory=memory, + content_hash=memory_content_hash(content), + entry_versions=(version,), + projections=(MemoryProjection(entry_version=version, searchable_text=analyze_text(text)),), + ) + + +async def _commit_memory(backend: RelationalMemoryBackend, commit: MemoryCommit) -> Memory: + async with backend.begin() as unit_of_work: + return await unit_of_work.commit(commit) + + def _database_config( database_kind: DatabaseKind, mode: MemorySearchMode,