test(agent adapter): 重构单元测试用例的mock创建逻辑
1. 抽取出公共的_make_chat_service_mock工具函数简化测试代码 2. 使用sys.modules patch替代直接patch目标函数,更贴合真实依赖注入场景 3. 删除冗余的async_iter辅助函数
This commit is contained in:
parent
6c95dc006a
commit
395475628c
@ -1,10 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import yuxi.services.chat_service # noqa: F401
|
||||
from yuxi.channel.domain.model.message.stream_chat_request import StreamChatRequest
|
||||
from yuxi.channel.domain.model.shared.channel_type import ChannelType
|
||||
from yuxi.channel.infrastructure.agent.agent_adapter import AgentAdapter
|
||||
@ -36,6 +36,17 @@ def stream_request() -> StreamChatRequest:
|
||||
)
|
||||
|
||||
|
||||
def _make_chat_service_mock(chunks: list[str]):
|
||||
mock_svc = MagicMock()
|
||||
|
||||
async def _stream_gen(*args, **kwargs):
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
mock_svc.stream_agent_chat = _stream_gen
|
||||
return mock_svc
|
||||
|
||||
|
||||
class TestAgentAdapter:
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_service_user_found(self, adapter: AgentAdapter, fake_session_factory) -> None:
|
||||
@ -89,11 +100,11 @@ class TestAgentAdapter:
|
||||
result_mock.scalar_one_or_none.return_value = fake_user
|
||||
db.execute.return_value = result_mock
|
||||
|
||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
||||
mock_stream.return_value = async_iter([
|
||||
'{"response": "hello", "status": "success"}',
|
||||
'{"response": " world", "status": "success"}',
|
||||
])
|
||||
mock_svc = _make_chat_service_mock([
|
||||
'{"response": "hello", "status": "success"}',
|
||||
'{"response": " world", "status": "success"}',
|
||||
])
|
||||
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||
result = await adapter.stream_chat(stream_request)
|
||||
assert result == "hello world"
|
||||
|
||||
@ -104,8 +115,10 @@ class TestAgentAdapter:
|
||||
result_mock.scalar_one_or_none.return_value = None
|
||||
db.execute.return_value = result_mock
|
||||
|
||||
result = await adapter.stream_chat(stream_request)
|
||||
assert "Error: no service user found" in result
|
||||
mock_svc = _make_chat_service_mock([])
|
||||
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||
result = await adapter.stream_chat(stream_request)
|
||||
assert "Error: no service user found" in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_chat_error_status(self, adapter: AgentAdapter, stream_request: StreamChatRequest, fake_session_factory) -> None:
|
||||
@ -117,11 +130,11 @@ class TestAgentAdapter:
|
||||
result_mock.scalar_one_or_none.return_value = fake_user
|
||||
db.execute.return_value = result_mock
|
||||
|
||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
||||
mock_stream.return_value = async_iter([
|
||||
'{"response": "hello", "status": "success"}',
|
||||
'{"status": "error", "error_message": "something wrong"}',
|
||||
])
|
||||
mock_svc = _make_chat_service_mock([
|
||||
'{"response": "hello", "status": "success"}',
|
||||
'{"status": "error", "error_message": "something wrong"}',
|
||||
])
|
||||
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||
result = await adapter.stream_chat(stream_request)
|
||||
assert "hello" in result
|
||||
|
||||
@ -135,14 +148,9 @@ class TestAgentAdapter:
|
||||
result_mock.scalar_one_or_none.return_value = fake_user
|
||||
db.execute.return_value = result_mock
|
||||
|
||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
||||
mock_stream.return_value = async_iter([
|
||||
'{"status": "error", "error_message": "something wrong"}',
|
||||
])
|
||||
mock_svc = _make_chat_service_mock([
|
||||
'{"status": "error", "error_message": "something wrong"}',
|
||||
])
|
||||
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||
result = await adapter.stream_chat(stream_request)
|
||||
assert "Error: something wrong" in result
|
||||
|
||||
|
||||
async def async_iter(items: list[str]):
|
||||
for item in items:
|
||||
yield item
|
||||
|
||||
Loading…
Reference in New Issue
Block a user