From 891b2dc537cc73e1ce06e74ef510c2a8d7733fb6 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 21 Jul 2026 10:51:24 +0900 Subject: [PATCH] test: use SQLite sessions in unit misc (#39120) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../layers/test_pause_state_persist_layer.py | 65 +++++++++++-------- api/tests/unit_tests/tools/test_api_tool.py | 34 +++++++--- api/tests/unit_tests/tools/test_mcp_tool.py | 34 ++++++---- 3 files changed, 85 insertions(+), 48 deletions(-) diff --git a/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py b/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py index ff7f27f5efa..b8bceb3a988 100644 --- a/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py +++ b/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py @@ -4,6 +4,8 @@ from time import time from unittest.mock import Mock import pytest +from sqlalchemy import Engine +from sqlalchemy.orm import Session, sessionmaker from core.app.app_config.entities import WorkflowUIBasedAppConfig from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity @@ -32,6 +34,13 @@ from models.model import AppMode from repositories.factory import DifyAPIRepositoryFactory +@pytest.fixture +def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + """Provide the real session factory injected into the persistence layer.""" + + return sessionmaker(sqlite_engine, expire_on_commit=False) + + def _create_initialized_response_stream_filter() -> ResponseStreamFilter: """Build a `ResponseStreamFilter` that has already run `initialize()`. @@ -211,27 +220,25 @@ class TestPauseStatePersistenceLayer: workflow_execution_id=workflow_execution_id, ) - def test_init_with_dependency_injection(self): - session_factory = Mock(name="session_factory") + def test_init_with_dependency_injection(self, sqlite_session_factory: sessionmaker[Session]): state_owner_user_id = "user-123" layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id=state_owner_user_id, generate_entity=self._create_generate_entity(), response_stream_filter=ResponseStreamFilter(), ) - assert layer._session_maker is session_factory + assert layer._session_maker is sqlite_session_factory assert layer._state_owner_user_id == state_owner_user_id with pytest.raises(GraphEngineLayerNotInitializedError): _ = layer.graph_runtime_state assert layer.command_channel is None - def test_initialize_sets_dependencies(self): - session_factory = Mock(name="session_factory") + def test_initialize_sets_dependencies(self, sqlite_session_factory: sessionmaker[Session]): layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner", generate_entity=self._create_generate_entity(), response_stream_filter=ResponseStreamFilter(), @@ -245,11 +252,12 @@ class TestPauseStatePersistenceLayer: assert layer.graph_runtime_state is graph_runtime_state assert layer.command_channel is command_channel - def test_on_event_with_graph_run_paused_event(self, monkeypatch: pytest.MonkeyPatch): - session_factory = Mock(name="session_factory") + def test_on_event_with_graph_run_paused_event( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] + ): generate_entity = self._create_generate_entity(workflow_execution_id="run-123") layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=generate_entity, response_stream_filter=_create_initialized_response_stream_filter(), @@ -272,7 +280,7 @@ class TestPauseStatePersistenceLayer: layer.on_event(event) - mock_factory.assert_called_once_with(session_factory) + mock_factory.assert_called_once_with(sqlite_session_factory) assert mock_repo.create_workflow_pause.call_count == 1 call_kwargs = mock_repo.create_workflow_pause.call_args.kwargs assert call_kwargs["workflow_run_id"] == "run-123" @@ -285,11 +293,12 @@ class TestPauseStatePersistenceLayer: assert isinstance(pause_reasons, list) - def test_on_event_enriches_hitl_pause_reasons_before_persisting(self, monkeypatch: pytest.MonkeyPatch): - session_factory = Mock(name="session_factory") + def test_on_event_enriches_hitl_pause_reasons_before_persisting( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] + ): generate_entity = self._create_generate_entity(workflow_execution_id="run-123") layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=generate_entity, response_stream_filter=_create_initialized_response_stream_filter(), @@ -343,10 +352,11 @@ class TestPauseStatePersistenceLayer: ) assert mock_repo.create_workflow_pause.call_args.kwargs["pause_reasons"] == [enriched_reason] - def test_on_event_ignores_non_paused_events(self, monkeypatch: pytest.MonkeyPatch): - session_factory = Mock(name="session_factory") + def test_on_event_ignores_non_paused_events( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] + ): layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=self._create_generate_entity(), response_stream_filter=ResponseStreamFilter(), @@ -372,10 +382,11 @@ class TestPauseStatePersistenceLayer: mock_factory.assert_not_called() mock_repo.create_workflow_pause.assert_not_called() - def test_on_event_raises_when_graph_runtime_state_is_uninitialized(self): - session_factory = Mock(name="session_factory") + def test_on_event_raises_when_graph_runtime_state_is_uninitialized( + self, sqlite_session_factory: sessionmaker[Session] + ): layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=self._create_generate_entity(), response_stream_filter=ResponseStreamFilter(), @@ -386,10 +397,11 @@ class TestPauseStatePersistenceLayer: with pytest.raises(GraphEngineLayerNotInitializedError): layer.on_event(event) - def test_on_event_asserts_when_workflow_execution_id_missing(self, monkeypatch: pytest.MonkeyPatch): - session_factory = Mock(name="session_factory") + def test_on_event_asserts_when_workflow_execution_id_missing( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] + ): layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=self._create_generate_entity(), response_stream_filter=_create_initialized_response_stream_filter(), @@ -494,12 +506,13 @@ def test_workflow_resumption_context_dumps_loads_roundtrip(state: WorkflowResump assert restored_entity.extras["trace_session_id"] == "session-1" -def test_on_event_persists_response_stream_filter_dump(monkeypatch: pytest.MonkeyPatch) -> None: - session_factory = Mock(name="session_factory") +def test_on_event_persists_response_stream_filter_dump( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: generate_entity = TestPauseStatePersistenceLayer._create_generate_entity(workflow_execution_id="run-123") response_stream_filter = _create_initialized_response_stream_filter() layer = PauseStatePersistenceLayer( - session_factory=session_factory, + session_factory=sqlite_session_factory, state_owner_user_id="owner-123", generate_entity=generate_entity, response_stream_filter=response_stream_filter, diff --git a/api/tests/unit_tests/tools/test_api_tool.py b/api/tests/unit_tests/tools/test_api_tool.py index e62f74b8019..0c5c75121a3 100644 --- a/api/tests/unit_tests/tools/test_api_tool.py +++ b/api/tests/unit_tests/tools/test_api_tool.py @@ -1,9 +1,12 @@ import json import operator +from collections.abc import Iterator from unittest.mock import Mock, patch import httpx import pytest +from sqlalchemy import Engine +from sqlalchemy.orm import Session from core.tools.__base.tool_runtime import ToolRuntime from core.tools.custom_tool.tool import ApiTool @@ -16,6 +19,13 @@ from core.tools.entities.tool_entities import ( ) +@pytest.fixture +def database_session(sqlite_engine: Engine) -> Iterator[Session]: + """Yield the live ORM session required by the shared tool invocation contract.""" + with Session(sqlite_engine) as session: + yield session + + def _get_message_by_type[T](msgs: list[ToolInvokeMessage], msg_type: type[T]) -> ToolInvokeMessage | None: return next((i for i in msgs if isinstance(i.message, msg_type)), None) @@ -57,7 +67,9 @@ class TestApiToolInvoke: ) @patch("core.tools.custom_tool.tool.ssrf_proxy.get") - def test_invoke_with_json_response_creates_text_message_with_serialized_json(self, mock_get: Mock) -> None: + def test_invoke_with_json_response_creates_text_message_with_serialized_json( + self, mock_get: Mock, database_session: Session + ) -> None: """Test that when upstream returns JSON, the output Text message contains JSON-serialized string.""" # Setup mock response with JSON content json_response_data = { @@ -74,7 +86,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=database_session, user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -135,7 +147,7 @@ class TestApiToolInvoke: ids=operator.itemgetter(0), ) def test_invoke_with_non_dict_json_response_creates_text_message_with_serialized_json( - self, mock_get: Mock, test_case + self, mock_get: Mock, test_case, database_session: Session ) -> None: """Test that when upstream returns a non-dict JSON, the output Text message contains JSON-serialized string.""" # Setup mock response with non-dict JSON content @@ -149,7 +161,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=database_session, user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -173,7 +185,9 @@ class TestApiToolInvoke: assert json_message is None, "_invoke should not yield a JSON message for JSON array response" @patch("core.tools.custom_tool.tool.ssrf_proxy.get") - def test_invoke_with_text_response_creates_text_message_with_original_text(self, mock_get: Mock) -> None: + def test_invoke_with_text_response_creates_text_message_with_original_text( + self, mock_get: Mock, database_session: Session + ) -> None: """Test that when upstream returns plain text, the output Text message contains the original text.""" # Setup mock response with plain text content text_response_data = "This is a plain text response" @@ -186,7 +200,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=database_session, user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -202,7 +216,7 @@ class TestApiToolInvoke: assert message.message.text == text_response_data @patch("core.tools.custom_tool.tool.ssrf_proxy.get") - def test_invoke_with_empty_response(self, mock_get: Mock) -> None: + def test_invoke_with_empty_response(self, mock_get: Mock, database_session: Session) -> None: """Test that empty responses are handled correctly.""" # Setup mock response with empty content mock_response = Mock(spec=httpx.Response) @@ -212,7 +226,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=database_session, user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -228,7 +242,7 @@ class TestApiToolInvoke: assert "Empty response from the tool" in message.message.text @patch("core.tools.custom_tool.tool.ssrf_proxy.get") - def test_invoke_with_error_response(self, mock_get: Mock) -> None: + def test_invoke_with_error_response(self, mock_get: Mock, database_session: Session) -> None: """Test that error responses are handled correctly.""" # Setup mock response with error status code mock_response = Mock(spec=httpx.Response) @@ -236,7 +250,7 @@ class TestApiToolInvoke: mock_response.text = "Not Found" mock_get.return_value = mock_response - result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=database_session, user_id="test_user", tool_parameters={}) # Invoke the tool and expect an error with pytest.raises(Exception) as exc_info: diff --git a/api/tests/unit_tests/tools/test_mcp_tool.py b/api/tests/unit_tests/tools/test_mcp_tool.py index 5984b3b6744..eb7bb34dbe4 100644 --- a/api/tests/unit_tests/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/tools/test_mcp_tool.py @@ -1,9 +1,12 @@ import base64 +from collections.abc import Iterator from decimal import Decimal from typing import Any from unittest.mock import Mock, patch import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from core.mcp.types import ( AudioContent, @@ -21,6 +24,13 @@ from core.tools.mcp_tool.tool import MCPTool from graphon.model_runtime.entities.llm_entities import LLMUsage +@pytest.fixture +def orm_session(sqlite_engine: Engine) -> Iterator[Session]: + """Use a real ORM session while MCP transport remains mocked.""" + with Session(sqlite_engine) as session: + yield session + + def _make_mcp_tool(output_schema: dict[str, Any] | None = None) -> MCPTool: identity = ToolIdentity( author="test", @@ -56,7 +66,7 @@ class TestMCPToolInvoke: ), ], ) - def test_invoke_image_or_audio_yields_blob(self, content_factory, mime_type) -> None: + def test_invoke_image_or_audio_yields_blob(self, content_factory, mime_type, orm_session: Session) -> None: tool = _make_mcp_tool() raw = b"\x00\x01test-bytes\x02" b64 = base64.b64encode(raw).decode() @@ -64,7 +74,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -73,14 +83,14 @@ class TestMCPToolInvoke: assert msg.message.blob == raw assert msg.meta == {"mime_type": mime_type} - def test_invoke_embedded_text_resource_yields_text(self) -> None: + def test_invoke_embedded_text_resource_yields_text(self, orm_session: Session) -> None: tool = _make_mcp_tool() text_resource = TextResourceContents(uri="file://test.txt", mimeType="text/plain", text="hello world") content = EmbeddedResource(type="resource", resource=text_resource) result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -92,7 +102,7 @@ class TestMCPToolInvoke: ("mime_type", "expected_mime"), [("application/pdf", "application/pdf"), (None, "application/octet-stream")], ) - def test_invoke_embedded_blob_resource_yields_blob(self, mime_type, expected_mime) -> None: + def test_invoke_embedded_blob_resource_yields_blob(self, mime_type, expected_mime, orm_session: Session) -> None: tool = _make_mcp_tool() raw = b"binary-data" b64 = base64.b64encode(raw).decode() @@ -101,7 +111,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -110,12 +120,12 @@ class TestMCPToolInvoke: assert msg.message.blob == raw assert msg.meta == {"mime_type": expected_mime} - def test_invoke_yields_variables_when_structured_content_and_schema(self) -> None: + def test_invoke_yields_variables_when_structured_content_and_schema(self, orm_session: Session) -> None: tool = _make_mcp_tool(output_schema={"type": "object"}) result = CallToolResult(content=[], structuredContent={"a": 1, "b": "x"}) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) # Expect two variable messages corresponding to keys a and b assert len(messages) == 2 @@ -266,7 +276,7 @@ class TestMCPToolUsageExtraction: assert usage.prompt_tokens == 100 assert usage.completion_tokens == 50 - def test_invoke_sets_latest_usage_from_meta(self) -> None: + def test_invoke_sets_latest_usage_from_meta(self, orm_session: Session) -> None: """Test that _invoke sets _latest_usage from result meta.""" tool = _make_mcp_tool() meta = { @@ -281,7 +291,7 @@ class TestMCPToolUsageExtraction: result = CallToolResult(content=[TextContent(type="text", text="test")], _meta=meta) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) # Verify latest_usage was set correctly assert tool.latest_usage.prompt_tokens == 200 @@ -289,13 +299,13 @@ class TestMCPToolUsageExtraction: assert tool.latest_usage.total_tokens == 300 assert tool.latest_usage.total_price == Decimal("0.003") - def test_invoke_with_no_meta_returns_empty_usage(self) -> None: + def test_invoke_with_no_meta_returns_empty_usage(self, orm_session: Session) -> None: """Test that _invoke returns empty usage when no meta is present.""" tool = _make_mcp_tool() result = CallToolResult(content=[TextContent(type="text", text="test")], _meta=None) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) + list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) # Verify latest_usage is empty assert tool.latest_usage.total_tokens == 0