Memory
Auditable long-term memory lifecycle for conversational RAG agents.
"""Auditable long-term memory lifecycle for conversational RAG agents."""
import math
from dataclasses import dataclass, replace
from typing import Iterable, List, Mapping, Optional, Sequence, Set, Tuple
from .temporal import parse_instant
from .text import content_terms, jaccard
@dataclass(frozen=True)
class MemoryRecord:
"""Externalized memory with provenance, permissions, and update semantics."""
identifier: str
text: str
created_at: str
source_turn: str
kind: str = "episodic"
importance: float = 0.5
principals: Tuple[str, ...] = ()
supersedes: Tuple[str, ...] = ()
expires_at: str = ""
deleted_at: str = ""
last_accessed_at: str = ""
access_count: int = 0
def __post_init__(self) -> None:
if not self.identifier or not self.text.strip() or not self.source_turn:
raise ValueError("memory identifier, text, and source turn must be non-empty")
if self.kind not in {"episodic", "semantic", "procedural", "profile", "summary"}:
raise ValueError("unsupported memory kind")
if not 0.0 <= self.importance <= 1.0:
raise ValueError("importance must be between zero and one")
if self.access_count < 0:
raise ValueError("access_count must be non-negative")
parse_instant(self.created_at)
if self.expires_at:
parse_instant(self.expires_at)
if self.deleted_at:
parse_instant(self.deleted_at)
def accessible_to(self, principals: Iterable[str]) -> bool:
allowed = set(self.principals)
return not allowed or bool(allowed & set(principals))
def active_at(self, instant: str) -> bool:
point = parse_instant(instant)
if parse_instant(self.created_at) > point:
return False
if self.deleted_at and parse_instant(self.deleted_at) <= point:
return False
if self.expires_at and parse_instant(self.expires_at) <= point:
return False
return True
@dataclass(frozen=True)
class MemoryHit:
"""Retrieved memory plus decomposed relevance features."""
record: MemoryRecord
score: float
lexical: float
recency: float
importance: float
@dataclass(frozen=True)
class WriteDecision:
"""Memory-write policy decision with inspectable reasons."""
write: bool
novelty: float
reasons: Tuple[str, ...]
class MemoryStore:
"""Append/update/delete store that keeps memory lifecycle observable."""
def __init__(self, records: Sequence[MemoryRecord] = ()) -> None:
self._records: List[MemoryRecord] = []
self._ids: Set[str] = set()
for record in records:
self.append(record)
@property
def records(self) -> Tuple[MemoryRecord, ...]:
return tuple(self._records)
def append(self, record: MemoryRecord) -> None:
if record.identifier in self._ids:
raise ValueError("duplicate memory identifier: " + record.identifier)
missing = set(record.supersedes) - self._ids
if missing:
raise ValueError("memory supersedes unknown records: " + ", ".join(sorted(missing)))
self._records.append(record)
self._ids.add(record.identifier)
def delete(self, identifier: str, deleted_at: str) -> None:
"""Logical deletion retained in the audit log."""
parse_instant(deleted_at)
for index, record in enumerate(self._records):
if record.identifier == identifier:
if record.deleted_at:
raise ValueError("memory is already deleted")
self._records[index] = replace(record, deleted_at=deleted_at)
return
raise KeyError(identifier)
def retrieve(
self,
query: str,
at: str,
principals: Iterable[str] = (),
k: int = 5,
recency_half_life_days: float = 30.0,
) -> Tuple[MemoryHit, ...]:
"""Retrieve active, authorized, non-superseded memories."""
if k <= 0:
return ()
if recency_half_life_days <= 0:
raise ValueError("recency half-life must be positive")
point = parse_instant(at)
active = [
record
for record in self._records
if record.active_at(at) and record.accessible_to(principals)
]
superseded = {identifier for record in active for identifier in record.supersedes}
query_terms = set(content_terms(query))
hits = []
for record in active:
if record.identifier in superseded:
continue
terms = set(content_terms(record.text))
lexical = len(query_terms & terms) / max(1, len(query_terms))
age_days = max(0.0, (point - parse_instant(record.created_at)).total_seconds() / 86400)
recency = math.exp(-math.log(2.0) * age_days / recency_half_life_days)
score = 0.65 * lexical + 0.20 * record.importance + 0.15 * recency
if lexical > 0 or record.importance >= 0.8:
hits.append(MemoryHit(record, score, lexical, recency, record.importance))
hits.sort(key=lambda hit: (-hit.score, hit.record.identifier))
return tuple(hits[:k])
def mark_accessed(self, identifiers: Iterable[str], at: str) -> None:
"""Update access metadata without altering memory creation time."""
parse_instant(at)
requested = set(identifiers)
found: Set[str] = set()
for index, record in enumerate(self._records):
if record.identifier in requested:
self._records[index] = replace(
record,
last_accessed_at=at,
access_count=record.access_count + 1,
)
found.add(record.identifier)
missing = requested - found
if missing:
raise KeyError(", ".join(sorted(missing)))
def purgeable(self, at: str) -> Tuple[str, ...]:
"""IDs whose retention period expired or whose logical deletion is effective."""
point = parse_instant(at)
purge = []
for record in self._records:
expired = bool(record.expires_at and parse_instant(record.expires_at) <= point)
deleted = bool(record.deleted_at and parse_instant(record.deleted_at) <= point)
if expired or deleted:
purge.append(record.identifier)
return tuple(sorted(purge))
def write_decision(
candidate_text: str,
existing: Sequence[MemoryRecord],
importance: float,
durable_signal: bool = False,
novelty_threshold: float = 0.35,
) -> WriteDecision:
"""Decide whether a turn deserves durable memory.
A write is justified by novelty plus importance, or by an explicit durable
signal such as a user preference/correction. This prevents storing every
turn while keeping policy decisions explainable.
"""
if not 0.0 <= importance <= 1.0 or not 0.0 <= novelty_threshold <= 1.0:
raise ValueError("importance and novelty threshold must be between zero and one")
terms = set(content_terms(candidate_text))
similarity = max(
(jaccard(terms, set(content_terms(record.text))) for record in existing),
default=0.0,
)
novelty = 1.0 - similarity
reasons = []
if novelty >= novelty_threshold:
reasons.append("novel")
if importance >= 0.7:
reasons.append("important")
if durable_signal:
reasons.append("explicit-durable-signal")
write = durable_signal or (novelty >= novelty_threshold and importance >= 0.5)
if not write:
reasons.append("write-threshold-not-met")
return WriteDecision(write, novelty, tuple(reasons))
def consolidation_groups(
records: Sequence[MemoryRecord], threshold: float = 0.6
) -> Tuple[Tuple[str, ...], ...]:
"""Connected components of semantically overlapping memories.
This returns candidates for a summarizer; it deliberately does not generate
a lossy summary or delete originals. Human/audited consolidation can then
create a new record whose ``supersedes`` field preserves lineage.
"""
if not 0.0 <= threshold <= 1.0:
raise ValueError("threshold must be between zero and one")
terms: Mapping[str, Set[str]] = {
record.identifier: set(content_terms(record.text)) for record in records
}
adjacency = {record.identifier: set() for record in records}
for left_index, left in enumerate(records):
for right in records[left_index + 1 :]:
if jaccard(terms[left.identifier], terms[right.identifier]) >= threshold:
adjacency[left.identifier].add(right.identifier)
adjacency[right.identifier].add(left.identifier)
groups = []
seen: Set[str] = set()
for identifier in sorted(adjacency):
if identifier in seen:
continue
stack = [identifier]
component = []
while stack:
current = stack.pop()
if current in seen:
continue
seen.add(current)
component.append(current)
stack.extend(sorted(adjacency[current] - seen, reverse=True))
if len(component) > 1:
groups.append(tuple(sorted(component)))
return tuple(groups)