Files
group_xinghuo_jinrong/tests/test_step12_memory_service.py
T

253 lines
8.8 KiB
Python
Raw Normal View History

"""Step 12 测试:MemoryService — 多轮对话记忆服务"""
import pytest
from unittest.mock import Mock
from datetime import datetime
from app.service.memory_service import MemoryService
class TestMemoryService:
"""测试 MemoryService 基本功能"""
def test_create_session_returns_session_id(self):
"""create_session 应返回新的 session_id"""
service = MemoryService(redis_client=Mock(), db_session=Mock())
session_id = service.create_session(user_id="advisor_1")
assert session_id is not None
assert isinstance(session_id, str)
assert len(session_id) > 0
def test_create_session_stores_metadata(self):
"""create_session 应存储会话元数据"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = service.create_session(user_id="advisor_1", session_type="advisor")
# 验证 Redis 中存储了会话元数据
assert mock_redis.hset.called
call_args = mock_redis.hset.call_args
assert "session:" in call_args[0][0]
assert call_args[1]["user_id"] == "advisor_1"
assert call_args[1]["session_type"] == "advisor"
def test_add_message_appends_to_redis(self):
"""add_message 应将消息追加到 Redis 列表"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session_123"
service.add_message(
session_id=session_id,
role="user",
content="你好",
)
# 验证消息被追加到 Redis
assert mock_redis.rpush.called
call_args = mock_redis.rpush.call_args
assert session_id in call_args[0][0]
def test_add_message_persists_to_db(self):
"""add_message 应持久化到数据库"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session_123"
service.add_message(
session_id=session_id,
role="assistant",
content="您好,我是您的投资顾问",
metadata={"tool_calls": []},
)
# 验证数据库写入
assert mock_db.add.called or mock_db.commit.called
def test_get_context_returns_recent_messages(self):
"""get_context 应返回最近的 N 条消息"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
# Mock Redis 返回消息列表
mock_messages = [
'{"role":"user","content":"第一条"}',
'{"role":"assistant","content":"回复一"}',
'{"role":"user","content":"第二条"}',
'{"role":"assistant","content":"回复二"}',
]
mock_redis.lrange.return_value = mock_messages
session_id = "test_session_123"
context = service.get_context(session_id, max_turns=2)
# 验证返回了最近 2 轮对话(4 条消息)
assert len(context) == 4
assert mock_redis.lrange.called
def test_get_context_handles_empty_session(self):
"""get_context 应处理空会话"""
mock_redis = Mock()
mock_redis.lrange.return_value = []
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "empty_session"
context = service.get_context(session_id)
assert context == []
def test_get_context_limits_by_max_turns(self):
"""get_context 应根据 max_turns 限制返回消息数"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
# Mock 10 条消息(5 轮对话)
mock_messages = []
for i in range(10):
role = "user" if i % 2 == 0 else "assistant"
mock_messages.append(f'{{"role":"{role}","content":"消息{i}"}}')
mock_redis.lrange.return_value = mock_messages
session_id = "test_session"
context = service.get_context(session_id, max_turns=3)
# 应返回最近 3 轮(6 条消息)
assert len(context) == 6
def test_get_full_history_loads_from_db(self):
"""get_full_history 应从数据库加载完整历史"""
mock_redis = Mock()
mock_db = Mock()
# Mock 数据库查询结果
mock_messages = [
Mock(role="user", content="消息1", created_at=datetime.now()),
Mock(role="assistant", content="回复1", created_at=datetime.now()),
]
mock_db.query.return_value.filter.return_value.all.return_value = mock_messages
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session"
history = service.get_full_history(session_id)
assert len(history) == 2
assert mock_db.query.called
def test_delete_session_removes_from_redis(self):
"""delete_session 应从 Redis 删除会话"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session_123"
service.delete_session(session_id)
# 验证 Redis 删除
assert mock_redis.delete.called
def test_delete_session_removes_from_db(self):
"""delete_session 应从数据库删除会话"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session_123"
service.delete_session(session_id)
# 验证数据库删除
assert mock_db.delete.called or mock_db.commit.called
def test_session_exists_returns_true_for_valid_session(self):
"""session_exists 应对有效会话返回 True"""
mock_redis = Mock()
mock_redis.exists.return_value = 1
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session_123"
exists = service.session_exists(session_id)
assert exists is True
def test_session_exists_returns_false_for_invalid_session(self):
"""session_exists 应对无效会话返回 False"""
mock_redis = Mock()
mock_redis.exists.return_value = 0
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "invalid_session"
exists = service.session_exists(session_id)
assert exists is False
class TestMemoryServiceIntegration:
"""测试 MemoryService 集成功能"""
def test_add_message_with_tool_calls(self):
"""add_message 应支持 tool_calls 元数据"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id = "test_session"
tool_calls = [
{"name": "query_holdings", "args": {"customer_id": "CUST-001"}},
]
service.add_message(
session_id=session_id,
role="assistant",
content="",
metadata={"tool_calls": tool_calls},
)
# 验证消息被正确存储
assert mock_redis.rpush.called
def test_get_context_with_metadata(self):
"""get_context 应返回包含元数据的消息"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
# Mock 包含元数据的消息
mock_messages = [
'{"role":"assistant","content":"","metadata":{"tool_calls":[{"name":"test"}]}}',
]
mock_redis.lrange.return_value = mock_messages
session_id = "test_session"
context = service.get_context(session_id)
assert len(context) == 1
assert "metadata" in context[0]
def test_multiple_sessions_isolated(self):
"""多个会话应相互隔离"""
mock_redis = Mock()
mock_db = Mock()
service = MemoryService(redis_client=mock_redis, db_session=mock_db)
session_id_1 = "session_1"
session_id_2 = "session_2"
service.add_message(session_id_1, "user", "会话1的消息")
service.add_message(session_id_2, "user", "会话2的消息")
# 验证 rpush 被调用两次,且使用不同的 key
assert mock_redis.rpush.call_count == 2
calls = mock_redis.rpush.call_args_list
assert calls[0][0][0] != calls[1][0][0]