diff --git a/backend/test/unit/channel/infrastructure/agent/test_agent_adapter.py b/backend/test/unit/channel/infrastructure/agent/test_agent_adapter.py index c9b1c6e7..2369724a 100644 --- a/backend/test/unit/channel/infrastructure/agent/test_agent_adapter.py +++ b/backend/test/unit/channel/infrastructure/agent/test_agent_adapter.py @@ -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