feat: knowledge-base memory, Dockerfile, docker-compose, CI/release workflows, PR template

This commit is contained in:
copilot-swe-agent[bot] 2026-07-25 12:29:21 +00:00 committed by GitHub
parent c74c8b0d8b
commit 867542dd49
11 changed files with 543 additions and 22 deletions

37
.dockerignore Normal file
View File

@ -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

21
.github/pull_request_template.md vendored Normal file
View File

@ -0,0 +1,21 @@
## What does this PR do?
<!-- A brief description of the change and why it's needed. -->
## 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
<!-- Anything reviewers should pay special attention to, or context that doesn't fit above. -->
---
> **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.

62
.github/workflows/ci.yml vendored Normal file
View File

@ -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 }}

62
.github/workflows/release.yml vendored Normal file
View File

@ -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 }}

View File

@ -0,0 +1,3 @@
{
".": "0.1.0"
}

23
Dockerfile Normal file
View File

@ -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"]

14
docker-compose.yml Normal file
View File

@ -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

View File

@ -39,6 +39,20 @@ _FLUSH_SYSTEM_PROMPT = (
"Be precise. Omit pleasantries." "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: 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.""" """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" "/help – show available commands\n"
"/clear – reset conversation history\n" "/clear – reset conversation history\n"
"/flush – summarise and archive this thread's memory\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", "/analyse – run a manual API analysis right now",
parse_mode=ParseMode.MARKDOWN, 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" "/clear – reset conversation history for this context\n"
"/flush – summarise the current thread, store the summary, and compress memory\n" "/flush – summarise the current thread, store the summary, and compress memory\n"
" _(only available inside a message thread)_\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", "/analyse – trigger an immediate API analysis and proposal",
parse_mode=ParseMode.MARKDOWN, parse_mode=ParseMode.MARKDOWN,
) )
@ -183,12 +197,20 @@ async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
system_prompt=_FLUSH_SYSTEM_PROMPT, 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") message_count = sum(1 for m in history if m["role"] == "user")
thread_summary = ThreadSummary( thread_summary = ThreadSummary(
chat_id=chat_id, chat_id=chat_id,
thread_id=thread_id, thread_id=thread_id,
summary=summary_text, summary=summary_text,
message_count=message_count, message_count=message_count,
tags=tags,
) )
store.save(thread_summary) 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: 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). With a query argument (e.g. ``/recall api design``): searches all stored
Outside a thread: lists all stored summaries (newest first). 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"] settings: Settings = context.bot_data["settings"]
store: ThreadMemoryStore = context.bot_data["thread_store"] 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): if user is None or not _is_allowed(user.id, settings):
return 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) key = _thread_key(update)
if key is not None: if key is not None:
chat_id, thread_id = key chat_id, thread_id = key
@ -243,6 +291,8 @@ async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
date = s.flushed_at[:10] date = s.flushed_at[:10]
first_line = s.summary.split("\n")[0][:80] first_line = s.summary.split("\n")[0][:80]
lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026") 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)) 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()) 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: async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
"""Handle plain text messages – forward to LLM and reply. """Handle plain text messages – forward to LLM and reply.
Thread messages: history is stored unbounded under the (chat_id, thread_id) key. 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. 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"] settings: Settings = context.bot_data["settings"]
llm: LLMClient = context.bot_data["llm"] llm: LLMClient = context.bot_data["llm"]
store: ThreadMemoryStore = context.bot_data["thread_store"]
user = update.effective_user user = update.effective_user
if user is None or not _is_allowed(user.id, settings): if user is None or not _is_allowed(user.id, settings):
return return
@ -287,13 +360,15 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
if key is not None: if key is not None:
# Thread message: unbounded history # Thread message: unbounded history
history = _thread_history[key] 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": "user", "content": text})
history.append({"role": "assistant", "content": reply}) history.append({"role": "assistant", "content": reply})
else: else:
# Non-thread message: capped history per user # Non-thread message: capped history per user
history = _history[user.id] 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": "user", "content": text})
history.append({"role": "assistant", "content": reply}) history.append({"role": "assistant", "content": reply})
if len(history) > _MAX_HISTORY * 2: if len(history) > _MAX_HISTORY * 2:

View File

@ -1,4 +1,4 @@
"""Persistent storage for flushed thread summaries.""" """Persistent storage for flushed thread summaries (knowledge-base style)."""
from __future__ import annotations from __future__ import annotations
@ -13,18 +13,36 @@ logger = logging.getLogger(__name__)
@dataclass @dataclass
class ThreadSummary: 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 chat_id: int
thread_id: int thread_id: int
summary: str summary: str
message_count: int message_count: int
flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat()) flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat())
tags: list[str] = field(default_factory=list)
@property @property
def key(self) -> str: def key(self) -> str:
return f"{self.chat_id}:{self.thread_id}" 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: def format_for_telegram(self) -> str:
"""Return a concise Telegram-formatted recall card.""" """Return a concise Telegram-formatted recall card."""
flushed = self.flushed_at[:19].replace("T", " ") flushed = self.flushed_at[:19].replace("T", " ")
@ -32,21 +50,24 @@ class ThreadSummary:
f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n" f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n"
f"_Flushed: {flushed} UTC \u2014 {self.message_count} messages_" 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: 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 Summaries are indexed by keyword tags so they can be recalled contextually
any previous summary for the same (chat_id, thread_id) pair. (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: def __init__(self, path: str | Path = "thread_memory.json") -> None:
self._path = Path(path) 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(): if self._path.exists():
try: try:
raw = json.loads(self._path.read_text(encoding="utf-8")) 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}") raw = self._data.get(f"{chat_id}:{thread_id}")
if raw is None: if raw is None:
return None return None
return ThreadSummary(**raw) # type: ignore[arg-type] return ThreadSummary.from_dict(raw)
def all(self) -> list[ThreadSummary]: def all(self) -> list[ThreadSummary]:
"""Return all stored summaries, newest first.""" """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) 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)

