diff --git a/api/core/datasource/datasource_file_manager.py b/api/core/datasource/datasource_file_manager.py index 521791206db..317f362d662 100644 --- a/api/core/datasource/datasource_file_manager.py +++ b/api/core/datasource/datasource_file_manager.py @@ -12,8 +12,8 @@ from uuid import uuid4 import httpx from configs import dify_config +from core.db.session_factory import session_factory from core.file import remote_fetcher -from extensions.ext_database import db from extensions.ext_storage import storage from extensions.storage.storage_type import StorageType from models.enums import CreatorUserRole @@ -54,6 +54,7 @@ class DatasourceFileManager: mimetype: str, filename: str | None = None, ) -> UploadFile: + """Persist an uploaded datasource file and its storage payload.""" extension = guess_extension(mimetype) or ".bin" unique_name = uuid4().hex unique_filename = f"{unique_name}{extension}" @@ -82,9 +83,10 @@ class DatasourceFileManager: created_at=datetime.now(), ) - db.session.add(upload_file) - db.session.commit() - db.session.refresh(upload_file) + with session_factory.create_session() as session: + session.add(upload_file) + session.commit() + session.refresh(upload_file) return upload_file @@ -95,6 +97,7 @@ class DatasourceFileManager: file_url: str, conversation_id: str | None = None, ) -> ToolFile: + """Download a remote file and persist its tool-file metadata.""" # try to download image try: response = remote_fetcher.make_request("GET", file_url) @@ -125,8 +128,9 @@ class DatasourceFileManager: size=len(blob), ) - db.session.add(tool_file) - db.session.commit() + with session_factory.create_session() as session: + session.add(tool_file) + session.commit() return tool_file @@ -139,7 +143,8 @@ class DatasourceFileManager: :return: the binary of the file, mime type """ - upload_file: UploadFile | None = db.session.get(UploadFile, id) + with session_factory.create_session() as session: + upload_file: UploadFile | None = session.get(UploadFile, id) if not upload_file: return None @@ -157,21 +162,24 @@ class DatasourceFileManager: :return: the binary of the file, mime type """ - message_file: MessageFile | None = db.session.get(MessageFile, id) + with session_factory.create_session() as session: + message_file: MessageFile | None = session.get(MessageFile, id) - # Check if message_file is not None - if message_file is not None: - # get tool file id - if message_file.url is not None: - tool_file_id = message_file.url.split("/")[-1] - # trim extension - tool_file_id = tool_file_id.split(".")[0] + # Check if message_file is not None + if message_file is not None: + # get tool file id + if message_file.url is not None: + tool_file_id = message_file.url.split("/")[-1] + # trim extension + tool_file_id = tool_file_id.split(".")[0] + else: + tool_file_id = None else: tool_file_id = None - else: - tool_file_id = None - tool_file: ToolFile | None = db.session.get(ToolFile, tool_file_id) + if not tool_file_id: + return None + tool_file: ToolFile | None = session.get(ToolFile, tool_file_id) if not tool_file: return None @@ -185,11 +193,12 @@ class DatasourceFileManager: """ get file binary - :param tool_file_id: the id of the tool file + :param upload_file_id: the id of the upload file :return: the binary of the file, mime type """ - upload_file: UploadFile | None = db.session.get(UploadFile, upload_file_id) + with session_factory.create_session() as session: + upload_file: UploadFile | None = session.get(UploadFile, upload_file_id) if not upload_file: return None, None diff --git a/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py b/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py index 827950fae2f..3bbc8db1b66 100644 --- a/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py +++ b/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py @@ -1,14 +1,62 @@ -from datetime import datetime +from datetime import UTC, datetime from unittest.mock import MagicMock, patch import httpx import pytest +from sqlalchemy.orm import Session from core.datasource.datasource_file_manager import DatasourceFileManager +from extensions.storage.storage_type import StorageType +from models.enums import CreatorUserRole from models.model import MessageFile, UploadFile from models.tools import ToolFile +def _upload_file(id: str, *, key: str, mime_type: str) -> UploadFile: + upload_file = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key=key, + name="file.png", + size=4, + extension=".png", + mime_type=mime_type, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + created_at=datetime.now(UTC).replace(tzinfo=None), + used=False, + ) + upload_file.id = id + return upload_file + + +def _tool_file(id: str, *, key: str = "tool_key", mimetype: str = "image/png") -> ToolFile: + tool_file = ToolFile( + tenant_id="tenant-1", + user_id="user-1", + conversation_id=None, + file_key=key, + mimetype=mimetype, + name="tool.png", + size=4, + ) + tool_file.id = id + return tool_file + + +def _message_file(id: str, *, url: str | None) -> MessageFile: + message_file = MessageFile( + message_id="message-1", + type="image", + transfer_method="remote_url", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + url=url, + ) + message_file.id = id + return message_file + + class TestDatasourceFileManager: @patch("core.datasource.datasource_file_manager.time.time") @patch("core.datasource.datasource_file_manager.os.urandom") @@ -32,11 +80,10 @@ class TestDatasourceFileManager: assert f"nonce={mock_urandom.return_value.hex()}" in signed_url assert "sign=" in signed_url - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, sqlite_session: Session): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -64,14 +111,14 @@ class TestDatasourceFileManager: assert upload_file.key == f"datasources/{tenant_id}/unique_hex.png" mock_storage.save.assert_called_once_with(upload_file.key, file_binary) - mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + persisted_file = sqlite_session.get(UploadFile, upload_file.id) + assert persisted_file is not None + assert persisted_file.key == upload_file.key - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -94,12 +141,11 @@ class TestDatasourceFileManager: # Verify assert upload_file.name == "test.png" # Should append extension - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") @patch("core.datasource.datasource_file_manager.guess_extension") - def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage): # Setup mock_guess_ext.return_value = None # Cannot guess mock_uuid.return_value = MagicMock(hex="unique_hex") @@ -118,11 +164,10 @@ class TestDatasourceFileManager: assert upload_file.extension == ".bin" assert upload_file.name == "unique_hex.bin" - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -141,10 +186,9 @@ class TestDatasourceFileManager: assert upload_file.extension == ".pdf" @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -154,17 +198,18 @@ class TestDatasourceFileManager: # Execute tool_file = DatasourceFileManager.create_file_by_url( - user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.png" + user_id="user_123", + tenant_id="tenant_456", + file_url="https://example.com/photo.png", ) # Verify assert tool_file.mimetype == "image/png" # Guessed from .png in URL @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -183,10 +228,9 @@ class TestDatasourceFileManager: assert tool_file.mimetype == "application/octet-stream" @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -196,7 +240,9 @@ class TestDatasourceFileManager: # Execute tool_file = DatasourceFileManager.create_file_by_url( - user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.jpg" + user_id="user_123", + tenant_id="tenant_456", + file_url="https://example.com/photo.jpg", ) # Verify @@ -216,152 +262,59 @@ class TestDatasourceFileManager: user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/large.file" ) - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary(self, mock_storage, mock_db): - # Setup - mock_upload_file = UploadFile( - tenant_id="tenant-id", - storage_type="opendal", - key="some_key", - name="test.txt", - size=0, - extension="txt", - mime_type="image/png", - created_by_role="account", - created_by="account-id", - created_at=datetime.now(), - used=False, - ) - - mock_db.session.get.return_value = mock_upload_file + def test_get_file_binary(self, mock_storage, sqlite_session: Session): + sqlite_session.add(_upload_file("file_id", key="some_key", mime_type="image/png")) + sqlite_session.commit() mock_storage.load_once.return_value = b"file content" - # Execute result = DatasourceFileManager.get_file_binary("file_id") # Verify assert result == (b"file content", "image/png") - # Case: Not found - mock_db.session.get.return_value = None assert DatasourceFileManager.get_file_binary("unknown") is None - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id(self, mock_storage, mock_db): - # Setup - mock_message_file = MessageFile( - message_id="message-id", - type="image", - transfer_method="local_file", - created_by_role="user", - created_by="user-id", - url="http://localhost/files/tools/tool_id.png", + def test_get_file_binary_by_message_file_id(self, mock_storage, sqlite_session: Session): + sqlite_session.add_all( + [ + _message_file("msg_file_id", url="http://localhost/files/tools/tool_id.png"), + _tool_file("tool_id"), + ] ) - - mock_tool_file = ToolFile( - user_id="user-id", - tenant_id="tenant-id", - conversation_id="conversation-id", - file_key="tool_key", - mimetype="image/png", - ) - - def mock_get(model, id): - if model == MessageFile: - return mock_message_file - elif model == ToolFile: - return mock_tool_file - return None - - mock_db.session.get.side_effect = mock_get + sqlite_session.commit() mock_storage.load_once.return_value = b"tool content" - # Execute result = DatasourceFileManager.get_file_binary_by_message_file_id("msg_file_id") # Verify assert result == (b"tool content", "image/png") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, mock_db): - # Test that it correctly parses tool_id even with extension in URL - mock_message_file = MessageFile( - message_id="message-id", - type="image", - transfer_method="local_file", - created_by_role="user", - created_by="user-id", - url="http://localhost/files/tools/abcdef.png", + def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, sqlite_session: Session): + sqlite_session.add_all( + [_message_file("m", url="http://localhost/files/tools/abcdef.png"), _tool_file("abcdef", key="tk")] ) - - mock_tool_file = ToolFile( - user_id="user-id", - tenant_id="tenant-id", - conversation_id="conversation-id", - file_key="tk", - mimetype="image/png", - ) - mock_tool_file.id = "abcdef" - - def mock_get(model, id): - if model == MessageFile: - return mock_message_file - return mock_tool_file - - mock_db.session.get.side_effect = mock_get + sqlite_session.commit() mock_storage.load_once.return_value = b"bits" result = DatasourceFileManager.get_file_binary_by_message_file_id("m") assert result == (b"bits", "image/png") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id_failures(self, mock_storage, mock_db): - # Case 1: Message file not found - mock_db.session.get.return_value = None + def test_get_file_binary_by_message_file_id_failures(self, mock_storage, sqlite_session: Session): assert DatasourceFileManager.get_file_binary_by_message_file_id("none") is None - # Case 2: Message file found but tool file not found - mock_message_file = MessageFile( - message_id="message-id", - type="image", - transfer_method="local_file", - created_by_role="user", - created_by="user-id", - url=None, - ) - - def mock_get_v2(model, id): - if model == MessageFile: - return mock_message_file - return None - - mock_db.session.get.side_effect = mock_get_v2 + sqlite_session.add(_message_file("msg_id", url=None)) + sqlite_session.commit() assert DatasourceFileManager.get_file_binary_by_message_file_id("msg_id") is None - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_generator_by_upload_file_id(self, mock_storage, mock_db): - # Setup - mock_upload_file = UploadFile( - tenant_id="tenant-id", - storage_type="opendal", - key="upload_key", - name="test.txt", - size=0, - extension="txt", - mime_type="text/plain", - created_by_role="account", - created_by="account-id", - created_at=datetime.now(), - used=False, - ) - - mock_db.session.get.return_value = mock_upload_file + def test_get_file_generator_by_upload_file_id(self, mock_storage, sqlite_session: Session): + sqlite_session.add(_upload_file("upload_id", key="upload_key", mime_type="text/plain")) + sqlite_session.commit() mock_storage.load_stream.return_value = iter([b"chunk1", b"chunk2"]) @@ -372,8 +325,6 @@ class TestDatasourceFileManager: assert mimetype == "text/plain" assert list(stream) == [b"chunk1", b"chunk2"] - # Case: Not found - mock_db.session.get.return_value = None stream, mimetype = DatasourceFileManager.get_file_generator_by_upload_file_id("none") assert stream is None assert mimetype is None