diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..a3820c2 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,37 @@ +# Python +__pycache__/ +*.py[cod] +*.egg-info/ +.eggs/ +dist/ +build/ + +# Virtual environments +.venv/ +venv/ +env/ + +# Environment secrets — never bake these into the image +.env + +# Test / lint caches +.pytest_cache/ +.coverage +htmlcov/ +.mypy_cache/ +.ruff_cache/ + +# IDE +.vscode/ +.idea/ + +# OS +.DS_Store + +# Git +.git/ +.github/ + +# Data directory (runtime artefact, not part of the image) +data/ +thread_memory.json diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md new file mode 100644 index 0000000..571afbc --- /dev/null +++ b/.github/pull_request_template.md @@ -0,0 +1,21 @@ +## What does this PR do? + + + +## Type of change + +- [ ] `fix:` Bug fix (patch release) +- [ ] `feat:` New feature (minor release) +- [ ] `feat!:` / `BREAKING CHANGE:` Breaking change (major release) +- [ ] `docs:` Documentation only +- [ ] `chore:` / `refactor:` / `test:` No release + +## Notes + + + +--- + +> **Commit message format** – this project uses [Conventional Commits](https://www.conventionalcommits.org/) +> to drive automatic semantic versioning via `release-please`. +> Use `fix:`, `feat:`, or `feat!:` as your commit prefix so the version bump is picked up correctly. diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..53a5017 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,62 @@ +name: CI + +on: + push: + branches-ignore: + - main + pull_request: + +permissions: + contents: read + packages: write + +jobs: + lint-and-test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-python@v5 + with: + python-version: "3.12" + cache: pip + + - run: pip install -e ".[dev]" + - run: python -m ruff check steward/ tests/ + - run: python -m pytest tests/ -v + + docker-build: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Log in to GHCR + # Only log in when pushing (not on PRs from forks) + if: github.event_name == 'push' + uses: docker/login-action@v3 + with: + registry: ghcr.io + username: ${{ github.actor }} + password: ${{ secrets.GITHUB_TOKEN }} + + - name: Docker metadata + id: meta + uses: docker/metadata-action@v5 + with: + images: ghcr.io/${{ github.repository }} + tags: | + # feature-branch push → ghcr.io/…/steward:my-feature-branch + type=ref,event=branch + # pull-request → ghcr.io/…/steward:pr-42 + type=ref,event=pr + # always → ghcr.io/…/steward:sha-abc1234 + type=sha,prefix=sha- + + - name: Build and push + uses: docker/build-push-action@v6 + with: + context: . + # Push on branch pushes; only validate (no push) on PRs + push: ${{ github.event_name == 'push' }} + tags: ${{ steps.meta.outputs.tags }} + labels: ${{ steps.meta.outputs.labels }} diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..601cc7c --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,62 @@ +name: Release + +# Runs on every push to main. +# release-please analyses conventional commits since the last release and either: +# • creates/updates a "Release PR" that bumps pyproject.toml + CHANGELOG.md, or +# • (when that PR is merged) tags the repo and creates a GitHub Release. +# Only when a GitHub Release is actually created does the publish job run. + +on: + push: + branches: + - main + +permissions: + contents: write + pull-requests: write + packages: write + +jobs: + release-please: + runs-on: ubuntu-latest + outputs: + release_created: ${{ steps.release.outputs.release_created }} + tag_name: ${{ steps.release.outputs.tag_name }} + steps: + - uses: googleapis/release-please-action@v4 + id: release + with: + release-type: python + + publish: + needs: release-please + if: needs.release-please.outputs.release_created == 'true' + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Log in to GHCR + uses: docker/login-action@v3 + with: + registry: ghcr.io + username: ${{ github.actor }} + password: ${{ secrets.GITHUB_TOKEN }} + + - name: Docker metadata + id: meta + uses: docker/metadata-action@v5 + with: + images: ghcr.io/${{ github.repository }} + tags: | + # semver tag → ghcr.io/…/steward:1.2.3 + type=raw,value=${{ needs.release-please.outputs.tag_name }} + # always push latest on a real release + type=raw,value=latest + + - name: Build and push + uses: docker/build-push-action@v6 + with: + context: . + push: true + tags: ${{ steps.meta.outputs.tags }} + labels: ${{ steps.meta.outputs.labels }} diff --git a/.release-please-manifest.json b/.release-please-manifest.json new file mode 100644 index 0000000..466df71 --- /dev/null +++ b/.release-please-manifest.json @@ -0,0 +1,3 @@ +{ + ".": "0.1.0" +} diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..3e41714 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,23 @@ +FROM python:3.12-slim + +WORKDIR /app + +# Install dependencies (separate layer for cache efficiency) +COPY pyproject.toml README.md ./ +COPY steward/ steward/ +RUN pip install --no-cache-dir . + +# Create a non-root user (UID/GID 1000) and persistent data directory +RUN addgroup --gid 1000 steward \ + && adduser --uid 1000 --gid 1000 --no-create-home --disabled-password --gecos "" steward \ + && mkdir /data \ + && chown steward:steward /data + +USER steward + +# Default path for the thread memory store — override via env or docker-compose +ENV THREAD_MEMORY_PATH=/data/thread_memory.json + +VOLUME ["/data"] + +CMD ["steward"] diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..4cae27d --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,14 @@ +services: + steward: + image: ghcr.io/djw4/steward:latest + # To build locally instead: uncomment the next line and comment out image above + # build: . + user: "1000:1000" # matches the UID/GID created in the Dockerfile + env_file: .env # copy .env.example → .env and fill in your values + environment: + THREAD_MEMORY_PATH: /data/thread_memory.json + volumes: + # Bind-mount a local ./data directory for persistent storage. + # Create it before the first run: mkdir -p data + - ./data:/data + restart: unless-stopped diff --git a/steward/bot/telegram.py b/steward/bot/telegram.py index 034d94d..fe73549 100644 --- a/steward/bot/telegram.py +++ b/steward/bot/telegram.py @@ -39,6 +39,20 @@ _FLUSH_SYSTEM_PROMPT = ( "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." +) + 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.""" @@ -83,7 +97,7 @@ async def start_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N "/help – show available commands\n" "/clear – reset conversation history\n" "/flush – summarise and archive this thread's memory\n" - "/recall – retrieve archived thread summaries\n" + "/recall [query] – retrieve archived thread summaries\n" "/analyse – run a manual API analysis right now", parse_mode=ParseMode.MARKDOWN, ) @@ -103,7 +117,7 @@ async def help_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> No "/clear – reset conversation history for this context\n" "/flush – summarise the current thread, store the summary, and compress memory\n" " _(only available inside a message thread)_\n" - "/recall – show the stored summary for this thread, or list all summaries\n" + "/recall [query] – show this thread's summary, list all summaries, or search by keyword\n" "/analyse – trigger an immediate API analysis and proposal", parse_mode=ParseMode.MARKDOWN, ) @@ -183,12 +197,20 @@ async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N 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()][:10] + 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) @@ -206,10 +228,14 @@ async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /recall – retrieve stored thread summaries. + """Handle /recall [query] – retrieve stored thread summaries. - Inside a thread: shows the stored summary for this thread (if any). - Outside a thread: lists all stored summaries (newest first). + 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). """ settings: Settings = context.bot_data["settings"] store: ThreadMemoryStore = context.bot_data["thread_store"] @@ -218,6 +244,28 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> if user is None or not _is_allowed(user.id, settings): return + # If the user supplied a keyword query, search the knowledge base + args: list[str] = context.args or [] # type: ignore[assignment] + if args: + query = " ".join(args).strip() + results = store.search(query) + if not results: + await update.message.reply_text( # type: ignore[union-attr] + f"No memories found matching *{query}*. " + "Try a different keyword or use /flush to add more summaries.", + parse_mode=ParseMode.MARKDOWN, + ) + return + lines = [f"\U0001f50d *Knowledge base search: {query}*\n"] + for s in results: + date = s.flushed_at[:10] + first_line = s.summary.split("\n")[0][:80] + lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026") + if s.tags: + lines.append(f" \U0001f3f7 {', '.join(s.tags)}") + await _send_long(update, "\n".join(lines)) + return + key = _thread_key(update) if key is not None: chat_id, thread_id = key @@ -243,6 +291,8 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> date = s.flushed_at[:10] first_line = s.summary.split("\n")[0][:80] lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026") + if s.tags: + lines.append(f" \U0001f3f7 {', '.join(s.tags)}") await _send_long(update, "\n".join(lines)) @@ -267,14 +317,37 @@ 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, str]], + relevant: list[ThreadSummary], +) -> list[dict[str, str]]: + """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. """ settings: Settings = context.bot_data["settings"] llm: LLMClient = context.bot_data["llm"] + store: ThreadMemoryStore = context.bot_data["thread_store"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): return @@ -287,13 +360,15 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> if key is not None: # Thread message: unbounded history history = _thread_history[key] - reply = await llm.chat(text, history=history) + call_history = _with_kb_context(history, store.search(text)) + reply = await llm.chat(text, history=call_history) history.append({"role": "user", "content": text}) history.append({"role": "assistant", "content": reply}) else: # Non-thread message: capped history per user history = _history[user.id] - reply = await llm.chat(text, history=history) + call_history = _with_kb_context(history, store.search(text)) + reply = await llm.chat(text, history=call_history) history.append({"role": "user", "content": text}) history.append({"role": "assistant", "content": reply}) if len(history) > _MAX_HISTORY * 2: diff --git a/steward/memory/thread_store.py b/steward/memory/thread_store.py index 2bf2a6e..7b43ffe 100644 --- a/steward/memory/thread_store.py +++ b/steward/memory/thread_store.py @@ -1,4 +1,4 @@ -"""Persistent storage for flushed thread summaries.""" +"""Persistent storage for flushed thread summaries (knowledge-base style).""" from __future__ import annotations @@ -13,18 +13,36 @@ logger = logging.getLogger(__name__) @dataclass class ThreadSummary: - """A persisted summary of a flushed Telegram message thread.""" + """A persisted summary of a flushed Telegram message thread. + + ``tags`` is a list of short lowercase keywords extracted by the LLM at flush + time. They are used to index the knowledge base so summaries can be recalled + contextually without being kept permanently in the conversation context. + """ chat_id: int thread_id: int summary: str message_count: int flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat()) + tags: list[str] = field(default_factory=list) @property def key(self) -> str: return f"{self.chat_id}:{self.thread_id}" + @classmethod + def from_dict(cls, data: dict[str, object]) -> ThreadSummary: + """Deserialise from a raw dict, tolerating missing optional fields.""" + return cls( + chat_id=int(data["chat_id"]), # type: ignore[arg-type] + thread_id=int(data["thread_id"]), # type: ignore[arg-type] + summary=str(data["summary"]), + message_count=int(data["message_count"]), # type: ignore[arg-type] + flushed_at=str(data.get("flushed_at", datetime.now(UTC).isoformat())), + tags=list(data.get("tags", [])), # type: ignore[arg-type] + ) + def format_for_telegram(self) -> str: """Return a concise Telegram-formatted recall card.""" flushed = self.flushed_at[:19].replace("T", " ") @@ -32,21 +50,24 @@ class ThreadSummary: f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n" f"_Flushed: {flushed} UTC \u2014 {self.message_count} messages_" ) - return f"{header}\n\n{self.summary}" + tag_line = f"\U0001f3f7 Tags: {', '.join(self.tags)}" if self.tags else "" + parts = [header, tag_line, self.summary] if tag_line else [header, self.summary] + return "\n\n".join(parts) class ThreadMemoryStore: - """JSON-backed store for flushed thread summaries. + """JSON-backed knowledge-base store for flushed thread summaries. - The store is intentionally lightweight for the MVP. Each flush overwrites - any previous summary for the same (chat_id, thread_id) pair. + Summaries are indexed by keyword tags so they can be recalled contextually + (via :meth:`search`) without being kept permanently in the LLM context. + Each flush overwrites the previous summary for the same (chat_id, thread_id) pair. """ def __init__(self, path: str | Path = "thread_memory.json") -> None: self._path = Path(path) - self._data: dict[str, dict[str, int | str]] = self._load() + self._data: dict[str, dict[str, object]] = self._load() - def _load(self) -> dict[str, dict[str, int | str]]: + def _load(self) -> dict[str, dict[str, object]]: if self._path.exists(): try: raw = json.loads(self._path.read_text(encoding="utf-8")) @@ -76,9 +97,25 @@ class ThreadMemoryStore: raw = self._data.get(f"{chat_id}:{thread_id}") if raw is None: return None - return ThreadSummary(**raw) # type: ignore[arg-type] + return ThreadSummary.from_dict(raw) def all(self) -> list[ThreadSummary]: """Return all stored summaries, newest first.""" - entries = [ThreadSummary(**v) for v in self._data.values()] # type: ignore[arg-type] + entries = [ThreadSummary.from_dict(v) for v in self._data.values()] return sorted(entries, key=lambda s: s.flushed_at, reverse=True) + + def search(self, query: str) -> list[ThreadSummary]: + """Return summaries whose tags overlap with words in *query*, newest first. + + The match is case-insensitive and word-based. Summaries without tags + are not returned even if the query is broad. + """ + query_words = {w.lower() for w in query.split() if w} + if not query_words: + return [] + results = [ + ThreadSummary.from_dict(v) + for v in self._data.values() + if query_words & {t.lower() for t in v.get("tags", [])} # type: ignore[union-attr] + ] + return sorted(results, key=lambda s: s.flushed_at, reverse=True) diff --git a/tests/test_bot.py b/tests/test_bot.py index 9ecda8e..6d830f7 100644 --- a/tests/test_bot.py +++ b/tests/test_bot.py @@ -58,11 +58,16 @@ def _make_context( store: ThreadMemoryStore | None = None, ) -> CallbackContext: # type: ignore[type-arg] ctx = MagicMock(spec=CallbackContext) + if store is None: + mock_store = MagicMock(spec=ThreadMemoryStore) + mock_store.search.return_value = [] + store = mock_store ctx.bot_data = { "settings": settings, "llm": llm, - "thread_store": store or MagicMock(spec=ThreadMemoryStore), + "thread_store": store, } + ctx.args = [] return ctx diff --git a/tests/test_thread_memory.py b/tests/test_thread_memory.py index 6df8de0..8c08682 100644 --- a/tests/test_thread_memory.py +++ b/tests/test_thread_memory.py @@ -55,11 +55,16 @@ def _make_context( store: ThreadMemoryStore | None = None, ) -> CallbackContext: # type: ignore[type-arg] ctx = MagicMock(spec=CallbackContext) + if store is None: + mock_store = MagicMock(spec=ThreadMemoryStore) + mock_store.search.return_value = [] + store = mock_store ctx.bot_data = { "settings": settings, "llm": llm or MagicMock(spec=LLMClient), - "thread_store": store or MagicMock(spec=ThreadMemoryStore), + "thread_store": store, } + ctx.args = [] return ctx @@ -122,6 +127,72 @@ class TestThreadMemoryStore: 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"], + ) + text = s.format_for_telegram() + assert "api" in text + assert "auth" in text + + def test_from_dict_tolerates_missing_tags(self, tmp_path: Path): + """from_dict must handle legacy entries that predate the tags field.""" + raw = {"chat_id": 1, "thread_id": 2, "summary": "old", "message_count": 3, + "flushed_at": "2026-01-01T00:00:00+00:00"} + s = ThreadSummary.from_dict(raw) + assert s.tags == [] + + 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"])) + store.save(ThreadSummary(chat_id=1, thread_id=2, summary="Database work", message_count=2, + tags=["database", "schema"])) + results = store.search("api") + assert len(results) == 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"])) + 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"])) + 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"])) + 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=[])) + assert store.search("api") == [] + + def test_legacy_store_roundtrip(self, tmp_path: Path): + """Summaries without tags survive a save/load cycle via from_dict.""" + path = tmp_path / "mem.json" + import json as _json + path.write_text( + _json.dumps({ + "1:1": {"chat_id": 1, "thread_id": 1, "summary": "old", + "message_count": 1, "flushed_at": "2026-01-01T00:00:00+00:00"} + }), + encoding="utf-8", + ) + store = ThreadMemoryStore(path) + s = store.get(1, 1) + assert s is not None + assert s.tags == [] + # --------------------------------------------------------------------------- # Thread-aware message_handler tests @@ -221,7 +292,8 @@ async def test_flush_summarises_stores_and_compresses(): settings = _make_settings() store = MagicMock(spec=ThreadMemoryStore) mock_llm = MagicMock(spec=LLMClient) - mock_llm.chat = AsyncMock(return_value="Great summary of the thread.") + # 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] = [ @@ -235,8 +307,8 @@ async def test_flush_summarises_stores_and_compresses(): ctx = _make_context(settings, llm=mock_llm, store=store) await flush_handler(update, ctx) - # LLM called to summarise - mock_llm.chat.assert_awaited_once() + # LLM called twice: once for summary, once for tags + assert mock_llm.chat.await_count == 2 # Summary persisted store.save.assert_called_once() saved: ThreadSummary = store.save.call_args.args[0] @@ -244,6 +316,7 @@ async def test_flush_summarises_stores_and_compresses(): assert saved.thread_id == 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 @@ -316,3 +389,112 @@ async def test_recall_outside_thread_empty_store(tmp_path): update.message.reply_text.assert_awaited_once() text = update.message.reply_text.call_args.args[0] assert "flush" in text.lower() + + +@pytest.mark.asyncio +async def test_recall_with_query_returns_matching_summaries(tmp_path): + """recall_handler with a query argument searches the knowledge base by keyword.""" + settings = _make_settings() + store = ThreadMemoryStore(tmp_path / "mem.json") + store.save(ThreadSummary( + chat_id=1, thread_id=1, summary="API authentication discussion", + message_count=2, tags=["api", "auth"], + )) + store.save(ThreadSummary( + chat_id=1, thread_id=2, summary="Database schema planning", + message_count=3, tags=["database", "schema"], + )) + + update = _make_update(thread_id=None) + ctx = _make_context(settings, store=store) + ctx.args = ["api"] + await recall_handler(update, ctx) + + update.message.reply_text.assert_awaited_once() + text = update.message.reply_text.call_args.args[0] + assert "API authentication" in text + assert "Database" not in text + + +@pytest.mark.asyncio +async def test_recall_with_query_no_match(tmp_path): + """recall_handler with a query that matches nothing tells the user.""" + settings = _make_settings() + store = ThreadMemoryStore(tmp_path / "mem.json") + store.save(ThreadSummary( + chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["database"], + )) + + update = _make_update(thread_id=None) + ctx = _make_context(settings, store=store) + ctx.args = ["deployment"] + await recall_handler(update, ctx) + + update.message.reply_text.assert_awaited_once() + text = update.message.reply_text.call_args.args[0] + assert "deployment" in text.lower() or "no memories" in text.lower() + + +@pytest.mark.asyncio +async def test_message_handler_injects_kb_context_transiently(tmp_path): + """message_handler injects relevant KB summaries as transient context without storing them.""" + settings = _make_settings() + mock_llm = MagicMock(spec=LLMClient) + mock_llm.chat = AsyncMock(return_value="reply with context") + + store = ThreadMemoryStore(tmp_path / "mem.json") + store.save(ThreadSummary( + chat_id=1, thread_id=1, summary="Previous API discussion", + message_count=2, tags=["api"], + )) + + user_id = 999 + from steward.bot.telegram import _history + _history[user_id].clear() + + 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 + mock_llm.chat.assert_awaited_once() + call_kwargs = mock_llm.chat.call_args + # The history passed to the LLM should contain the KB context message + history_arg = call_kwargs.kwargs.get("history") or ( + call_kwargs.args[1] if len(call_kwargs.args) > 1 else None + ) + 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] + ) + + +@pytest.mark.asyncio +async def test_message_handler_no_kb_injection_when_no_match(tmp_path): + """message_handler does not inject KB context when no summaries match.""" + settings = _make_settings() + mock_llm = MagicMock(spec=LLMClient) + mock_llm.chat = AsyncMock(return_value="plain reply") + + store = ThreadMemoryStore(tmp_path / "mem.json") + # Store a summary with unrelated tags + store.save(ThreadSummary( + chat_id=1, thread_id=1, summary="Database recap", message_count=1, tags=["database"], + )) + + user_id = 888 + from steward.bot.telegram import _history + _history[user_id].clear() + + 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 + history_arg = call_kwargs.kwargs.get("history") or [] + # No KB context injected + assert not any("knowledge base" in m.get("content", "").lower() for m in history_arg)