View File

@ -58,11 +58,16 @@ def _make_context(
store: ThreadMemoryStore | None = None, store: ThreadMemoryStore | None = None,
) -> CallbackContext: # type: ignore[type-arg] ) -> CallbackContext: # type: ignore[type-arg]
ctx = MagicMock(spec=CallbackContext) ctx = MagicMock(spec=CallbackContext)
if store is None:
mock_store = MagicMock(spec=ThreadMemoryStore)
mock_store.search.return_value = []
store = mock_store
ctx.bot_data = { ctx.bot_data = {
"settings": settings, "settings": settings,
"llm": llm, "llm": llm,
"thread_store": store or MagicMock(spec=ThreadMemoryStore), "thread_store": store,
} }
ctx.args = []
return ctx return ctx

View File

@ -55,11 +55,16 @@ def _make_context(
store: ThreadMemoryStore | None = None, store: ThreadMemoryStore | None = None,
) -> CallbackContext: # type: ignore[type-arg] ) -> CallbackContext: # type: ignore[type-arg]
ctx = MagicMock(spec=CallbackContext) ctx = MagicMock(spec=CallbackContext)
if store is None:
mock_store = MagicMock(spec=ThreadMemoryStore)
mock_store.search.return_value = []
store = mock_store
ctx.bot_data = { ctx.bot_data = {
"settings": settings, "settings": settings,
"llm": llm or MagicMock(spec=LLMClient), "llm": llm or MagicMock(spec=LLMClient),
"thread_store": store or MagicMock(spec=ThreadMemoryStore), "thread_store": store,
} }
ctx.args = []
return ctx return ctx
@ -122,6 +127,72 @@ class TestThreadMemoryStore:
assert "77" in text assert "77" in text
assert "recap" 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 # Thread-aware message_handler tests
@ -221,7 +292,8 @@ async def test_flush_summarises_stores_and_compresses():
settings = _make_settings() settings = _make_settings()
store = MagicMock(spec=ThreadMemoryStore) store = MagicMock(spec=ThreadMemoryStore)
mock_llm = MagicMock(spec=LLMClient) 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) key = (400, 99)
_thread_history[key] = [ _thread_history[key] = [
@ -235,8 +307,8 @@ async def test_flush_summarises_stores_and_compresses():
ctx = _make_context(settings, llm=mock_llm, store=store) ctx = _make_context(settings, llm=mock_llm, store=store)
await flush_handler(update, ctx) await flush_handler(update, ctx)
# LLM called to summarise # LLM called twice: once for summary, once for tags
mock_llm.chat.assert_awaited_once() assert mock_llm.chat.await_count == 2
# Summary persisted # Summary persisted
store.save.assert_called_once() store.save.assert_called_once()
saved: ThreadSummary = store.save.call_args.args[0] 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.thread_id == 99
assert saved.summary == "Great summary of the thread." assert saved.summary == "Great summary of the thread."
assert saved.message_count == 2 # 2 user turns assert saved.message_count == 2 # 2 user turns
assert saved.tags == ["api", "design", "testing"]
# In-memory history replaced with compressed context # In-memory history replaced with compressed context
assert len(_thread_history[key]) == 1 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() update.message.reply_text.assert_awaited_once()
text = update.message.reply_text.call_args.args[0] text = update.message.reply_text.call_args.args[0]
assert "flush" in text.lower() 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)