"""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]