122 lines
4.7 KiB
Python
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)
|