"""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 from steward.bot.thread_key import ThreadKey logger = logging.getLogger(__name__) @dataclass class ThreadSummary: """A persisted summary of a flushed conversation 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. ``platform``/``scope``/``thread`` normalise the conversation identity across chat platforms (e.g. Telegram chat+thread, or a Matrix room). """ platform: str scope: str thread: str | None 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: parts = [self.platform, self.scope] if self.thread: parts.append(self.thread) return ":".join(parts) @property def thread_id(self) -> str: """Human-readable thread identifier for display (falls back to scope).""" return self.thread or self.scope @staticmethod def _extract_tags(data: dict[str, object]) -> list[str]: raw = data.get("tags") if not isinstance(raw, list): return [] return [str(t) for t in raw] @classmethod def from_dict(cls, data: dict[str, object]) -> ThreadSummary: """Deserialise from a raw dict, tolerating missing optional fields. Legacy records stored ``chat_id``/``thread_id`` (Telegram-only). Those are mapped to ``platform="telegram"``, ``scope=str(chat_id)`` and ``thread=str(thread_id)`` for backward compatibility. """ if "platform" in data: return cls( platform=str(data["platform"]), scope=str(data["scope"]), thread=str(data["thread"]) if data.get("thread") else None, summary=str(data["summary"]), message_count=int(str(data["message_count"])), flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())), tags=cls._extract_tags(data), ) return cls( platform="telegram", scope=str(data["chat_id"]), thread=str(data["thread_id"]), summary=str(data["summary"]), message_count=int(str(data["message_count"])), flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())), tags=cls._extract_tags(data), ) 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 self._migrate_legacy_keys(raw) except (json.JSONDecodeError, OSError): logger.warning( "Could not read thread memory store at %s; starting fresh", self._path ) return {} @staticmethod def _migrate_legacy_keys( raw: dict[str, object], ) -> dict[str, dict[str, object]]: """Convert legacy ``"chat_id:thread_id"`` keys to the platform-scoped format. Legacy records predate multi-platform support and stored keys as ``":"`` with ``chat_id``/``thread_id`` fields. These are migrated to ``"telegram::"`` so they remain addressable via :class:`~steward.bot.thread_key.ThreadKey`. """ migrated: dict[str, dict[str, object]] = {} for key, value in raw.items(): if not isinstance(value, dict): continue if "platform" in value: migrated[key] = value continue parts = str(key).split(":") if len(parts) == 2 and parts[0].lstrip("-").isdigit() and parts[1].isdigit(): new_key = f"telegram:{parts[0]}:{parts[1]}" migrated[new_key] = value else: migrated[key] = value return migrated 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, key: ThreadKey) -> ThreadSummary | None: """Return the stored summary for a conversation scope, or None if not found.""" raw = self._data.get(str(key)) 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 ThreadSummary._extract_tags(v)} ] return sorted(results, key=lambda s: s.flushed_at, reverse=True)