steward_mirror/steward/memory/thread_store.py

122 lines
4.7 KiB
Python

"""Persistent storage for flushed thread summaries (knowledge-base style)."""
from __future__ import annotations
import json
import logging
from dataclasses import asdict, dataclass, field
from datetime import UTC, datetime
from pathlib import Path
logger = logging.getLogger(__name__)
@dataclass
class ThreadSummary:
"""A persisted summary of a flushed Telegram message thread.
``tags`` is a list of short lowercase keywords extracted by the LLM at flush
time. They are used to index the knowledge base so summaries can be recalled
contextually without being kept permanently in the conversation context.
"""
chat_id: int
thread_id: int
summary: str
message_count: int
flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat())
tags: list[str] = field(default_factory=list)
@property
def key(self) -> str:
return f"{self.chat_id}:{self.thread_id}"
@classmethod
def from_dict(cls, data: dict[str, object]) -> ThreadSummary:
"""Deserialise from a raw dict, tolerating missing optional fields."""
return cls(
chat_id=int(data["chat_id"]), # type: ignore[arg-type]
thread_id=int(data["thread_id"]), # type: ignore[arg-type]
summary=str(data["summary"]),
message_count=int(data["message_count"]), # type: ignore[arg-type]
flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())),
tags=list(data.get("tags", [])), # type: ignore[arg-type]
)
def format_for_telegram(self) -> str:
"""Return a concise Telegram-formatted recall card."""
flushed = self.flushed_at[:19].replace("T", " ")
header = (
f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n"
f"_Flushed: {flushed} UTC \u2014 {self.message_count} messages_"
)
tag_line = f"\U0001f3f7 Tags: {', '.join(self.tags)}" if self.tags else ""
parts = [header, tag_line, self.summary] if tag_line else [header, self.summary]
return "\n\n".join(parts)
class ThreadMemoryStore:
"""JSON-backed knowledge-base store for flushed thread summaries.
Summaries are indexed by keyword tags so they can be recalled contextually
(via :meth:`search`) without being kept permanently in the LLM context.
Each flush overwrites the previous summary for the same (chat_id, thread_id) pair.
"""
def __init__(self, path: str | Path = "thread_memory.json") -> None:
self._path = Path(path)
self._data: dict[str, dict[str, object]] = self._load()
def _load(self) -> dict[str, dict[str, object]]:
if self._path.exists():
try:
raw = json.loads(self._path.read_text(encoding="utf-8"))
if isinstance(raw, dict):
return raw # type: ignore[return-value]
except (json.JSONDecodeError, OSError):
logger.warning(
"Could not read thread memory store at %s; starting fresh", self._path
)
return {}
def _save(self) -> None:
try:
self._path.write_text(
json.dumps(self._data, indent=2, ensure_ascii=False), encoding="utf-8"
)
except OSError:
logger.exception("Failed to write thread memory store to %s", self._path)
def save(self, summary: ThreadSummary) -> None:
"""Persist a thread summary, replacing any previous entry for this thread."""
self._data[summary.key] = asdict(summary)
self._save()
def get(self, chat_id: int, thread_id: int) -> ThreadSummary | None:
"""Return the stored summary for a thread, or None if not found."""
raw = self._data.get(f"{chat_id}:{thread_id}")
if raw is None:
return None
return ThreadSummary.from_dict(raw)
def all(self) -> list[ThreadSummary]:
"""Return all stored summaries, newest first."""
entries = [ThreadSummary.from_dict(v) for v in self._data.values()]
return sorted(entries, key=lambda s: s.flushed_at, reverse=True)
def search(self, query: str) -> list[ThreadSummary]:
"""Return summaries whose tags overlap with words in *query*, newest first.
The match is case-insensitive and word-based. Summaries without tags
are not returned even if the query is broad.
"""
query_words = {w.lower() for w in query.split() if w}
if not query_words:
return []
results = [
ThreadSummary.from_dict(v)
for v in self._data.values()
if query_words & {t.lower() for t in v.get("tags", [])} # type: ignore[union-attr]
]
return sorted(results, key=lambda s: s.flushed_at, reverse=True)