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."
|
"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:
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user