test: use SQLite sessions in unit misc (#39120)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato 2026-07-21 10:51:24 +09:00 committed by GitHub
parent 75f7069541
commit 891b2dc537
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 85 additions and 48 deletions

View File

@ -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,

View File

@ -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:

View File

@ -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