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
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import sys
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
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.message.stream_chat_request import StreamChatRequest
|
||||||
from yuxi.channel.domain.model.shared.channel_type import ChannelType
|
from yuxi.channel.domain.model.shared.channel_type import ChannelType
|
||||||
from yuxi.channel.infrastructure.agent.agent_adapter import AgentAdapter
|
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:
|
class TestAgentAdapter:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_get_service_user_found(self, adapter: AgentAdapter, fake_session_factory) -> None:
|
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
|
result_mock.scalar_one_or_none.return_value = fake_user
|
||||||
db.execute.return_value = result_mock
|
db.execute.return_value = result_mock
|
||||||
|
|
||||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
mock_svc = _make_chat_service_mock([
|
||||||
mock_stream.return_value = async_iter([
|
'{"response": "hello", "status": "success"}',
|
||||||
'{"response": "hello", "status": "success"}',
|
'{"response": " world", "status": "success"}',
|
||||||
'{"response": " world", "status": "success"}',
|
])
|
||||||
])
|
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||||
result = await adapter.stream_chat(stream_request)
|
result = await adapter.stream_chat(stream_request)
|
||||||
assert result == "hello world"
|
assert result == "hello world"
|
||||||
|
|
||||||
@ -104,8 +115,10 @@ class TestAgentAdapter:
|
|||||||
result_mock.scalar_one_or_none.return_value = None
|
result_mock.scalar_one_or_none.return_value = None
|
||||||
db.execute.return_value = result_mock
|
db.execute.return_value = result_mock
|
||||||
|
|
||||||
result = await adapter.stream_chat(stream_request)
|
mock_svc = _make_chat_service_mock([])
|
||||||
assert "Error: no service user found" in result
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_stream_chat_error_status(self, adapter: AgentAdapter, stream_request: StreamChatRequest, fake_session_factory) -> None:
|
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
|
result_mock.scalar_one_or_none.return_value = fake_user
|
||||||
db.execute.return_value = result_mock
|
db.execute.return_value = result_mock
|
||||||
|
|
||||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
mock_svc = _make_chat_service_mock([
|
||||||
mock_stream.return_value = async_iter([
|
'{"response": "hello", "status": "success"}',
|
||||||
'{"response": "hello", "status": "success"}',
|
'{"status": "error", "error_message": "something wrong"}',
|
||||||
'{"status": "error", "error_message": "something wrong"}',
|
])
|
||||||
])
|
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||||
result = await adapter.stream_chat(stream_request)
|
result = await adapter.stream_chat(stream_request)
|
||||||
assert "hello" in result
|
assert "hello" in result
|
||||||
|
|
||||||
@ -135,14 +148,9 @@ class TestAgentAdapter:
|
|||||||
result_mock.scalar_one_or_none.return_value = fake_user
|
result_mock.scalar_one_or_none.return_value = fake_user
|
||||||
db.execute.return_value = result_mock
|
db.execute.return_value = result_mock
|
||||||
|
|
||||||
with patch("yuxi.services.chat_service.stream_agent_chat") as mock_stream:
|
mock_svc = _make_chat_service_mock([
|
||||||
mock_stream.return_value = async_iter([
|
'{"status": "error", "error_message": "something wrong"}',
|
||||||
'{"status": "error", "error_message": "something wrong"}',
|
])
|
||||||
])
|
with patch.dict(sys.modules, {"yuxi.services.chat_service": mock_svc}):
|
||||||
result = await adapter.stream_chat(stream_request)
|
result = await adapter.stream_chat(stream_request)
|
||||||
assert "Error: something wrong" in result
|
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