From 76777c98eb2c01f65156760f027a835c87e505d5 Mon Sep 17 00:00:00 2001 From: Andrew Ridgway Date: Tue, 18 Aug 2026 21:38:01 +1000 Subject: [PATCH] feat: add Matrix appservice bot and platform-agnostic conversation core Introduce a shared ConversationService (steward/bot/core.py) that owns the LLM call, history, knowledge-base search, and thread-memory keying behind a normalized ThreadKey, so both Telegram and Matrix drive the same pipeline. - Add steward/bot/matrix.py: a mautrix-python appservice bot that receives Synapse transactions and replies via the client-server API. - Refactor telegram.py handlers into thin wrappers over ConversationService. - Generalize ThreadMemoryStore/ThreadSummary to platform-scoped keys with legacy chat_id:thread_id migration. - Add a matrix config section (homeserver, tokens, room/user allowlists). - Rewrite main.py as async, starting Telegram and/or Matrix on one event loop. - Add mautrix>=0.21.0 dependency. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus --- pyproject.toml | 3 +- steward/bot/core.py | 179 ++++++++++++++++++++++++ steward/bot/matrix.py | 126 +++++++++++++++++ steward/bot/telegram.py | 245 ++++++++------------------------- steward/bot/thread_key.py | 24 ++++ steward/config.py | 34 +++++ steward/config_schema.yaml | 33 +++++ steward/main.py | 104 ++++++++++---- steward/memory/thread_store.py | 91 ++++++++++-- tests/test_bot.py | 58 ++++---- tests/test_thread_memory.py | 200 +++++++++++++-------------- 11 files changed, 749 insertions(+), 348 deletions(-) create mode 100644 steward/bot/core.py create mode 100644 steward/bot/matrix.py create mode 100644 steward/bot/thread_key.py diff --git a/pyproject.toml b/pyproject.toml index a3ad2db..03f7127 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,6 +17,7 @@ dependencies = [ "pydantic>=2.7", "pydantic-settings>=2.3", "omegaconf>=2.3", + "mautrix>=0.21.0", ] [project.optional-dependencies] @@ -31,7 +32,7 @@ dev = [ ] [project.scripts] -steward = "steward.main:main" +steward = "steward.main:run" [tool.setuptools.packages.find] where = ["."] diff --git a/steward/bot/core.py b/steward/bot/core.py new file mode 100644 index 0000000..744f57d --- /dev/null +++ b/steward/bot/core.py @@ -0,0 +1,179 @@ +"""Platform-agnostic conversation core for Steward. + +This module owns the shared message pipeline (LLM call, history management, +knowledge-base search, and thread-memory keying) so that both the Telegram +and Matrix adapters can drive the same behaviour without duplicating logic. +""" + +from __future__ import annotations + +import logging +from collections.abc import Callable +from typing import Any + +from steward.bot.thread_key import ThreadKey +from steward.config import Settings +from steward.llm.client import LLMClient +from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary +from steward.tools.client import ToolClient + +logger = logging.getLogger(__name__) + +_MAX_HISTORY = 20 + +_FLUSH_SYSTEM_PROMPT = ( + "You are Steward. The following is a complete conversation thread. " + "Produce a concise but comprehensive summary that captures:\n" + "- The main topics discussed\n" + "- Key decisions or conclusions reached\n" + "- Any outstanding actions or open questions\n" + "- Important context that would help recall this conversation later\n\n" + "Be precise. Omit pleasantries." +) + +_TAGS_SYSTEM_PROMPT = ( + "You are a keyword tagger for a knowledge base. " + "Extract 5-8 short, lowercase keyword tags from the following conversation summary. " + "Tags should represent the main topics, entities, and concepts discussed. " + "Return ONLY a comma-separated list of tags with no other text or punctuation. " + "Example output: api design, authentication, database schema, user roles, caching" +) + +_KB_CONTEXT_HEADER = ( + "The following are relevant past conversation summaries from your knowledge base. " + "Use them as background context if they relate to the current question, " + "but do not repeat their contents unless directly asked." +) + + +class ConversationService: + """Owns the shared message pipeline used by all platform adapters.""" + + def __init__( + self, + settings: Settings, + llm: LLMClient, + thread_store: ThreadMemoryStore, + tool_client: ToolClient | None = None, + ) -> None: + self._settings = settings + self._llm = llm + self._store = thread_store + self._tool_client = tool_client + self._histories: dict[ThreadKey, list[dict[str, Any]]] = {} + + def _history_for(self, key: ThreadKey) -> list[dict[str, Any]]: + return self._histories.setdefault(key, []) + + @property + def llm(self) -> LLMClient: + """The shared LLM client (used by platform adapters for ad-hoc calls).""" + return self._llm + + def _with_kb_context( + self, + history: list[dict[str, Any]], + relevant: list[ThreadSummary], + ) -> list[dict[str, Any]]: + if not relevant: + return history + snippets = [f"[Thread {s.thread_id}] {s.summary[:400]}" for s in relevant[:3]] + kb_msg = _KB_CONTEXT_HEADER + "\n\n" + "\n\n---\n\n".join(snippets) + return [{"role": "system", "content": kb_msg}, *history] + + async def process_message( + self, + key: ThreadKey, + user_id: str, + text: str, + system_prompt: str, + history_cap: int | None = None, + history_formatter: Callable[[str], str] | None = None, + ) -> str: + """Process a user message and return the raw assistant reply text. + + The reply is the raw LLM output; platform adapters are responsible for + parsing and rendering it (e.g. Telegram polls/multi-message markers). + ``history_formatter`` transforms the raw reply into the text stored in + conversation history (defaults to the raw reply). + """ + history = self._history_for(key) + call_history = self._with_kb_context(history, self._store.search(text)) + if self._tool_client is not None: + reply = await self._llm.chat_with_tools( + text, + self._tool_client, + history=call_history, + system_prompt=system_prompt, + ) + else: + reply = await self._llm.chat(text, history=call_history, system_prompt=system_prompt) + history.append({"role": "user", "content": text}) + stored_reply = history_formatter(reply) if history_formatter else reply + history.append({"role": "assistant", "content": stored_reply}) + if history_cap is not None and len(history) > history_cap: + self._histories[key] = history[-history_cap:] + return reply + + def clear_history(self, key: ThreadKey) -> None: + """Wipe the in-memory conversation history for a scope.""" + self._histories.pop(key, None) + + def has_history(self, key: ThreadKey) -> bool: + """Return True if the scope has any in-memory conversation history.""" + return bool(self._histories.get(key)) + + async def flush_history(self, key: ThreadKey) -> ThreadSummary | None: + """Summarise the scope's history, persist it, and compress in-memory history. + + Returns the persisted :class:`ThreadSummary`, or ``None`` if there was no + history to flush. + """ + history = self._history_for(key) + if not history: + return None + + transcript_lines = [] + for msg in history: + role_label = "User" if msg["role"] == "user" else "Steward" + transcript_lines.append(f"{role_label}: {msg['content']}") + transcript = "\n".join(transcript_lines) + + summary_text = await self._llm.chat( + f"Thread transcript:\n\n{transcript}", + system_prompt=_FLUSH_SYSTEM_PROMPT, + ) + + tags_raw = await self._llm.chat( + f"Summary to tag:\n\n{summary_text}", + system_prompt=_TAGS_SYSTEM_PROMPT, + ) + tags = [t.strip().lower() for t in tags_raw.split(",") if t.strip()][:8] + + message_count = sum(1 for m in history if m["role"] == "user") + thread_summary = ThreadSummary( + platform=key.platform, + scope=key.scope, + thread=key.thread, + summary=summary_text, + message_count=message_count, + tags=tags, + ) + self._store.save(thread_summary) + + self._histories[key] = [ + {"role": "system", "content": f"Summary of earlier conversation:\n{summary_text}"} + ] + return thread_summary + + def get_summary(self, key: ThreadKey) -> ThreadSummary | None: + """Return the stored summary for a scope, or None if not found.""" + return self._store.get(key) + + def search_summaries(self, query: str) -> list[ThreadSummary]: + """Search the knowledge base for summaries whose tags overlap the query.""" + return self._store.search(query) + + def all_summaries(self) -> list[ThreadSummary]: + """Return all stored summaries, newest first.""" + return self._store.all() diff --git a/steward/bot/matrix.py b/steward/bot/matrix.py new file mode 100644 index 0000000..e04aeb3 --- /dev/null +++ b/steward/bot/matrix.py @@ -0,0 +1,126 @@ +"""Matrix appservice bot interface for Steward. + +Registers as a Synapse application service. Synapse pushes room events to the +appservice HTTP server (``/_matrix/app/v1/transactions/{txnId}``); the bot +replies through the client-server API using the appservice's ``as_token``. +""" + +from __future__ import annotations + +import asyncio +import logging + +from mautrix.appservice import AppService +from mautrix.types import ( + Event, + EventType, + MessageEvent, + MessageType, + TextMessageEventContent, +) + +from steward.bot.core import ConversationService +from steward.bot.thread_key import ThreadKey +from steward.config import Settings + +logger = logging.getLogger(__name__) + +_MATRIX_SYSTEM_APPENDIX = """ +Matrix conversation guidance: +- Reply like a thoughtful software engineer in a chat, not like a one-shot FAQ bot. +- Ask a brief follow-up question when the requested outcome, constraints, or preferred + option is unclear. +- Keep replies to a single message unless splitting genuinely helps readability. +""".strip() + + +class StewardMatrixBot: + """Matrix appservice bot that drives the shared conversation pipeline.""" + + def __init__(self, settings: Settings, service: ConversationService) -> None: + self._settings = settings + self._service = service + self._az = AppService( + server=settings.matrix.homeserver_url, + domain=settings.matrix.homeserver_domain, + as_token=settings.matrix.as_token, + hs_token=settings.matrix.hs_token, + bot_localpart=settings.matrix.bot_localpart, + id=settings.matrix.appservice_id, + ) + self._az.matrix_event_handler(self._on_event) + + @property + def bot_mxid(self) -> str: + return self._az.bot_mxid + + def _is_allowed_user(self, sender: str) -> bool: + allowed = self._settings.matrix.allowed_user_ids + if not allowed: + return True + return sender in allowed + + def _is_allowed_room(self, room_id: str) -> bool: + allowed = self._settings.matrix.allowed_room_ids + if not allowed: + return True + return room_id in allowed + + def _matrix_system_prompt(self) -> str: + base_prompt = self._settings.openai_system_prompt.strip() + if not base_prompt: + return _MATRIX_SYSTEM_APPENDIX + return f"{base_prompt}\n\n{_MATRIX_SYSTEM_APPENDIX}" + + async def _on_event(self, evt: Event) -> None: + if not isinstance(evt, MessageEvent): + return + if evt.sender == self.bot_mxid: + return + if evt.type != EventType.ROOM_MESSAGE: + return + if not isinstance(evt.content, TextMessageEventContent): + return + if not self._is_allowed_user(evt.sender): + logger.info("Ignoring message from unauthorized user %s", evt.sender) + return + if not self._is_allowed_room(evt.room_id): + logger.info("Ignoring message in unauthorized room %s", evt.room_id) + return + + body = (evt.content.body or "").strip() + if not body: + return + + key = ThreadKey(platform="matrix", scope=evt.room_id) + reply = await self._service.process_message( + key, + evt.sender, + body, + self._matrix_system_prompt(), + history_cap=40, + ) + if not reply.strip(): + return + + content = TextMessageEventContent(msgtype=MessageType.TEXT, body=reply) + await self._az.intent.send_message(evt.room_id, content) + + async def run(self) -> None: + """Start the appservice HTTP server and keep the event loop alive.""" + await self._az.start( + host=self._settings.matrix.listen_host, + port=self._settings.matrix.listen_port, + ) + logger.info( + "Matrix appservice listening on %s:%s (bot %s)", + self._settings.matrix.listen_host, + self._settings.matrix.listen_port, + self.bot_mxid, + ) + await self._az.intent.set_displayname("Steward") + while True: + await asyncio.sleep(3600) + + async def stop(self) -> None: + await self._az.stop() diff --git a/steward/bot/telegram.py b/steward/bot/telegram.py index 6d2e837..b3e05bf 100644 --- a/steward/bot/telegram.py +++ b/steward/bot/telegram.py @@ -2,7 +2,6 @@ import logging import re -from collections import defaultdict from dataclasses import dataclass, field from typing import Any @@ -16,23 +15,13 @@ from telegram.ext import ( filters, ) +from steward.bot.core import ConversationService +from steward.bot.thread_key import ThreadKey from steward.config import Settings -from steward.llm.client import LLMClient -from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary from steward.proposals.generator import Proposal, ProposalGenerator -from steward.tools.client import ToolClient logger = logging.getLogger(__name__) -# Per-user conversation history for non-threaded messages (capped at _MAX_HISTORY turns). -_history: dict[int, list[dict[str, Any]]] = defaultdict(list) -_MAX_HISTORY = 20 - -# Per-thread conversation history: (chat_id, thread_id) → full history (unbounded). -# Messages belonging to a Telegram message thread are kept in their entirety here -# until explicitly flushed by the /flush command. -_thread_history: dict[tuple[int, int], list[dict[str, Any]]] = defaultdict(list) - _FLUSH_SYSTEM_PROMPT = ( "You are Steward. The following is a complete Telegram message thread conversation. " "Produce a concise but comprehensive summary that captures:\n" @@ -45,18 +34,12 @@ _FLUSH_SYSTEM_PROMPT = ( _TAGS_SYSTEM_PROMPT = ( "You are a keyword tagger for a knowledge base. " - "Extract 5–8 short, lowercase keyword tags from the following conversation summary. " + "Extract 5\u20138 short, lowercase keyword tags from the following conversation summary. " "Tags should represent the main topics, entities, and concepts discussed. " "Return ONLY a comma-separated list of tags with no other text or punctuation. " "Example output: api design, authentication, database schema, user roles, caching" ) -_KB_CONTEXT_HEADER = ( - "The following are relevant past conversation summaries from your knowledge base. " - "Use them as background context if they relate to the current question, " - "but do not repeat their contents unless directly asked." -) - _CONVERSATION_SYSTEM_APPENDIX = """ Telegram conversation guidance: - Reply like a thoughtful software engineer in a chat, not like a one-shot FAQ bot. @@ -111,8 +94,8 @@ class TelegramResponsePlan: return [action for action in self.actions if isinstance(action, PollRequest)] -def _thread_key(update: Update) -> tuple[int, int] | None: - """Return the (chat_id, thread_id) key if the message is part of a thread, else None.""" +def _thread_key(update: Update) -> ThreadKey | None: + """Return the ThreadKey if the message is part of a Telegram thread, else None.""" msg = update.message chat = update.effective_chat if msg is None or chat is None: @@ -120,7 +103,7 @@ def _thread_key(update: Update) -> tuple[int, int] | None: thread_id = msg.message_thread_id if thread_id is None: return None - return (chat.id, thread_id) + return ThreadKey(platform="telegram", scope=str(chat.id), thread=str(thread_id)) def _is_allowed(user_id: int, settings: Settings) -> bool: @@ -307,11 +290,11 @@ async def start_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N "Hello, I'm *Steward* \U0001f916\n\n" "I'm your AI-assisted personal operations platform.\n" "Talk to me naturally, or use:\n" - "/help – show available commands\n" - "/clear – reset conversation history\n" - "/flush – summarise and archive this thread's memory\n" - "/recall [query] – retrieve archived thread summaries\n" - "/analyse – run a manual API analysis right now", + "/help \u2013 show available commands\n" + "/clear \u2013 reset conversation history\n" + "/flush \u2013 summarise and archive this thread's memory\n" + "/recall [query] \u2013 retrieve archived thread summaries\n" + "/analyse \u2013 run a manual API analysis right now", parse_mode=ParseMode.MARKDOWN, ) @@ -329,24 +312,26 @@ async def help_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> No await update.message.reply_text( # type: ignore[union-attr] "*Steward commands*\n\n" - "/start – greeting\n" - "/help – this message\n" - "/clear – reset conversation history for this context\n" - "/flush – summarise the current thread, store the summary, and compress memory\n" + "/start \u2013 greeting\n" + "/help \u2013 this message\n" + "/clear \u2013 reset conversation history for this context\n" + "/flush \u2013 summarise the current thread, store the summary, and compress memory\n" " _(only available inside a message thread)_\n" - "/recall [query] – show this thread's summary, list all summaries, or search by keyword\n" - "/analyse – trigger an immediate API analysis and proposal", + "/recall [query] \u2013 show this thread's summary, list all summaries, or search\n" + " by keyword\n" + "/analyse \u2013 trigger an immediate API analysis and proposal", parse_mode=ParseMode.MARKDOWN, ) async def clear_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /clear – wipe conversation history for this context. + """Handle /clear \u2013 wipe conversation history for this context. Inside a message thread: clears the thread's unbounded history. Outside a thread: clears the per-user capped history. """ settings: Settings = context.bot_data["settings"] + service: ConversationService = context.bot_data["service"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): return @@ -357,30 +342,21 @@ async def clear_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N key = _thread_key(update) if key is not None: - _thread_history[key].clear() + service.clear_history(key) await update.message.reply_text( # type: ignore[union-attr] "Thread conversation history cleared." ) else: - _history[user.id].clear() + service.clear_history(ThreadKey(platform="telegram", scope=str(user.id))) await update.message.reply_text( # type: ignore[union-attr] "Conversation history cleared." ) async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /flush – summarise thread memory, persist it, and compress in-memory history. - - Steps: - 1. Verify the command is issued inside a message thread. - 2. Summarise the full thread history via the LLM. - 3. Persist the summary in the ThreadMemoryStore (keyed by chat_id + thread_id). - 4. Replace the in-memory thread history with a single compressed context message - so conversation can continue with the summary as background. - """ + """Handle /flush \u2013 summarise thread memory, persist it, and compress in-memory history.""" settings: Settings = context.bot_data["settings"] - llm: LLMClient = context.bot_data["llm"] - store: ThreadMemoryStore = context.bot_data["thread_store"] + service: ConversationService = context.bot_data["service"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): @@ -397,10 +373,7 @@ async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N ) return - chat_id, thread_id = key - history = _thread_history[key] - - if not history: + if not service.has_history(key): await update.message.reply_text( # type: ignore[union-attr] "This thread has no conversation history to flush." ) @@ -410,60 +383,25 @@ async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N "\U0001f4be Summarising thread memory\u2026" ) - # Build a readable transcript for the LLM to summarise - transcript_lines = [] - for msg in history: - role_label = "User" if msg["role"] == "user" else "Steward" - transcript_lines.append(f"{role_label}: {msg['content']}") - transcript = "\n".join(transcript_lines) - - summary_text = await llm.chat( - f"Thread transcript:\n\n{transcript}", - system_prompt=_FLUSH_SYSTEM_PROMPT, - ) - - # Extract keyword tags for knowledge-base indexing (second LLM call, lightweight) - tags_raw = await llm.chat( - f"Summary to tag:\n\n{summary_text}", - system_prompt=_TAGS_SYSTEM_PROMPT, - ) - tags = [t.strip().lower() for t in tags_raw.split(",") if t.strip()][:8] - - message_count = sum(1 for m in history if m["role"] == "user") - thread_summary = ThreadSummary( - chat_id=chat_id, - thread_id=thread_id, - summary=summary_text, - message_count=message_count, - tags=tags, - ) - store.save(thread_summary) - - # Compress: replace history with a single system-context entry so the thread - # can continue with the summary as background knowledge. - _thread_history[key] = [ - {"role": "system", "content": f"Summary of earlier conversation:\n{summary_text}"} - ] + thread_summary = await service.flush_history(key) + if thread_summary is None: + await update.message.reply_text( # type: ignore[union-attr] + "This thread has no conversation history to flush." + ) + return await update.message.reply_text( # type: ignore[union-attr] f"\u2705 Thread memory flushed and stored " - f"(thread `{thread_id}`, {message_count} messages summarised).", + f"(thread `{thread_summary.thread_id}`, " + f"{thread_summary.message_count} messages summarised).", parse_mode=ParseMode.MARKDOWN, ) async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /recall [query] – retrieve stored thread summaries. - - With a query argument (e.g. ``/recall api design``): searches all stored - summaries whose tags overlap with the query keywords and returns matches. - - Without a query: - - Inside a thread: shows the stored summary for this thread (if any). - - Outside a thread: lists all stored summaries (newest first). - """ + """Handle /recall [query] \u2013 retrieve stored thread summaries.""" settings: Settings = context.bot_data["settings"] - store: ThreadMemoryStore = context.bot_data["thread_store"] + service: ConversationService = context.bot_data["service"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): @@ -473,11 +411,10 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> if chat is None or not _is_chat_enabled(chat, settings): return - # If the user supplied a keyword query, search the knowledge base - args: list[str] = context.args or [] # type: ignore[assignment] + args: list[str] = context.args or [] if args: query = " ".join(args).strip() - results = store.search(query) + results = service.search_summaries(query) if not results: await update.message.reply_text( # type: ignore[union-attr] f"No memories found matching *{query}*. " @@ -497,8 +434,7 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> key = _thread_key(update) if key is not None: - chat_id, thread_id = key - stored = store.get(chat_id, thread_id) + stored = service.get_summary(key) if stored is None: await update.message.reply_text( # type: ignore[union-attr] "No stored summary for this thread yet. Use /flush to create one." @@ -507,8 +443,7 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> await _send_long(update, stored.format_for_telegram()) return - # Outside a thread: list all stored summaries - all_summaries = store.all() + all_summaries = service.all_summaries() if not all_summaries: await update.message.reply_text( # type: ignore[union-attr] "No thread summaries stored yet. Use /flush inside a message thread." @@ -526,9 +461,8 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /analyse – run the proposal generator on demand.""" + """Handle /analyse \u2013 run the proposal generator on demand.""" settings: Settings = context.bot_data["settings"] - llm: LLMClient = context.bot_data["llm"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): return @@ -538,7 +472,7 @@ async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> return await update.message.reply_text("Running analysis, please wait...") # type: ignore[union-attr] - generator = ProposalGenerator(settings, llm) + generator = ProposalGenerator(settings, context.bot_data["llm"]) proposal = await generator.run() if proposal is None: await update.message.reply_text( # type: ignore[union-attr] @@ -549,42 +483,10 @@ async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> await _send_long(update, proposal.format_for_telegram()) -def _with_kb_context( - history: list[dict[str, Any]], - relevant: list[ThreadSummary], -) -> list[dict[str, Any]]: - """Prepend relevant knowledge-base summaries as a transient system context message. - - The returned list is a *new* list — the original *history* is not mutated. - The injected message is never appended to the stored history, so it does not - permanently consume the context window. - """ - if not relevant: - return history - snippets = [f"[Thread {s.thread_id}] {s.summary[:400]}" for s in relevant[:3]] - kb_msg = _KB_CONTEXT_HEADER + "\n\n" + "\n\n---\n\n".join(snippets) - return [{"role": "system", "content": kb_msg}, *history] - - async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle plain text messages – forward to LLM and reply. - - Thread messages: history is stored unbounded under the (chat_id, thread_id) key. - Non-thread messages: history is capped at _MAX_HISTORY turns per user. - - Before each LLM call the knowledge base is searched for summaries whose tags - overlap with keywords in the current message. Any matches are injected as - transient context — they are NOT stored in the rolling history, so they do - not permanently consume the context window. - - If a :class:`~steward.tools.client.ToolClient` is registered in - ``context.bot_data``, the LLM is invoked with tool calling support so it - can take actions on the configured MCP/OpenAPI tool server. - """ + """Handle plain text messages \u2013 forward to the shared pipeline and reply.""" settings: Settings = context.bot_data["settings"] - llm: LLMClient = context.bot_data["llm"] - store: ThreadMemoryStore = context.bot_data["thread_store"] - tool_client: ToolClient | None = context.bot_data.get("tool_client") + service: ConversationService = context.bot_data["service"] user = update.effective_user chat = update.effective_chat if user is None or not _is_allowed(user.id, settings): @@ -598,60 +500,35 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> system_prompt = _telegram_system_prompt(settings) key = _thread_key(update) - if key is not None: - # Thread message: unbounded history - history: list[dict[str, Any]] = _thread_history[key] - call_history = _with_kb_context(history, store.search(text)) - if tool_client is not None: - reply = await llm.chat_with_tools( - text, - tool_client, - history=call_history, - system_prompt=system_prompt, - ) - else: - reply = await llm.chat(text, history=call_history, system_prompt=system_prompt) - response_plan = _parse_telegram_response(reply) - history.append({"role": "user", "content": text}) - history.append( - {"role": "assistant", "content": _format_response_for_history(response_plan)} - ) + if key is None: + key = ThreadKey(platform="telegram", scope=str(user.id)) + history_cap = 40 else: - # Non-thread message: capped history per user - user_history: list[dict[str, Any]] = _history[user.id] - call_history = _with_kb_context(user_history, store.search(text)) - if tool_client is not None: - reply = await llm.chat_with_tools( - text, - tool_client, - history=call_history, - system_prompt=system_prompt, - ) - else: - reply = await llm.chat(text, history=call_history, system_prompt=system_prompt) - response_plan = _parse_telegram_response(reply) - user_history.append({"role": "user", "content": text}) - user_history.append( - {"role": "assistant", "content": _format_response_for_history(response_plan)} - ) - if len(user_history) > _MAX_HISTORY * 2: - _history[user.id] = user_history[-(_MAX_HISTORY * 2) :] + history_cap = None + reply = await service.process_message( + key, + str(user.id), + text, + system_prompt, + history_cap=history_cap, + history_formatter=lambda raw: _format_response_for_history(_parse_telegram_response(raw)), + ) + response_plan = _parse_telegram_response(reply) await _send_telegram_response(update, response_plan) def build_application( settings: Settings, - llm: LLMClient, - thread_store: ThreadMemoryStore | None = None, - tool_client: ToolClient | None = None, + service: ConversationService, + generator: ProposalGenerator | None = None, ) -> Application: # type: ignore[type-arg] """Build and return the Telegram Application.""" app = Application.builder().token(settings.telegram_bot_token).build() app.bot_data["settings"] = settings - app.bot_data["llm"] = llm - app.bot_data["thread_store"] = thread_store or ThreadMemoryStore(settings.thread_memory_path) - app.bot_data["tool_client"] = tool_client # None when tools are not configured + app.bot_data["service"] = service + app.bot_data["llm"] = service.llm + app.bot_data["generator"] = generator logger.info("Configured allowed users: %s", settings.telegram_allowed_user_ids) logger.info("Configured group IDs: %s", settings.telegram_group_ids) diff --git a/steward/bot/thread_key.py b/steward/bot/thread_key.py new file mode 100644 index 0000000..6db2bba --- /dev/null +++ b/steward/bot/thread_key.py @@ -0,0 +1,24 @@ +"""Normalised conversation identity shared across platforms and storage.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True, slots=True) +class ThreadKey: + """Normalised identity of a conversation scope across platforms. + + ``platform`` is ``"telegram"`` or ``"matrix"``. ``scope`` is the chat/room/user + identifier as a string. ``thread`` is an optional sub-thread identifier. + """ + + platform: str + scope: str + thread: str | None = None + + def __str__(self) -> str: + parts = [self.platform, self.scope] + if self.thread: + parts.append(self.thread) + return ":".join(parts) diff --git a/steward/config.py b/steward/config.py index 4a04dcc..46144fa 100644 --- a/steward/config.py +++ b/steward/config.py @@ -104,6 +104,34 @@ class ToolsConfig(BaseModel): mcp_server_api_key: str = Field(default="", description="MCP server API key") +class MatrixConfig(BaseModel): + """Matrix appservice configuration.""" + + homeserver_url: str = Field(default="", description="Homeserver client-server base URL") + homeserver_domain: str = Field(default="", description="Homeserver domain (server_name)") + as_token: str = Field( + default="", description="Appservice token for authenticating to the homeserver" + ) + hs_token: str = Field( + default="", description="Homeserver token for authenticating incoming transactions" + ) + bot_localpart: str = Field(default="steward", description="Localpart of the bot user") + appservice_id: str = Field(default="steward", description="Unique appservice ID") + listen_host: str = Field( + default="0.0.0.0", description="Host the appservice HTTP server listens on" + ) + listen_port: int = Field(default=8000, description="Port the appservice HTTP server listens on") + allowed_room_ids: list[str] = Field(default_factory=list, description="Allowed room IDs") + allowed_user_ids: list[str] = Field(default_factory=list, description="Allowed user MXIDs") + + @field_validator("allowed_room_ids", "allowed_user_ids", mode="before") + @classmethod + def parse_str_list(cls, value: Any) -> Any: + if isinstance(value, str): + return [item.strip() for item in value.split(",") if item.strip()] + return value + + class Settings(BaseModel): """Application settings with OmegaConf and pydantic integration.""" @@ -114,6 +142,7 @@ class Settings(BaseModel): analysis: AnalysisConfig = Field(default_factory=AnalysisConfig) memory: MemoryConfig = Field(default_factory=MemoryConfig) tools: ToolsConfig = Field(default_factory=ToolsConfig) + matrix: MatrixConfig = Field(default_factory=MatrixConfig) # Compatibility properties for existing code @property @@ -186,6 +215,11 @@ class Settings(BaseModel): """Legacy property for backward compatibility.""" return self.tools.mcp_server_api_key + @property + def matrix_enabled(self) -> bool: + """Return True if Matrix appservice integration is configured.""" + return bool(self.matrix.homeserver_url and self.matrix.as_token and self.matrix.hs_token) + _settings_instance: Settings | None = None diff --git a/steward/config_schema.yaml b/steward/config_schema.yaml index ac08813..a45e8d9 100644 --- a/steward/config_schema.yaml +++ b/steward/config_schema.yaml @@ -49,3 +49,36 @@ tools: # API key for MCP server (optional) mcp_server_api_key: "" + +matrix: + # Matrix appservice integration (optional - set to enable the Matrix bot) + # The bot registers as a Synapse application service and receives events + # via HTTP transactions, then replies through the client-server API. + + # Homeserver client-server base URL (e.g. http://matrix:8008) + homeserver_url: "" + + # Homeserver domain (server_name), e.g. matrix.aridgwayweb.com + homeserver_domain: "" + + # Appservice token used by the bot to authenticate to the homeserver + as_token: "" + + # Homeserver token used to authenticate incoming transactions + hs_token: "" + + # Localpart of the bot user (becomes @:) + bot_localpart: "steward" + + # Unique appservice ID + appservice_id: "steward" + + # Host/port the appservice HTTP server listens on + listen_host: "0.0.0.0" + listen_port: 8000 + + # List of room IDs the bot is allowed to operate in (empty = all rooms) + allowed_room_ids: [] + + # List of user IDs (MXIDs) allowed to talk to the bot (empty = all users) + allowed_user_ids: [] diff --git a/steward/main.py b/steward/main.py index 927b982..dffa582 100644 --- a/steward/main.py +++ b/steward/main.py @@ -1,11 +1,16 @@ """Steward application entry point.""" +import asyncio import logging +import signal import sys from datetime import time +from typing import Any -from telegram.ext import ContextTypes +from telegram.ext import Application, ContextTypes +from steward.bot.core import ConversationService +from steward.bot.matrix import StewardMatrixBot from steward.bot.telegram import build_application, send_proposal from steward.config import get_settings from steward.llm.client import LLMClient @@ -34,12 +39,37 @@ async def _run_scheduled_analysis(context: ContextTypes.DEFAULT_TYPE) -> None: logger.info("No proposal generated (generator returned None)") -def main() -> None: - """Start Steward.""" +async def _run_telegram(app: Application, shutdown_event: asyncio.Event) -> None: # type: ignore[type-arg] + """Start the Telegram bot and wait for shutdown.""" + await app.initialize() + await app.start() + if app.updater is not None: + await app.updater.start_polling(allowed_updates=["message"]) + logger.info("Telegram bot started") + try: + await shutdown_event.wait() + finally: + if app.updater is not None: + await app.updater.stop() + await app.stop() + await app.shutdown() + + +async def _run_matrix(bot: StewardMatrixBot, shutdown_event: asyncio.Event) -> None: + """Start the Matrix appservice and wait for shutdown.""" + await bot.run() + try: + await shutdown_event.wait() + finally: + await bot.stop() + + +async def main() -> None: + """Start Steward (Telegram and/or Matrix).""" settings = get_settings() - if not settings.telegram_bot_token: - logger.error("TELEGRAM_BOT_TOKEN is not set – cannot start") + if not settings.telegram_bot_token and not settings.matrix_enabled: + logger.error("No platform configured – set TELEGRAM_BOT_TOKEN or Matrix appservice config") sys.exit(1) if not settings.openai_api_key: @@ -59,27 +89,55 @@ def main() -> None: else: logger.info("No MCP_SERVER_URL configured – tool calling disabled") - app = build_application(settings, llm, thread_store, tool_client) - generator = ProposalGenerator(settings, llm) - app.bot_data["generator"] = generator - app.bot_data["user_ids"] = settings.telegram_allowed_user_ids + service = ConversationService(settings, llm, thread_store, tool_client) - if app.job_queue is not None: - app.job_queue.run_daily( - _run_scheduled_analysis, - time=time(hour=settings.analysis_cron_hour, minute=settings.analysis_cron_minute), + shutdown_event = asyncio.Event() + loop = asyncio.get_running_loop() + for sig in (signal.SIGINT, signal.SIGTERM): + loop.add_signal_handler(sig, shutdown_event.set) + + tasks: list[asyncio.Task[Any]] = [] + + if settings.telegram_bot_token: + app = build_application(settings, service) + generator = ProposalGenerator(settings, llm) + app.bot_data["generator"] = generator + app.bot_data["user_ids"] = settings.telegram_allowed_user_ids + + if app.job_queue is not None: + app.job_queue.run_daily( + _run_scheduled_analysis, + time=time(hour=settings.analysis_cron_hour, minute=settings.analysis_cron_minute), + ) + else: + logger.warning("JobQueue not available – scheduled analysis disabled") + + logger.info( + "Steward starting: model=%s analysis_url=%s tools=%s", + settings.openai_model, + settings.analysis_target_url or "(none)", + settings.mcp_server_url or "(none)", ) - else: - logger.warning("JobQueue not available – scheduled analysis disabled") + tasks.append(asyncio.create_task(_run_telegram(app, shutdown_event))) - logger.info( - "Steward starting: model=%s analysis_url=%s tools=%s", - settings.openai_model, - settings.analysis_target_url or "(none)", - settings.mcp_server_url or "(none)", - ) - app.run_polling(allowed_updates=["message"]) + if settings.matrix_enabled: + matrix_bot = StewardMatrixBot(settings, service) + tasks.append(asyncio.create_task(_run_matrix(matrix_bot, shutdown_event))) + + if not tasks: + logger.error("No platform started") + sys.exit(1) + + try: + await asyncio.gather(*tasks) + except asyncio.CancelledError: + pass if __name__ == "__main__": - main() + asyncio.run(main()) + + +def run() -> None: + """Synchronous entry point for the ``steward`` console script.""" + asyncio.run(main()) diff --git a/steward/memory/thread_store.py b/steward/memory/thread_store.py index 7b43ffe..035930d 100644 --- a/steward/memory/thread_store.py +++ b/steward/memory/thread_store.py @@ -8,20 +8,26 @@ 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 Telegram message thread. + """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). """ - chat_id: int - thread_id: int + platform: str + scope: str + thread: str | None summary: str message_count: int flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat()) @@ -29,18 +35,49 @@ class ThreadSummary: @property def key(self) -> str: - return f"{self.chat_id}:{self.thread_id}" + 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.""" + """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( - chat_id=int(data["chat_id"]), # type: ignore[arg-type] - thread_id=int(data["thread_id"]), # type: ignore[arg-type] + platform="telegram", + scope=str(data["chat_id"]), + thread=str(data["thread_id"]), summary=str(data["summary"]), - message_count=int(data["message_count"]), # type: ignore[arg-type] + message_count=int(str(data["message_count"])), flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())), - tags=list(data.get("tags", [])), # type: ignore[arg-type] + tags=cls._extract_tags(data), ) def format_for_telegram(self) -> str: @@ -72,13 +109,39 @@ class ThreadMemoryStore: try: raw = json.loads(self._path.read_text(encoding="utf-8")) if isinstance(raw, dict): - return raw # type: ignore[return-value] + 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( @@ -92,9 +155,9 @@ class ThreadMemoryStore: 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}") + 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) @@ -116,6 +179,6 @@ class ThreadMemoryStore: 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] + if query_words & {t.lower() for t in ThreadSummary._extract_tags(v)} ] return sorted(results, key=lambda s: s.flushed_at, reverse=True) diff --git a/tests/test_bot.py b/tests/test_bot.py index cac9efb..2abf1c9 100644 --- a/tests/test_bot.py +++ b/tests/test_bot.py @@ -6,15 +6,15 @@ import pytest from telegram import Chat, Message, Update, User from telegram.ext import CallbackContext +from steward.bot.core import ConversationService from steward.bot.telegram import ( - _history, _is_allowed, _is_chat_enabled, - _thread_history, clear_handler, message_handler, start_handler, ) +from steward.bot.thread_key import ThreadKey from steward.config import Settings from steward.llm.client import LLMClient from steward.memory.thread_store import ThreadMemoryStore @@ -68,10 +68,13 @@ def _make_context( mock_store = MagicMock(spec=ThreadMemoryStore) mock_store.search.return_value = [] store = mock_store + if llm is None: + llm = MagicMock(spec=LLMClient) + service = ConversationService(settings, llm, store) ctx.bot_data = { "settings": settings, + "service": service, "llm": llm, - "thread_store": store, } ctx.args = [] return ctx @@ -140,14 +143,15 @@ async def test_clear_handler_clears_user_history(): """clear_handler without a thread clears per-user history.""" settings = _make_settings() user_id = 42 - _history[user_id] = [{"role": "user", "content": "old msg"}] + ctx = _make_context(settings) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope=str(user_id)) + service._histories[key] = [{"role": "user", "content": "old msg"}] update = _make_update(user_id=user_id) # no thread_id - ctx = _make_context(settings) - await clear_handler(update, ctx) - assert _history[user_id] == [] + assert not service.has_history(key) update.message.reply_text.assert_awaited_once() @@ -155,15 +159,15 @@ async def test_clear_handler_clears_user_history(): async def test_clear_handler_clears_thread_history(): """clear_handler inside a thread clears that thread's history.""" settings = _make_settings() - key = (100, 7) - _thread_history[key] = [{"role": "user", "content": "thread msg"}] + ctx = _make_context(settings) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope="100", thread="7") + service._histories[key] = [{"role": "user", "content": "thread msg"}] update = _make_update(user_id=1, chat_id=100, thread_id=7) - ctx = _make_context(settings) - await clear_handler(update, ctx) - assert _thread_history[key] == [] + assert not service.has_history(key) update.message.reply_text.assert_awaited_once() @@ -175,15 +179,17 @@ async def test_message_handler_calls_llm_and_replies(): update = _make_update(user_id=77, text="What is the weather?") ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope="77") - _history[77].clear() await message_handler(update, ctx) mock_llm.chat.assert_awaited_once() update.message.reply_text.assert_awaited_once() - assert len(_history[77]) == 2 - assert _history[77][0]["role"] == "user" - assert _history[77][1]["role"] == "assistant" + history = service._histories[key] + assert len(history) == 2 + assert history[0]["role"] == "user" + assert history[1]["role"] == "assistant" @pytest.mark.asyncio @@ -196,15 +202,16 @@ async def test_message_handler_sends_multiple_reply_messages(): update = _make_update(user_id=78, text="Can you help with this vague thing?") ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope="78") - _history[78].clear() await message_handler(update, ctx) assert update.message.reply_text.await_count == 2 assert update.message.reply_text.await_args_list[0].args[0] == "First thought." assert update.message.reply_text.await_args_list[1].args[0] == "What outcome do you want?" update.message.reply_poll.assert_not_awaited() - assert _history[78][1]["content"] == "First thought.\n\nWhat outcome do you want?" + assert service._histories[key][1]["content"] == "First thought.\n\nWhat outcome do you want?" @pytest.mark.asyncio @@ -225,8 +232,9 @@ async def test_message_handler_sends_native_poll_from_directive(): update = _make_update(user_id=79, text="Should we do a minimal change or redesign?") ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope="79") - _history[79].clear() await message_handler(update, ctx) assert update.message.reply_text.await_count == 2 @@ -235,7 +243,7 @@ async def test_message_handler_sends_native_poll_from_directive(): options=["Minimal change", "Full redesign"], is_anonymous=False, ) - assert _history[79][1]["content"] == ( + assert service._histories[key][1]["content"] == ( "I can turn that into a vote.\n\n" "Poll: Which implementation should we choose? (Minimal change; Full redesign)\n\n" "I'll use the winning option." @@ -245,22 +253,24 @@ async def test_message_handler_sends_native_poll_from_directive(): @pytest.mark.asyncio async def test_message_handler_non_thread_trims_history(): """Non-thread history is trimmed to _MAX_HISTORY turns.""" - from steward.bot.telegram import _MAX_HISTORY + from steward.bot.core import _MAX_HISTORY settings = _make_settings() mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="reply") user_id = 200 + ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope=str(user_id)) # Pre-fill exactly at the limit - _history[user_id] = [ + service._histories[key] = [ {"role": "user" if i % 2 == 0 else "assistant", "content": f"msg{i}"} for i in range(_MAX_HISTORY * 2) ] update = _make_update(user_id=user_id, text="new question") - ctx = _make_context(settings, llm=mock_llm) await message_handler(update, ctx) # After appending 2 new messages, should trim back to _MAX_HISTORY * 2 - assert len(_history[user_id]) == _MAX_HISTORY * 2 + assert len(service._histories[key]) == _MAX_HISTORY * 2 diff --git a/tests/test_thread_memory.py b/tests/test_thread_memory.py index 444b232..e749b24 100644 --- a/tests/test_thread_memory.py +++ b/tests/test_thread_memory.py @@ -7,12 +7,13 @@ import pytest from telegram import Chat, Message, Update, User from telegram.ext import CallbackContext +from steward.bot.core import ConversationService from steward.bot.telegram import ( - _thread_history, flush_handler, message_handler, recall_handler, ) +from steward.bot.thread_key import ThreadKey from steward.config import Settings from steward.llm.client import LLMClient from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary @@ -64,10 +65,13 @@ def _make_context( mock_store = MagicMock(spec=ThreadMemoryStore) mock_store.search.return_value = [] store = mock_store + if llm is None: + llm = MagicMock(spec=LLMClient) + service = ConversationService(settings, llm, store) ctx.bot_data = { "settings": settings, - "llm": llm or MagicMock(spec=LLMClient), - "thread_store": store, + "service": service, + "llm": llm, } ctx.args = [] return ctx @@ -82,50 +86,56 @@ class TestThreadMemoryStore: def _store(self, tmp_path: Path) -> ThreadMemoryStore: return ThreadMemoryStore(tmp_path / "mem.json") + def _key(self, scope: str = "1", thread: str | None = "2") -> ThreadKey: + return ThreadKey(platform="telegram", scope=scope, thread=thread) + + def _summary(self, scope: str = "1", thread: str | None = "2", **kwargs) -> ThreadSummary: + defaults = dict(summary="A recap.", message_count=5) + defaults.update(kwargs) + return ThreadSummary(platform="telegram", scope=scope, thread=thread, **defaults) + def test_get_missing_returns_none(self, tmp_path: Path): store = self._store(tmp_path) - assert store.get(1, 2) is None + assert store.get(self._key()) is None def test_save_and_get_roundtrip(self, tmp_path: Path): store = self._store(tmp_path) - summary = ThreadSummary(chat_id=1, thread_id=42, summary="A recap.", message_count=5) + summary = self._summary() store.save(summary) - retrieved = store.get(1, 42) + retrieved = store.get(self._key()) assert retrieved is not None assert retrieved.summary == "A recap." assert retrieved.message_count == 5 def test_save_overwrites_previous_entry(self, tmp_path: Path): store = self._store(tmp_path) - store.save(ThreadSummary(chat_id=1, thread_id=1, summary="old", message_count=2)) - store.save(ThreadSummary(chat_id=1, thread_id=1, summary="new", message_count=4)) - assert store.get(1, 1).summary == "new" # type: ignore[union-attr] + store.save(self._summary(summary="old", message_count=2)) + store.save(self._summary(summary="new", message_count=4)) + assert store.get(self._key()).summary == "new" # type: ignore[union-attr] def test_persisted_to_disk(self, tmp_path: Path): path = tmp_path / "mem.json" store = ThreadMemoryStore(path) - store.save(ThreadSummary(chat_id=5, thread_id=9, summary="saved", message_count=1)) + store.save(self._summary(scope="5", thread="9", summary="saved", message_count=1)) # Load a fresh store from the same file store2 = ThreadMemoryStore(path) - assert store2.get(5, 9) is not None + assert store2.get(self._key("5", "9")) is not None def test_all_returns_newest_first(self, tmp_path: Path): store = self._store(tmp_path) store.save( - ThreadSummary( - chat_id=1, - thread_id=1, + self._summary( + thread="1", summary="first", message_count=1, flushed_at="2026-01-01T00:00:00+00:00", ) ) store.save( - ThreadSummary( - chat_id=1, - thread_id=2, + self._summary( + thread="2", summary="second", message_count=1, flushed_at="2026-06-01T00:00:00+00:00", @@ -142,19 +152,13 @@ class TestThreadMemoryStore: assert store.all() == [] def test_format_for_telegram_contains_thread_id(self, tmp_path: Path): - s = ThreadSummary(chat_id=1, thread_id=77, summary="recap", message_count=3) + s = self._summary(thread="77", summary="recap", message_count=3) text = s.format_for_telegram() assert "77" in text assert "recap" in text def test_format_for_telegram_shows_tags(self, tmp_path: Path): - s = ThreadSummary( - chat_id=1, - thread_id=77, - summary="recap", - message_count=3, - tags=["api", "auth"], - ) + s = self._summary(thread="77", summary="recap", message_count=3, tags=["api", "auth"]) text = s.format_for_telegram() assert "api" in text assert "auth" in text @@ -174,55 +178,37 @@ class TestThreadMemoryStore: def test_search_returns_matching_summaries(self, tmp_path: Path): store = self._store(tmp_path) store.save( - ThreadSummary( - chat_id=1, - thread_id=1, - summary="API work", - message_count=2, - tags=["api", "design", "auth"], + self._summary( + thread="1", summary="API work", message_count=2, tags=["api", "design", "auth"] ) ) store.save( - ThreadSummary( - chat_id=1, - thread_id=2, - summary="Database work", - message_count=2, - tags=["database", "schema"], + self._summary( + thread="2", summary="Database work", message_count=2, tags=["database", "schema"] ) ) results = store.search("api") assert len(results) == 1 - assert results[0].thread_id == 1 + assert results[0].thread_id == "1" def test_search_no_match_returns_empty(self, tmp_path: Path): store = self._store(tmp_path) - store.save( - ThreadSummary( - chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["database"] - ) - ) + store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["database"])) assert store.search("deployment") == [] def test_search_empty_query_returns_empty(self, tmp_path: Path): store = self._store(tmp_path) - store.save( - ThreadSummary( - chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["database"] - ) - ) + store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["database"])) assert store.search("") == [] def test_search_case_insensitive(self, tmp_path: Path): store = self._store(tmp_path) - store.save( - ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["API"]) - ) + store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["API"])) assert len(store.search("api")) == 1 def test_search_skips_untagged_summaries(self, tmp_path: Path): store = self._store(tmp_path) - store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1, tags=[])) + store.save(self._summary(thread="1", summary="recap", message_count=1, tags=[])) assert store.search("api") == [] def test_legacy_store_roundtrip(self, tmp_path: Path): @@ -245,7 +231,7 @@ class TestThreadMemoryStore: encoding="utf-8", ) store = ThreadMemoryStore(path) - s = store.get(1, 1) + s = store.get(self._key("1", "1")) assert s is not None assert s.tags == [] @@ -257,49 +243,49 @@ class TestThreadMemoryStore: @pytest.mark.asyncio async def test_thread_message_stored_in_thread_history(): - """Messages in a thread go to _thread_history, not _history.""" - from steward.bot.telegram import _history - + """Messages in a thread go to the thread's history, not the user's.""" settings = _make_settings() mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="thread reply") - key = (100, 55) - _thread_history[key].clear() + ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + thread_key = ThreadKey(platform="telegram", scope="100", thread="55") + user_key = ThreadKey(platform="telegram", scope="1") update = _make_update(user_id=1, chat_id=100, thread_id=55, text="thread message") - ctx = _make_context(settings, llm=mock_llm) await message_handler(update, ctx) - assert len(_thread_history[key]) == 2 - assert _thread_history[key][0]["role"] == "user" + assert len(service._histories[thread_key]) == 2 + assert service._histories[thread_key][0]["role"] == "user" # Regular user history untouched - assert len(_history[1]) == 0 + assert not service.has_history(user_key) @pytest.mark.asyncio async def test_thread_history_is_unbounded(): """Thread history never gets trimmed regardless of how many turns there are.""" - from steward.bot.telegram import _MAX_HISTORY + from steward.bot.core import _MAX_HISTORY settings = _make_settings() mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="reply") - key = (200, 66) + ctx = _make_context(settings, llm=mock_llm) + service: ConversationService = ctx.bot_data["service"] + thread_key = ThreadKey(platform="telegram", scope="200", thread="66") # Pre-fill well beyond the cap used for non-thread history - _thread_history[key] = [ + service._histories[thread_key] = [ {"role": "user" if i % 2 == 0 else "assistant", "content": f"m{i}"} for i in range(_MAX_HISTORY * 4) # 4× the normal cap ] - prior_len = len(_thread_history[key]) + prior_len = len(service._histories[thread_key]) update = _make_update(user_id=2, chat_id=200, thread_id=66, text="more") - ctx = _make_context(settings, llm=mock_llm) await message_handler(update, ctx) # Should have grown by exactly 2 (user + assistant), never trimmed - assert len(_thread_history[key]) == prior_len + 2 + assert len(service._histories[thread_key]) == prior_len + 2 # --------------------------------------------------------------------------- @@ -333,9 +319,6 @@ async def test_flush_empty_thread_warns(): mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="summary") - key = (300, 88) - _thread_history[key].clear() - update = _make_update(chat_id=300, thread_id=88) ctx = _make_context(settings, llm=mock_llm, store=store) await flush_handler(update, ctx) @@ -353,8 +336,10 @@ async def test_flush_summarises_stores_and_compresses(): # First call → summary, second call → tags mock_llm.chat = AsyncMock(side_effect=["Great summary of the thread.", "api, design, testing"]) - key = (400, 99) - _thread_history[key] = [ + ctx = _make_context(settings, llm=mock_llm, store=store) + service: ConversationService = ctx.bot_data["service"] + key = ThreadKey(platform="telegram", scope="400", thread="99") + service._histories[key] = [ {"role": "user", "content": "question one"}, {"role": "assistant", "content": "answer one"}, {"role": "user", "content": "question two"}, @@ -362,7 +347,6 @@ async def test_flush_summarises_stores_and_compresses(): ] update = _make_update(chat_id=400, thread_id=99) - ctx = _make_context(settings, llm=mock_llm, store=store) await flush_handler(update, ctx) # LLM called twice: once for summary, once for tags @@ -370,16 +354,17 @@ async def test_flush_summarises_stores_and_compresses(): # Summary persisted store.save.assert_called_once() saved: ThreadSummary = store.save.call_args.args[0] - assert saved.chat_id == 400 - assert saved.thread_id == 99 + assert saved.platform == "telegram" + assert saved.scope == "400" + assert saved.thread == "99" assert saved.summary == "Great summary of the thread." assert saved.message_count == 2 # 2 user turns assert saved.tags == ["api", "design", "testing"] # In-memory history replaced with compressed context - assert len(_thread_history[key]) == 1 - assert _thread_history[key][0]["role"] == "system" - assert "Great summary" in _thread_history[key][0]["content"] + assert len(service._histories[key]) == 1 + assert service._histories[key][0]["role"] == "system" + assert "Great summary" in service._histories[key][0]["content"] # --------------------------------------------------------------------------- @@ -392,7 +377,11 @@ async def test_recall_in_thread_returns_stored_summary(tmp_path): """recall_handler inside a thread returns the stored summary for that thread.""" settings = _make_settings() store = ThreadMemoryStore(tmp_path / "mem.json") - store.save(ThreadSummary(chat_id=500, thread_id=11, summary="recap text", message_count=3)) + store.save( + ThreadSummary( + platform="telegram", scope="500", thread="11", summary="recap text", message_count=3 + ) + ) update = _make_update(chat_id=500, thread_id=11) ctx = _make_context(settings, store=store) @@ -423,8 +412,12 @@ async def test_recall_outside_thread_lists_all(tmp_path): """recall_handler outside a thread lists all stored summaries.""" settings = _make_settings() store = ThreadMemoryStore(tmp_path / "mem.json") - store.save(ThreadSummary(chat_id=1, thread_id=1, summary="alpha", message_count=1)) - store.save(ThreadSummary(chat_id=1, thread_id=2, summary="beta", message_count=2)) + store.save( + ThreadSummary(platform="telegram", scope="1", thread="1", summary="alpha", message_count=1) + ) + store.save( + ThreadSummary(platform="telegram", scope="1", thread="2", summary="beta", message_count=2) + ) update = _make_update(thread_id=None) # no thread ctx = _make_context(settings, store=store) @@ -457,8 +450,9 @@ async def test_recall_with_query_returns_matching_summaries(tmp_path): store = ThreadMemoryStore(tmp_path / "mem.json") store.save( ThreadSummary( - chat_id=1, - thread_id=1, + platform="telegram", + scope="1", + thread="1", summary="API authentication discussion", message_count=2, tags=["api", "auth"], @@ -466,8 +460,9 @@ async def test_recall_with_query_returns_matching_summaries(tmp_path): ) store.save( ThreadSummary( - chat_id=1, - thread_id=2, + platform="telegram", + scope="1", + thread="2", summary="Database schema planning", message_count=3, tags=["database", "schema"], @@ -492,8 +487,9 @@ async def test_recall_with_query_no_match(tmp_path): store = ThreadMemoryStore(tmp_path / "mem.json") store.save( ThreadSummary( - chat_id=1, - thread_id=1, + platform="telegram", + scope="1", + thread="1", summary="recap", message_count=1, tags=["database"], @@ -520,8 +516,9 @@ async def test_message_handler_injects_kb_context_transiently(tmp_path): store = ThreadMemoryStore(tmp_path / "mem.json") store.save( ThreadSummary( - chat_id=1, - thread_id=1, + platform="telegram", + scope="1", + thread="1", summary="Previous API discussion", message_count=2, tags=["api"], @@ -529,12 +526,11 @@ async def test_message_handler_injects_kb_context_transiently(tmp_path): ) user_id = 999 - from steward.bot.telegram import _history - - _history[user_id].clear() + ctx = _make_context(settings, llm=mock_llm, store=store) + service: ConversationService = ctx.bot_data["service"] + user_key = ThreadKey(platform="telegram", scope=str(user_id)) update = _make_update(user_id=user_id, text="tell me about the api work") - ctx = _make_context(settings, llm=mock_llm, store=store) await message_handler(update, ctx) # LLM was called @@ -547,8 +543,10 @@ async def test_message_handler_injects_kb_context_transiently(tmp_path): assert history_arg is not None assert any("knowledge base" in m.get("content", "").lower() for m in history_arg) - # But the KB message must NOT be stored in _history - assert all("knowledge base" not in m.get("content", "").lower() for m in _history[user_id]) + # But the KB message must NOT be stored in history + assert all( + "knowledge base" not in m.get("content", "").lower() for m in service._histories[user_key] + ) @pytest.mark.asyncio @@ -562,8 +560,9 @@ async def test_message_handler_no_kb_injection_when_no_match(tmp_path): # Store a summary with unrelated tags store.save( ThreadSummary( - chat_id=1, - thread_id=1, + platform="telegram", + scope="1", + thread="1", summary="Database recap", message_count=1, tags=["database"], @@ -571,12 +570,9 @@ async def test_message_handler_no_kb_injection_when_no_match(tmp_path): ) user_id = 888 - from steward.bot.telegram import _history - - _history[user_id].clear() + ctx = _make_context(settings, llm=mock_llm, store=store) update = _make_update(user_id=user_id, text="what is the weather like?") - ctx = _make_context(settings, llm=mock_llm, store=store) await message_handler(update, ctx) call_kwargs = mock_llm.chat.call_args