feat: knowledge-base memory, Dockerfile, docker-compose, CI/release workflows, PR template
This commit is contained in:
parent
c74c8b0d8b
commit
867542dd49
37
.dockerignore
Normal file
37
.dockerignore
Normal 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
21
.github/pull_request_template.md
vendored
Normal 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
62
.github/workflows/ci.yml
vendored
Normal 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
62
.github/workflows/release.yml
vendored
Normal 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 }}
|
||||
3
.release-please-manifest.json
Normal file
3
.release-please-manifest.json
Normal file
@ -0,0 +1,3 @@
|
||||
{
|
||||
".": "0.1.0"
|
||||
}
|
||||
23
Dockerfile
Normal file
23
Dockerfile
Normal 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
14
docker-compose.yml
Normal 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
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user