"""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)