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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
parent
260720dd10
commit
76777c98eb
@ -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 = ["."]
|
||||
|
||||
179
steward/bot/core.py
Normal file
179
steward/bot/core.py
Normal file
@ -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()
|
||||
126
steward/bot/matrix.py
Normal file
126
steward/bot/matrix.py
Normal file
@ -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()
|
||||
@ -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)
|
||||
|
||||
|
||||
24
steward/bot/thread_key.py
Normal file
24
steward/bot/thread_key.py
Normal file
@ -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)
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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 @<localpart>:<domain>)
|
||||
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: []
|
||||
|
||||
104
steward/main.py
104
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())
|
||||
|
||||
@ -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
|
||||
``"<chat_id>:<thread_id>"`` with ``chat_id``/``thread_id`` fields. These
|
||||
are migrated to ``"telegram:<chat_id>:<thread_id>"`` 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)
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user