mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 18:58:35 +08:00
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:
parent
75f7069541
commit
891b2dc537
@ -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,
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user