diff --git a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py index 501225fdbab..4b41e28ec8e 100644 --- a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py @@ -6,6 +6,8 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest +from sqlalchemy import Engine +from sqlalchemy.orm import sessionmaker from core.callback_handler.workflow_tool_callback_handler import DifyWorkflowCallbackHandler from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError @@ -26,7 +28,7 @@ from tests.workflow_test_utils import build_test_graph_init_params, build_test_v @pytest.fixture -def runtime(monkeypatch) -> DifyToolNodeRuntime: +def runtime(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> DifyToolNodeRuntime: module_name = "core.ops.ops_trace_manager" if module_name not in sys.modules: ops_stub = types.ModuleType(module_name) @@ -44,9 +46,7 @@ def runtime(monkeypatch) -> DifyToolNodeRuntime: invoke_from="debugger", call_depth=0, ) - session_maker = MagicMock() - session_maker.begin.return_value.__enter__.return_value = MagicMock(name="session") - session_maker.begin.return_value.__exit__.return_value = None + session_maker = sessionmaker(sqlite_engine, expire_on_commit=False) return DifyToolNodeRuntime(init_params.run_context, session_maker=session_maker) diff --git a/api/tests/unit_tests/core/workflow/test_node_runtime.py b/api/tests/unit_tests/core/workflow/test_node_runtime.py index 6190c7fb91c..adfd7ed2c5f 100644 --- a/api/tests/unit_tests/core/workflow/test_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/test_node_runtime.py @@ -1,8 +1,12 @@ +from collections.abc import Iterator +from datetime import UTC, datetime from types import SimpleNamespace from unittest.mock import MagicMock, Mock, sentinel from uuid import uuid4 import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext, InvokeFrom, UserFrom from core.app.file_access import FileAccessScope, bind_file_access_scope, grant_retriever_segment_access @@ -44,9 +48,62 @@ from graphon.model_runtime.model_providers.base.large_language_model import Larg from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol from graphon.nodes.tool.entities import ToolNodeData, ToolProviderType from graphon.variables.segments import ArrayFileSegment, FileSegment +from models.base import TypeBase +from models.dataset import SegmentAttachmentBinding +from models.enums import CreatorUserRole +from models.model import StorageType, UploadFile from tests.workflow_test_utils import build_test_run_context +@pytest.fixture +def attachment_session(sqlite_engine: Engine) -> Iterator[Session]: + """Provide real attachment and upload-file persistence to node runtime tests.""" + + TypeBase.metadata.create_all(sqlite_engine, tables=[SegmentAttachmentBinding.__table__, UploadFile.__table__]) + with Session(sqlite_engine, expire_on_commit=False) as session: + yield session + + +def _persist_attachment( + session: Session, + *, + segment_id: str, + upload_file_id: str, + upload_file_tenant_id: str = "tenant-id", +) -> UploadFile: + """Persist an attachment binding for the test tenant and its referenced upload file.""" + + upload_file = UploadFile( + tenant_id=upload_file_tenant_id, + storage_type=StorageType.LOCAL, + key="storage-key", + name="diagram.png", + size=128, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-id", + created_at=datetime.now(UTC).replace(tzinfo=None), + used=False, + source_url="https://example.com/diagram.png", + ) + upload_file.id = upload_file_id + session.add_all( + [ + upload_file, + SegmentAttachmentBinding( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id=segment_id, + attachment_id=upload_file_id, + ), + ] + ) + session.commit() + return upload_file + + def _build_model_schema(*, features: list[ModelFeature] | None = None) -> AIModelEntity: return AIModelEntity( model="gpt-4o-mini", @@ -348,29 +405,12 @@ def test_dify_prompt_message_serializer_delegates(monkeypatch: pytest.MonkeyPatc ) -def test_dify_retriever_attachment_loader_builds_graph_files(monkeypatch: pytest.MonkeyPatch) -> None: - upload_file = SimpleNamespace( - id="upload-file-id", - name="diagram.png", - extension="png", - mime_type="image/png", - source_url="https://example.com/diagram.png", - key="storage-key", - size=128, - ) - session = MagicMock() - session.execute.return_value.all.return_value = [(None, upload_file)] - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - +def test_dify_retriever_attachment_loader_builds_graph_files( + monkeypatch: pytest.MonkeyPatch, attachment_session: Session +) -> None: + _persist_attachment(attachment_session, segment_id="segment-id", upload_file_id="upload-file-id") build_from_mapping = MagicMock(return_value=sentinel.file) - monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr(node_runtime, "Session", MagicMock(return_value=_SessionContext())) + monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=attachment_session.get_bind())) loader = DifyRetrieverAttachmentLoader( file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping) ) @@ -388,39 +428,18 @@ def test_dify_retriever_attachment_loader_builds_graph_files(monkeypatch: pytest def test_dify_retriever_attachment_loader_grants_upload_files_for_allowed_segment( monkeypatch: pytest.MonkeyPatch, + attachment_session: Session, ) -> None: from factories.file_factory import builders as file_builders upload_file_id = str(uuid4()) segment_id = str(uuid4()) - upload_file = SimpleNamespace( - id=upload_file_id, - tenant_id="tenant-id", - name="diagram.png", - extension="png", - mime_type="image/png", - source_url="https://example.com/diagram.png", - key="storage-key", - size=128, - ) - attachment_session = MagicMock() - attachment_session.execute.return_value.all.return_value = [(None, upload_file)] - - class _AttachmentSessionContext: - def __enter__(self): - return attachment_session - - def __exit__(self, exc_type, exc, tb): - return False - - upload_session = MagicMock() - upload_session.__enter__.return_value = upload_session - upload_session.__exit__.return_value = False - upload_session.scalar.return_value = upload_file - - monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr(node_runtime, "Session", MagicMock(return_value=_AttachmentSessionContext())) - monkeypatch.setattr(file_builders, "session_factory", SimpleNamespace(create_session=lambda: upload_session)) + _persist_attachment(attachment_session, segment_id=segment_id, upload_file_id=upload_file_id) + engine = attachment_session.get_bind() + assert engine is not None + monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine)) + session_maker = sessionmaker(engine, expire_on_commit=False) + monkeypatch.setattr(file_builders.session_factory, "create_session", session_maker) loader = DifyRetrieverAttachmentLoader(file_reference_factory=DifyFileReferenceFactory(_build_run_context())) scope = FileAccessScope( @@ -435,18 +454,57 @@ def test_dify_retriever_attachment_loader_grants_upload_files_for_allowed_segmen files = loader.load(segment_id=segment_id) assert files[0].related_id == upload_file_id - stmt = upload_session.scalar.call_args.args[0] - whereclause = str(stmt.whereclause) - assert "upload_files.tenant_id" in whereclause - assert "upload_files.id IN" in whereclause + assert files[0].filename == "diagram.png" + + +def test_dify_retriever_attachment_loader_rejects_granted_upload_file_from_another_tenant( + monkeypatch: pytest.MonkeyPatch, + attachment_session: Session, +) -> None: + from factories.file_factory import builders as file_builders + + upload_file_id = str(uuid4()) + segment_id = str(uuid4()) + _persist_attachment( + attachment_session, + segment_id=segment_id, + upload_file_id=upload_file_id, + upload_file_tenant_id="other-tenant-id", + ) + engine = attachment_session.get_bind() + assert engine is not None + monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine)) + monkeypatch.setattr(file_builders.session_factory, "create_session", sessionmaker(engine, expire_on_commit=False)) + + loader = DifyRetrieverAttachmentLoader(file_reference_factory=DifyFileReferenceFactory(_build_run_context())) + scope = FileAccessScope( + tenant_id="tenant-id", + user_id="end-user-id", + user_from=UserFrom.END_USER, + invoke_from=InvokeFrom.WEB_APP, + ) + + with bind_file_access_scope(scope): + grant_retriever_segment_access([segment_id]) + with pytest.raises(ValueError, match="Invalid upload file"): + loader.load(segment_id=segment_id) def test_dify_retriever_attachment_loader_skips_ungranted_segment_for_end_user( monkeypatch: pytest.MonkeyPatch, + attachment_session: Session, ) -> None: build_from_mapping = MagicMock() - session_factory = MagicMock() - monkeypatch.setattr(node_runtime, "Session", session_factory) + engine = attachment_session.get_bind() + assert engine is not None + monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine)) + statement_count = 0 + + def count_statements(*_args, **_kwargs) -> None: + nonlocal statement_count + statement_count += 1 + + event.listen(engine, "before_cursor_execute", count_statements) loader = DifyRetrieverAttachmentLoader( file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping) ) @@ -460,19 +518,31 @@ def test_dify_retriever_attachment_loader_skips_ungranted_segment_for_end_user( with bind_file_access_scope(scope): files = loader.load(segment_id=str(uuid4())) - assert files == [] - session_factory.assert_not_called() - build_from_mapping.assert_not_called() + try: + assert files == [] + assert statement_count == 0 + build_from_mapping.assert_not_called() + finally: + event.remove(engine, "before_cursor_execute", count_statements) def test_dify_retriever_attachment_loader_skips_segment_rejected_by_checker( monkeypatch: pytest.MonkeyPatch, + attachment_session: Session, ) -> None: segment_id = str(uuid4()) build_from_mapping = MagicMock() - session_factory = MagicMock() segment_access_checker = MagicMock(return_value=False) - monkeypatch.setattr(node_runtime, "Session", session_factory) + engine = attachment_session.get_bind() + assert engine is not None + monkeypatch.setattr(node_runtime, "db", SimpleNamespace(engine=engine)) + statement_count = 0 + + def count_statements(*_args, **_kwargs) -> None: + nonlocal statement_count + statement_count += 1 + + event.listen(engine, "before_cursor_execute", count_statements) loader = DifyRetrieverAttachmentLoader( file_reference_factory=SimpleNamespace(build_from_mapping=build_from_mapping), segment_access_checker=segment_access_checker, @@ -488,10 +558,13 @@ def test_dify_retriever_attachment_loader_skips_segment_rejected_by_checker( grant_retriever_segment_access([segment_id]) files = loader.load(segment_id=segment_id) - assert files == [] - segment_access_checker.assert_called_once_with(segment_id) - session_factory.assert_not_called() - build_from_mapping.assert_not_called() + try: + assert files == [] + segment_access_checker.assert_called_once_with(segment_id) + assert statement_count == 0 + build_from_mapping.assert_not_called() + finally: + event.remove(engine, "before_cursor_execute", count_statements) def test_dify_tool_file_manager_resolves_conversation_id_for_tool_files(monkeypatch: pytest.MonkeyPatch) -> None: