Michel / test_memory_layers.py
Nanny7's picture
Add LCM-lite memory layers and fix conversation continuity
697288c
Raw History Blame Contribute Delete
5.88 kB
"""Smoke tests for the LCM-lite memory flow in ChatHandler.
This test avoids the real database and network by using a fake DB object and
monkeypatching requests.post. It verifies prompt assembly and post-exchange
memory sequencing.
"""
import asyncio
from backend.chat import ChatHandler
import backend.chat as chat_mod
class FakeResponse:
status_code = 200
text = ""
def iter_lines(self):
for chunk in [
b'data: {"choices":[{"delta":{"content":"Hello"}}]}',
b'data: {"choices":[{"delta":{"content":" there"}}]}',
b"data: [DONE]",
]:
yield chunk
class FakeDB:
def __init__(self):
self.messages = {
1: [
{
"id": 1,
"role": "user",
"content": "Old thread msg",
"timestamp": "t1",
},
{
"id": 2,
"role": "assistant",
"content": "Old reply",
"timestamp": "t2",
},
],
2: [
{
"id": 3,
"role": "user",
"content": "Other conversation noise",
"timestamp": "t3",
}
],
}
self.profile = "Likes Stoicism."
self.facts = [
{
"category": "preference",
"content": "Likes Stoicism.",
"importance": 0.8,
"is_pinned": False,
}
]
self.summary = "They were discussing friendship and doubt."
self.updated_profile = None
self.saved_summary = None
self.last_fact = None
self.title = None
def create_conversation(self, user_id):
return 1
def count_conversation_messages(self, conversation_id):
return len(self.messages.get(conversation_id, []))
def add_message(self, conversation_id, role, content):
items = self.messages.setdefault(conversation_id, [])
items.append(
{
"id": len(items) + 10,
"role": role,
"content": content,
"timestamp": "now",
}
)
def get_user_profile(self, user_id):
return self.profile
def get_user_facts(self, user_id, limit=5):
return self.facts[:limit]
def get_conversation_summary(self, conversation_id):
return self.summary
def get_conversation_turns(self, conversation_id, limit=10):
return self.messages[conversation_id][-limit:]
def refund_free_message(self, user_id):
return None
def get_conversation_title(self, conversation_id):
return None
def get_title_sample_messages(self, conversation_id):
return self.messages[conversation_id]
def update_conversation_title(self, conversation_id, title):
self.title = title
def add_user_fact(self, **kwargs):
self.last_fact = kwargs
def update_user_profile(self, user_id, summary):
self.updated_profile = summary
def save_conversation_summary(
self, conversation_id, user_id, summary, message_count
):
self.saved_summary = (conversation_id, user_id, summary, message_count)
def test_build_context_uses_active_thread_only():
db = FakeDB()
handler = ChatHandler(db)
db.add_message(1, "user", "Current prompt")
ctx = handler._build_context(123, 1)
current_count = sum(
1 for msg in ctx if msg["role"] == "user" and msg["content"] == "Current prompt"
)
other_noise = any(msg.get("content") == "Other conversation noise" for msg in ctx)
assert current_count == 1
assert not other_noise
assert any("Likes Stoicism." in msg["content"] for msg in ctx)
assert any("Summary of this conversation so far" in msg["content"] for msg in ctx)
def test_memory_update_runs_in_order():
db = FakeDB()
handler = ChatHandler(db)
order = []
def fake_extract(*args, **kwargs):
order.append("extract")
def fake_refresh(*args, **kwargs):
order.append("refresh")
handler._extract_user_facts = fake_extract
handler._refresh_user_profile = fake_refresh
handler._update_memory_after_exchange(1, 1, "Hello", "Hi")
assert order == ["extract", "refresh"]
def test_generate_response_streams_and_finishes():
db = FakeDB()
handler = ChatHandler(db)
orig_post = chat_mod.requests.post
orig_extract = ChatHandler._extract_user_facts
orig_refresh = ChatHandler._refresh_user_profile
orig_summary = ChatHandler._generate_conversation_summary
chat_mod.requests.post = lambda *args, **kwargs: FakeResponse()
ChatHandler._extract_user_facts = lambda self, *args, **kwargs: None
ChatHandler._refresh_user_profile = lambda self, *args, **kwargs: None
ChatHandler._generate_conversation_summary = lambda self, *args, **kwargs: None
async def run_once():
chunks = []
async for chunk in handler.generate_response(
123, "Another prompt", conversation_id=1, is_guest=False
):
chunks.append(chunk)
return chunks
try:
chunks = asyncio.run(run_once())
finally:
chat_mod.requests.post = orig_post
ChatHandler._extract_user_facts = orig_extract
ChatHandler._refresh_user_profile = orig_refresh
ChatHandler._generate_conversation_summary = orig_summary
assert any("Hello" in chunk for chunk in chunks)
assert any('"done": true' in chunk for chunk in chunks)
if __name__ == "__main__":
test_build_context_uses_active_thread_only()
test_memory_update_runs_in_order()
test_generate_response_streams_and_finishes()
print("memory layer smoke tests passed")