feat: knowledge-base memory, Dockerfile, docker-compose, CI/release workflows, PR template
This commit is contained in:
committed by
GitHub
parent
c74c8b0d8b
commit
867542dd49
+186
-4
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user