From 1e61078e93aed8dd065cca509ebdb5669fb084e1 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 28 Jul 2026 12:18:32 +0900 Subject: [PATCH] test: use SQLite sessions in services core (#39088) --- .../unit_tests/services/test_file_service.py | 235 +++++++++--------- 1 file changed, 111 insertions(+), 124 deletions(-) diff --git a/api/tests/unit_tests/services/test_file_service.py b/api/tests/unit_tests/services/test_file_service.py index 583509e20db..0278887e010 100644 --- a/api/tests/unit_tests/services/test_file_service.py +++ b/api/tests/unit_tests/services/test_file_service.py @@ -1,6 +1,8 @@ import base64 import hashlib import os +from collections.abc import Iterator +from datetime import UTC, datetime from unittest.mock import MagicMock, patch import pytest @@ -9,6 +11,8 @@ from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import NotFound from configs import dify_config +from extensions.storage.storage_type import StorageType +from models.base import TypeBase from models.enums import CreatorUserRole from models.model import Account, EndUser, UploadFile from services.errors.file import BlockedFileExtensionError, FileTooLargeError, UnsupportedFileTypeError @@ -17,31 +21,54 @@ from services.file_service import FileService class TestFileService: @pytest.fixture - def mock_db_session(self): - session = MagicMock(spec=Session) - # Mock context manager behavior - session.__enter__.return_value = session - return session + def sqlite_session_maker(self, sqlite_engine: Engine) -> sessionmaker[Session]: + TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[UploadFile.__tablename__]]) + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) @pytest.fixture - def mock_session_maker(self, mock_db_session): - maker = MagicMock(spec=sessionmaker) - maker.return_value = mock_db_session - return maker + def db_session(self, sqlite_session_maker: sessionmaker[Session]) -> Iterator[Session]: + with sqlite_session_maker() as session: + yield session @pytest.fixture - def file_service(self, mock_session_maker): - return FileService(session_factory=mock_session_maker) + def file_service(self, sqlite_session_maker: sessionmaker[Session]) -> FileService: + return FileService(session_factory=sqlite_session_maker) - def test_init_with_engine(self): - engine = MagicMock(spec=Engine) - service = FileService(session_factory=engine) + @staticmethod + def _persist_upload_file( + session: Session, + *, + file_id: str = "file_id", + tenant_id: str = "tenant_id", + extension: str = "txt", + mime_type: str = "text/plain", + key: str = "key", + ) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=key, + name=f"test.{extension}", + size=10, + extension=extension, + mime_type=mime_type, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user_id", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + used=False, + ) + upload_file.id = file_id + session.add(upload_file) + session.commit() + return upload_file + + def test_init_with_engine(self, sqlite_engine: Engine): + service = FileService(session_factory=sqlite_engine) assert isinstance(service._session_maker, sessionmaker) - def test_init_with_sessionmaker(self): - maker = MagicMock(spec=sessionmaker) - service = FileService(session_factory=maker) - assert service._session_maker == maker + def test_init_with_sessionmaker(self, sqlite_session_maker: sessionmaker[Session]): + service = FileService(session_factory=sqlite_session_maker) + assert service._session_maker == sqlite_session_maker def test_init_invalid_factory(self): with pytest.raises(AssertionError, match="must be a sessionmaker or an Engine."): @@ -52,11 +79,11 @@ class TestFileService: @patch("services.file_service.extract_tenant_id") @patch("services.file_service.file_helpers.get_signed_file_url") def test_upload_file_success( - self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, mock_db_session + self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, db_session: Session ): # Setup mock_tenant_id.return_value = "tenant_id" - mock_now.return_value = "2024-01-01" + mock_now.return_value = datetime(2024, 1, 1, tzinfo=UTC) mock_get_url.return_value = "http://signed-url" user = MagicMock(spec=Account) @@ -81,8 +108,9 @@ class TestFileService: assert result.source_url == "http://signed-url" mock_storage.save.assert_called_once() - mock_db_session.add.assert_called_once_with(result) - mock_db_session.commit.assert_called_once() + persisted = db_session.get(UploadFile, result.id) + assert persisted is not None + assert persisted.hash == result.hash def test_upload_file_uses_explicit_resource_tenant(self, file_service: FileService): user = MagicMock(spec=Account) @@ -109,7 +137,7 @@ class TestFileService: with pytest.raises(ValueError, match="Filename contains invalid characters"): file_service.upload_file(filename="invalid/file.txt", content=b"", mimetype="text/plain", user=MagicMock()) - def test_upload_file_long_filename(self, file_service: FileService, mock_db_session): + def test_upload_file_long_filename(self, file_service: FileService, db_session: Session): # Setup long_name = "a" * 210 + ".txt" user = MagicMock(spec=Account) @@ -124,6 +152,7 @@ class TestFileService: result = file_service.upload_file(filename=long_name, content=b"test", mimetype="text/plain", user=user) assert len(result.name) <= 205 # 200 + . + extension assert result.name.endswith(".txt") + assert db_session.get(UploadFile, result.id) is not None def test_upload_file_blocked_extension(self, file_service): with patch.object(dify_config, "inner_UPLOAD_FILE_EXTENSION_BLACKLIST", "exe"): @@ -145,7 +174,7 @@ class TestFileService: with pytest.raises(FileTooLargeError): file_service.upload_file(filename="test.jpg", content=content, mimetype="image/jpeg", user=MagicMock()) - def test_upload_file_end_user(self, file_service: FileService, mock_db_session): + def test_upload_file_end_user(self, file_service: FileService, db_session: Session): user = MagicMock(spec=EndUser) user.id = "end_user_id" @@ -157,6 +186,7 @@ class TestFileService: mock_tenant.return_value = "tenant" result = file_service.upload_file(filename="test.txt", content=b"test", mimetype="text/plain", user=user) assert result.created_by_role == CreatorUserRole.END_USER + assert db_session.get(UploadFile, result.id) is not None def test_is_file_size_within_limit(self): with ( @@ -181,12 +211,8 @@ class TestFileService: assert FileService.is_file_size_within_limit(extension="txt", file_size=5 * 1024 * 1024) is True assert FileService.is_file_size_within_limit(extension="pdf", file_size=6 * 1024 * 1024) is False - def test_get_file_base64_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "test_key" - mock_db_session.scalar.return_value = upload_file + def test_get_file_base64_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, key="test_key") with patch("services.file_service.storage") as mock_storage: mock_storage.load_once.return_value = b"test content" @@ -198,16 +224,17 @@ class TestFileService: assert result == base64.b64encode(b"test content").decode() mock_storage.load_once.assert_called_once_with("test_key") - def test_get_file_base64_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_base64_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_base64("non_existent") - def test_get_file_presigned_url_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.key = "upload_files/tenant_id/icon.png" - upload_file.mime_type = "image/png" - mock_db_session.scalar.return_value = upload_file + def test_get_file_presigned_url_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file( + db_session, + extension="png", + mime_type="image/png", + key="upload_files/tenant_id/icon.png", + ) with ( patch.object(dify_config, "FILES_ACCESS_TIMEOUT", 300), @@ -224,13 +251,11 @@ class TestFileService: content_type="image/png", ) - def test_get_file_presigned_url_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None - + def test_get_file_presigned_url_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_presigned_url(file_id="file_id", tenant_id="tenant_id") - def test_upload_text_success(self, file_service: FileService, mock_db_session): + def test_upload_text_success(self, file_service: FileService, db_session: Session): # Setup text = "sample text" text_name = "test.txt" @@ -249,21 +274,17 @@ class TestFileService: assert result.used is True assert result.extension == "txt" mock_storage.save.assert_called_once() - mock_db_session.add.assert_called_once() - mock_db_session.commit.assert_called_once() + assert db_session.get(UploadFile, result.id) is not None - def test_upload_text_long_name(self, file_service: FileService, mock_db_session): + def test_upload_text_long_name(self, file_service: FileService, db_session: Session): long_name = "a" * 210 with patch("services.file_service.storage"): result = file_service.upload_text("text", long_name, "user", "tenant") assert len(result.name) == 200 + assert db_session.get(UploadFile, result.id) is not None - def test_get_file_preview_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "pdf" - mock_db_session.scalar.return_value = upload_file + def test_get_file_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="pdf", mime_type="application/pdf") with patch("services.file_service.ExtractProcessor.load_from_upload_file") as mock_extract: mock_extract.return_value = "Extracted text content" @@ -274,27 +295,17 @@ class TestFileService: # Assert assert result == "Extracted text content" - def test_get_file_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_preview_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_preview("non_existent", "tenant_id") - def test_get_file_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "exe" - mock_db_session.scalar.return_value = upload_file + def test_get_file_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="exe", mime_type="application/octet-stream") with pytest.raises(UnsupportedFileTypeError): file_service.get_file_preview("file_id", "tenant_id") - def test_get_image_preview_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "jpg" - upload_file.mime_type = "image/jpeg" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_image_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="jpg", mime_type="image/jpeg") with ( patch("services.file_service.file_helpers.verify_image_signature") as mock_verify, @@ -316,28 +327,21 @@ class TestFileService: with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_image_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_image_preview_not_found(self, file_service: FileService): with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_image_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "txt" - mock_db_session.scalar.return_value = upload_file + def test_get_image_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(UnsupportedFileTypeError): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_file_generator_by_file_id_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_file_generator_by_file_id_success(self, file_service: FileService, db_session: Session): + upload_file = self._persist_upload_file(db_session) with ( patch("services.file_service.file_helpers.verify_file_signature") as mock_verify, @@ -348,7 +352,8 @@ class TestFileService: gen, file = file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") assert list(gen) == [b"chunk"] - assert file == upload_file + assert file.id == upload_file.id + assert file.key == upload_file.key def test_get_file_generator_by_file_id_invalid_sig(self, file_service): with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify: @@ -356,20 +361,14 @@ class TestFileService: with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") - def test_get_file_generator_by_file_id_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_generator_by_file_id_not_found(self, file_service: FileService): with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") - def test_get_public_image_preview_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "png" - upload_file.mime_type = "image/png" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_public_image_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="png", mime_type="image/png") with patch("services.file_service.storage") as mock_storage: mock_storage.load.return_value = b"image content" @@ -377,66 +376,56 @@ class TestFileService: assert gen == b"image content" assert mime == "image/png" - def test_get_public_image_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_public_image_preview_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_public_image_preview("file_id") - def test_get_public_image_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "txt" - mock_db_session.scalar.return_value = upload_file + def test_get_public_image_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with pytest.raises(UnsupportedFileTypeError): file_service.get_public_image_preview("file_id") - def test_get_file_content_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_file_content_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.storage") as mock_storage: mock_storage.load.return_value = b"hello world" result = file_service.get_file_content("file_id") assert result == "hello world" - def test_get_file_content_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_content_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_content("file_id") - def test_delete_file_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - # For session.scalar(select(...)) - mock_db_session.scalar.return_value = upload_file + def test_delete_file_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.storage") as mock_storage: file_service.delete_file("file_id") mock_storage.delete.assert_called_once_with("key") - mock_db_session.delete.assert_called_once_with(upload_file) + db_session.expire_all() + assert db_session.get(UploadFile, "file_id") is None - def test_delete_file_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_delete_file_not_found(self, file_service: FileService): file_service.delete_file("file_id") # Should return without doing anything - def test_get_upload_files_by_ids_empty(self): - session = MagicMock() - result = FileService.get_upload_files_by_ids("tenant_id", [], session=session) + def test_get_upload_files_by_ids_empty(self, db_session: Session): + result = FileService.get_upload_files_by_ids("tenant_id", [], session=db_session) assert result == {} - def test_get_upload_files_by_ids(self): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "550e8400-e29b-41d4-a716-446655440000" - upload_file.tenant_id = "tenant_id" - session = MagicMock() - session.scalars().all.return_value = [upload_file] + def test_get_upload_files_by_ids(self, db_session: Session): + upload_file = self._persist_upload_file(db_session, file_id="550e8400-e29b-41d4-a716-446655440000") + self._persist_upload_file( + db_session, + file_id="550e8400-e29b-41d4-a716-446655440001", + tenant_id="other-tenant", + ) result = FileService.get_upload_files_by_ids( - "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"], session=session + "tenant_id", + ["550e8400-e29b-41d4-a716-446655440000", "550e8400-e29b-41d4-a716-446655440001"], + session=db_session, ) assert result["550e8400-e29b-41d4-a716-446655440000"] == upload_file @@ -453,10 +442,8 @@ class TestFileService: used.add("a (1).txt") assert FileService._dedupe_zip_entry_name("a.txt", used) == "a (2).txt" - def test_build_upload_files_zip_tempfile(self): - upload_file = MagicMock(spec=UploadFile) - upload_file.name = "test.txt" - upload_file.key = "key" + def test_build_upload_files_zip_tempfile(self, db_session: Session): + upload_file = self._persist_upload_file(db_session) with ( patch("services.file_service.storage") as mock_storage,