From 749d1f8d04f5fde4ab2035ffb367c6863c268adc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Tue, 8 Sep 2026 02:03:24 +0000 Subject: [PATCH] refactor(api): extract upload file delivery service (#41772) --- api/.importlinter | 14 + api/controllers/files/__init__.py | 4 +- ...age_preview.py => upload_file_delivery.py} | 146 ++++----- api/extensions/ext_application_services.py | 7 + .../upload_file_delivery_repository.py | 78 +++++ api/services/account_service.py | 6 - api/services/file_service.py | 52 --- api/services/upload_file_delivery_service.py | 115 +++++++ .../services/test_account_service.py | 28 -- .../services/test_file_service.py | 244 -------------- api/tests/unit_tests/.ruff.toml | 1 - .../controllers/files/test_image_preview.py | 303 ------------------ .../files/test_upload_file_delivery.py | 236 ++++++++++++++ .../test_ext_application_services.py | 18 ++ api/tests/unit_tests/pyrefly.toml | 2 +- .../test_upload_file_delivery_repository.py | 129 ++++++++ .../unit_tests/services/test_file_service.py | 81 ----- .../test_upload_file_delivery_service.py | 175 ++++++++++ 18 files changed, 842 insertions(+), 797 deletions(-) rename api/controllers/files/{image_preview.py => upload_file_delivery.py} (52%) create mode 100644 api/repositories/upload_file_delivery_repository.py create mode 100644 api/services/upload_file_delivery_service.py delete mode 100644 api/tests/unit_tests/controllers/files/test_image_preview.py create mode 100644 api/tests/unit_tests/controllers/files/test_upload_file_delivery.py create mode 100644 api/tests/unit_tests/repositories/test_upload_file_delivery_repository.py create mode 100644 api/tests/unit_tests/services/test_upload_file_delivery_service.py diff --git a/api/.importlinter b/api/.importlinter index 9ec0e118b03..5f0c8ea0915 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -293,6 +293,20 @@ forbidden_modules = sqlalchemy werkzeug +[importlinter:contract:upload-file-delivery-service-boundary] +name = Upload file delivery application service is framework and persistence neutral +type = forbidden +source_modules = + services.upload_file_delivery_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + [importlinter:contract:webapp-access-query-service-boundary] name = Web app access query application service is framework and persistence neutral type = forbidden diff --git a/api/controllers/files/__init__.py b/api/controllers/files/__init__.py index 42b3761b92d..657fccdbb81 100644 --- a/api/controllers/files/__init__.py +++ b/api/controllers/files/__init__.py @@ -14,7 +14,7 @@ api = ExternalApi( files_ns = Namespace("files", description="File operations", path="/") -from . import appdeploy_files, image_preview, tool_files, upload +from . import appdeploy_files, tool_files, upload, upload_file_delivery api.add_namespace(files_ns) @@ -23,7 +23,7 @@ __all__ = [ "appdeploy_files", "bp", "files_ns", - "image_preview", "tool_files", "upload", + "upload_file_delivery", ] diff --git a/api/controllers/files/image_preview.py b/api/controllers/files/upload_file_delivery.py similarity index 52% rename from api/controllers/files/image_preview.py rename to api/controllers/files/upload_file_delivery.py index e27cf79abe1..044e2c1a8fa 100644 --- a/api/controllers/files/image_preview.py +++ b/api/controllers/files/upload_file_delivery.py @@ -7,14 +7,16 @@ from flask_restx import Resource from pydantic import BaseModel, Field from werkzeug.exceptions import NotFound -import services from controllers.common.errors import UnsupportedFileTypeError from controllers.common.file_response import enforce_download_for_html -from controllers.common.schema import register_schema_models +from controllers.common.schema import query_params_from_model, register_schema_models from controllers.files import files_ns -from extensions.ext_database import db -from services.account_service import TenantService -from services.file_service import FileService +from extensions.ext_application_services import application_services +from services.errors.file import UnsupportedFileTypeError as UnsupportedFileTypeServiceError +from services.upload_file_delivery_service import ( + UploadFileDelivery, + UploadFileDeliveryNotFoundError, +) class FileSignatureQuery(BaseModel): @@ -29,6 +31,21 @@ class FilePreviewQuery(FileSignatureQuery): register_schema_models(files_ns, FileSignatureQuery, FilePreviewQuery) +_RANGE_MEDIA_TYPES = frozenset( + { + "audio/aac", + "audio/flac", + "audio/mp4", + "audio/mpeg", + "audio/ogg", + "audio/wav", + "audio/x-m4a", + "video/mp4", + "video/quicktime", + "video/webm", + } +) + def _is_svg_content(mime_type: str | None, filename: str | None, extension: str | None) -> bool: normalized_mime_type = mime_type.split(";", 1)[0].strip().lower() if mime_type else "" @@ -51,37 +68,32 @@ class ImagePreviewApi(Resource): @files_ns.doc( params={ "file_id": "ID of the file to preview", - "timestamp": "Unix timestamp used in the signature", - "nonce": "Random string used in the signature", - "sign": "HMAC signature verifying the request", + **query_params_from_model(FileSignatureQuery), } ) @files_ns.doc( responses={ 200: "Image preview returned successfully", - 400: "Missing or invalid signature parameters", + 400: "Missing or invalid query parameters", + 404: "File not found or signature is invalid", 415: "Unsupported file type", } ) - def get(self, file_id: UUID): - file_id_str = str(file_id) - + def get(self, file_id: UUID) -> Response: args = FileSignatureQuery.model_validate(request.args.to_dict(flat=True)) - timestamp = args.timestamp - nonce = args.nonce - sign = args.sign - try: - generator, mimetype = FileService(db.engine).get_image_preview( - file_id=file_id_str, - timestamp=timestamp, - nonce=nonce, - sign=sign, + delivery = application_services().upload_file_delivery.get_signed_image_preview( + file_id=str(file_id), + timestamp=args.timestamp, + nonce=args.nonce, + sign=args.sign, ) - except services.errors.file.UnsupportedFileTypeError: - raise UnsupportedFileTypeError() + except UploadFileDeliveryNotFoundError as error: + raise NotFound(str(error) or None) from error + except UnsupportedFileTypeServiceError as error: + raise UnsupportedFileTypeError() from error - return Response(generator, mimetype=mimetype) + return Response(delivery.content, mimetype=delivery.file.mime_type) @files_ns.route("//file-preview") @@ -91,60 +103,42 @@ class FilePreviewApi(Resource): @files_ns.doc( params={ "file_id": "ID of the file to preview", - "timestamp": "Unix timestamp used in the signature", - "nonce": "Random string used in the signature", - "sign": "HMAC signature verifying the request", - "as_attachment": "Whether to download the file as an attachment", + **query_params_from_model(FilePreviewQuery), } ) @files_ns.doc( responses={ 200: "File stream returned successfully", - 400: "Missing or invalid signature parameters", - 404: "File not found", - 415: "Unsupported file type", + 400: "Missing or invalid query parameters", + 404: "File not found or signature is invalid", } ) - def get(self, file_id: UUID): - file_id_str = str(file_id) - + def get(self, file_id: UUID) -> Response: args = FilePreviewQuery.model_validate(request.args.to_dict(flat=True)) try: - generator, upload_file = FileService(db.engine).get_file_generator_by_file_id( - file_id=file_id_str, + delivery = application_services().upload_file_delivery.get_signed_file_preview( + file_id=str(file_id), timestamp=args.timestamp, nonce=args.nonce, sign=args.sign, ) - except services.errors.file.UnsupportedFileTypeError: - raise UnsupportedFileTypeError() + except UploadFileDeliveryNotFoundError as error: + raise NotFound(str(error) or None) from error - response = Response( - generator, - mimetype=upload_file.mime_type, - direct_passthrough=True, - headers={}, - ) - # add Accept-Ranges header for audio/video files - if upload_file.mime_type in [ - "audio/mpeg", - "audio/wav", - "audio/mp4", - "audio/ogg", - "audio/flac", - "audio/aac", - "video/mp4", - "video/webm", - "video/quicktime", - "audio/x-m4a", - ]: + return self._build_response(delivery=delivery, as_attachment=args.as_attachment) + + @staticmethod + def _build_response(*, delivery: UploadFileDelivery, as_attachment: bool) -> Response: + file = delivery.file + response = Response(delivery.content, mimetype=file.mime_type, direct_passthrough=True, headers={}) + if file.mime_type in _RANGE_MEDIA_TYPES: response.headers["Accept-Ranges"] = "bytes" - if upload_file.size > 0: - response.headers["Content-Length"] = str(upload_file.size) - is_svg = _is_svg_content(upload_file.mime_type, upload_file.name, upload_file.extension) - if args.as_attachment or is_svg: - encoded_filename = quote(upload_file.name) + if file.size > 0: + response.headers["Content-Length"] = str(file.size) + is_svg = _is_svg_content(file.mime_type, file.name, file.extension) + if as_attachment or is_svg: + encoded_filename = quote(file.name) response.headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{encoded_filename}" response.headers["Content-Type"] = "application/octet-stream" if is_svg: @@ -152,9 +146,9 @@ class FilePreviewApi(Resource): enforce_download_for_html( response, - mime_type=upload_file.mime_type, - filename=upload_file.name, - extension=upload_file.extension, + mime_type=file.mime_type, + filename=file.name, + extension=file.extension, ) return response @@ -176,20 +170,14 @@ class WorkspaceWebappLogoApi(Resource): 415: "Unsupported file type", } ) - def get(self, workspace_id: UUID): - workspace_id_str = str(workspace_id) - - custom_config = TenantService.get_custom_config(workspace_id_str) - webapp_logo_file_id = custom_config.get("replace_webapp_logo") if custom_config is not None else None - - if not webapp_logo_file_id: - raise NotFound("webapp logo is not found") - + def get(self, workspace_id: UUID) -> Response: try: - generator, mimetype = FileService(db.engine).get_public_image_preview( - webapp_logo_file_id, + delivery = application_services().upload_file_delivery.get_workspace_webapp_logo( + workspace_id=str(workspace_id), ) - except services.errors.file.UnsupportedFileTypeError: - raise UnsupportedFileTypeError() + except UploadFileDeliveryNotFoundError as error: + raise NotFound(str(error) or None) from error + except UnsupportedFileTypeServiceError as error: + raise UnsupportedFileTypeError() from error - return Response(generator, mimetype=mimetype) + return Response(delivery.content, mimetype=delivery.file.mime_type) diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 07842055b90..338ed8b0a48 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -58,6 +58,7 @@ from repositories.step_by_step_tour_repository import SQLAlchemyStepByStepTourSt from repositories.tag_repository import TagRepository from repositories.trial_app_query_repository import TrialAppQueryRepository from repositories.trial_app_usage_repository import TrialAppUsageRepository +from repositories.upload_file_delivery_repository import UploadFileDeliveryQueryRepository from repositories.web_passport_repository import WebPassportRepository from repositories.webapp_access_query_repository import WebAppAccessQueryRepository from repositories.workflow_app_log_query_repository import WorkflowAppLogQueryRepository @@ -186,6 +187,7 @@ from services.step_by_step_tour_service import StepByStepTourService from services.system_feature_service import SystemFeatureService from services.tag_application_service import TagApplicationService from services.trial_app_usage import TrialAppUsageRecorder +from services.upload_file_delivery_service import UploadFileDeliveryService from services.web_app_runtime_query_service import WebAppRuntimeQueryService from services.web_passport_gateways import ( DeploymentWebPassportAuthGateway, @@ -266,6 +268,7 @@ class ApplicationServices: files: FileService human_input_file_uploads: HumanInputFileUploadService message_file_previews: MessageFilePreviewService + upload_file_delivery: UploadFileDeliveryService oauth_server: OAuthServerService init_validation: InitValidationService notifications: NotificationService @@ -669,6 +672,10 @@ def build_application_services( files=MessageFilePreviewQueryRepository(session_factory=database_client), storage=storage, ), + upload_file_delivery=UploadFileDeliveryService( + files=UploadFileDeliveryQueryRepository(session_factory=database_client), + storage=storage, + ), oauth_server=_build_oauth_server_service(database_client=database_client, redis=redis), init_validation=InitValidationService( state=installation_state, diff --git a/api/repositories/upload_file_delivery_repository.py b/api/repositories/upload_file_delivery_repository.py new file mode 100644 index 00000000000..13a2a3518e9 --- /dev/null +++ b/api/repositories/upload_file_delivery_repository.py @@ -0,0 +1,78 @@ +"""SQLAlchemy query adapter for public UploadFile delivery endpoints.""" + +import json +from typing import cast, override + +from sqlalchemy import Row, Select, select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Tenant, TenantCustomConfigDict +from models.model import UploadFile +from services.upload_file_delivery_service import ( + UploadFileDeliveryNotFoundError, + UploadFileDeliveryQuery, + UploadFileDeliveryRecord, +) + + +class UploadFileDeliveryQueryRepository(UploadFileDeliveryQuery): + def __init__(self, *, session_factory: sessionmaker[Session]) -> None: + self._session_factory = session_factory + + @override + def get_by_id(self, *, file_id: str) -> UploadFileDeliveryRecord | None: + with self._session_factory() as session: + row = session.execute(self._file_query().where(UploadFile.id == file_id).limit(1)).one_or_none() + + return self._to_record(row) + + @override + def get_workspace_logo(self, *, workspace_id: str) -> UploadFileDeliveryRecord | None: + with self._session_factory() as session: + workspace_row = session.execute( + select(Tenant.custom_config).where(Tenant.id == workspace_id).limit(1) + ).one_or_none() + if workspace_row is None: + raise UploadFileDeliveryNotFoundError + + custom_config = ( + cast(TenantCustomConfigDict, json.loads(workspace_row.custom_config)) + if workspace_row.custom_config + else {} + ) + logo_file_id = custom_config.get("replace_webapp_logo") + if not logo_file_id: + raise UploadFileDeliveryNotFoundError("webapp logo is not found") + + file_row = session.execute( + self._file_query() + .where( + UploadFile.id == logo_file_id, + UploadFile.tenant_id == workspace_id, + ) + .limit(1) + ).one_or_none() + + return self._to_record(file_row) + + @staticmethod + def _file_query() -> Select[tuple[str, str, int, str, str | None]]: + return select( + UploadFile.key, + UploadFile.name, + UploadFile.size, + UploadFile.extension, + UploadFile.mime_type, + ) + + @staticmethod + def _to_record(row: Row[tuple[str, str, int, str, str | None]] | None) -> UploadFileDeliveryRecord | None: + if row is None: + return None + return UploadFileDeliveryRecord( + key=row.key, + name=row.name, + size=row.size, + extension=row.extension, + mime_type=row.mime_type, + ) diff --git a/api/services/account_service.py b/api/services/account_service.py index e33af75e385..81ee3aeb037 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -1680,12 +1680,6 @@ class TenantService: target_member_join.role = new_tenant_role session.commit() - @staticmethod - def get_custom_config(tenant_id: str): - tenant = db.get_or_404(Tenant, tenant_id) - - return tenant.custom_config_dict - @staticmethod def is_owner(account: Account, tenant: Tenant, *, session: Session) -> bool: return TenantService.get_user_role(account, tenant, session=session) == TenantAccountRole.OWNER diff --git a/api/services/file_service.py b/api/services/file_service.py index 9c387720002..9fe68ab75b3 100644 --- a/api/services/file_service.py +++ b/api/services/file_service.py @@ -252,58 +252,6 @@ class FileService: text = ExtractProcessor.load_from_upload_file(upload_file, return_text=True) return text[0:PREVIEW_WORDS_LIMIT] if text else "" - def get_image_preview(self, file_id: str, timestamp: str, nonce: str, sign: str): - result = file_helpers.verify_image_signature( - upload_file_id=file_id, timestamp=timestamp, nonce=nonce, sign=sign - ) - if not result: - raise NotFound("File not found or signature is invalid") - with self._session_maker(expire_on_commit=False) as session: - upload_file = session.scalar(select(UploadFile).where(UploadFile.id == file_id).limit(1)) - - if not upload_file: - raise NotFound("File not found or signature is invalid") - - # extract text from file - extension = upload_file.extension - if extension.lower() not in IMAGE_EXTENSIONS: - raise UnsupportedFileTypeError() - - generator = storage.load(upload_file.key, stream=True) - - return generator, upload_file.mime_type - - def get_file_generator_by_file_id(self, file_id: str, timestamp: str, nonce: str, sign: str): - result = file_helpers.verify_file_signature(upload_file_id=file_id, timestamp=timestamp, nonce=nonce, sign=sign) - if not result: - raise NotFound("File not found or signature is invalid") - - with self._session_maker(expire_on_commit=False) as session: - upload_file = session.scalar(select(UploadFile).where(UploadFile.id == file_id).limit(1)) - - if not upload_file: - raise NotFound("File not found or signature is invalid") - - generator = storage.load(upload_file.key, stream=True) - - return generator, upload_file - - def get_public_image_preview(self, file_id: str): - with self._session_maker(expire_on_commit=False) as session: - upload_file = session.scalar(select(UploadFile).where(UploadFile.id == file_id).limit(1)) - - if not upload_file: - raise NotFound("File not found or signature is invalid") - - # extract text from file - extension = upload_file.extension - if extension.lower() not in IMAGE_EXTENSIONS: - raise UnsupportedFileTypeError() - - generator = storage.load(upload_file.key) - - return generator, upload_file.mime_type - def get_file_content(self, file_id: str) -> str: with self._session_maker(expire_on_commit=False) as session: upload_file: UploadFile | None = session.scalar(select(UploadFile).where(UploadFile.id == file_id).limit(1)) diff --git a/api/services/upload_file_delivery_service.py b/api/services/upload_file_delivery_service.py new file mode 100644 index 00000000000..3ae6742c2fe --- /dev/null +++ b/api/services/upload_file_delivery_service.py @@ -0,0 +1,115 @@ +"""Application service for delivering uploaded files through public file endpoints.""" + +from collections.abc import Iterator +from typing import NamedTuple, Protocol + +from constants import IMAGE_EXTENSIONS +from graphon.file import helpers as file_helpers +from services.errors.file import UnsupportedFileTypeError + + +class UploadFileDeliveryNotFoundError(LookupError): + pass + + +class UploadFileDeliveryRecord(NamedTuple): + key: str + name: str + size: int + extension: str + mime_type: str | None + + +class UploadFileDeliveryQuery(Protocol): + def get_by_id(self, *, file_id: str) -> UploadFileDeliveryRecord | None: ... + + def get_workspace_logo(self, *, workspace_id: str) -> UploadFileDeliveryRecord | None: ... + + +class UploadFileStorage(Protocol): + def load_stream(self, filename: str) -> Iterator[bytes]: ... + + def load_once(self, filename: str) -> bytes: ... + + +class UploadFileDelivery(NamedTuple): + content: bytes | Iterator[bytes] + file: UploadFileDeliveryRecord + + +class UploadFileDeliveryService: + def __init__( + self, + *, + files: UploadFileDeliveryQuery, + storage: UploadFileStorage, + ) -> None: + self._files = files + self._storage = storage + + def get_signed_image_preview( + self, + *, + file_id: str, + timestamp: str, + nonce: str, + sign: str, + ) -> UploadFileDelivery: + if not file_helpers.verify_image_signature( + upload_file_id=file_id, + timestamp=timestamp, + nonce=nonce, + sign=sign, + ): + raise UploadFileDeliveryNotFoundError("File not found or signature is invalid") + + file = self._get_file(file_id=file_id) + self._ensure_image(file=file) + return UploadFileDelivery( + content=self._storage.load_stream(file.key), + file=file, + ) + + def get_signed_file_preview( + self, + *, + file_id: str, + timestamp: str, + nonce: str, + sign: str, + ) -> UploadFileDelivery: + if not file_helpers.verify_file_signature( + upload_file_id=file_id, + timestamp=timestamp, + nonce=nonce, + sign=sign, + ): + raise UploadFileDeliveryNotFoundError("File not found or signature is invalid") + + file = self._get_file(file_id=file_id) + return UploadFileDelivery( + content=self._storage.load_stream(file.key), + file=file, + ) + + def get_workspace_webapp_logo(self, *, workspace_id: str) -> UploadFileDelivery: + file = self._files.get_workspace_logo(workspace_id=workspace_id) + if file is None: + raise UploadFileDeliveryNotFoundError("File not found or signature is invalid") + + self._ensure_image(file=file) + return UploadFileDelivery( + content=self._storage.load_once(file.key), + file=file, + ) + + def _get_file(self, *, file_id: str) -> UploadFileDeliveryRecord: + file = self._files.get_by_id(file_id=file_id) + if file is None: + raise UploadFileDeliveryNotFoundError("File not found or signature is invalid") + return file + + @staticmethod + def _ensure_image(*, file: UploadFileDeliveryRecord) -> None: + if file.extension.lower() not in IMAGE_EXTENSIONS: + raise UnsupportedFileTypeError() diff --git a/api/tests/test_containers_integration_tests/services/test_account_service.py b/api/tests/test_containers_integration_tests/services/test_account_service.py index f7953f8dbf9..9bd696f4394 100644 --- a/api/tests/test_containers_integration_tests/services/test_account_service.py +++ b/api/tests/test_containers_integration_tests/services/test_account_service.py @@ -1701,34 +1701,6 @@ class TestTenantService: assert dataset_operators[0].email == operator_email assert dataset_operators[0].role == "dataset_operator" - def test_get_custom_config_success(self, db_session_with_containers: Session, mock_external_service_dependencies): - """ - Test getting custom config successfully. - """ - fake = Faker() - tenant_name = fake.company() - theme = fake.random_element(elements=("dark", "light")) - language = fake.random_element(elements=("zh-CN", "en-US")) - # Setup mocks - mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True - - # Create tenant with custom config - tenant = TenantService.create_tenant(name=tenant_name, session=db_session_with_containers) - - # Set custom config - custom_config = {"theme": theme, "language": language, "feature_flags": {"beta": True}} - tenant.custom_config_dict = custom_config - - db_session_with_containers.commit() - - # Get custom config - retrieved_config = TenantService.get_custom_config(tenant.id) - - assert retrieved_config == custom_config - assert retrieved_config["theme"] == theme - assert retrieved_config["language"] == language - assert retrieved_config["feature_flags"]["beta"] is True - class TestRegisterService: """Integration tests for RegisterService using testcontainers.""" diff --git a/api/tests/test_containers_integration_tests/services/test_file_service.py b/api/tests/test_containers_integration_tests/services/test_file_service.py index deb0d9d7d08..63d09e2fbaa 100644 --- a/api/tests/test_containers_integration_tests/services/test_file_service.py +++ b/api/tests/test_containers_integration_tests/services/test_file_service.py @@ -38,8 +38,6 @@ class TestFileService: mock_storage.save.return_value = None mock_storage.load.return_value = BytesIO(b"mock file content") mock_file_helpers.get_signed_file_url.return_value = "https://example.com/signed-url" - mock_file_helpers.verify_image_signature.return_value = True - mock_file_helpers.verify_file_signature.return_value = True mock_extract_processor.load_from_upload_file.return_value = "extracted text content" yield { @@ -577,248 +575,6 @@ class TestFileService: assert len(result) == 3000 # PREVIEW_WORDS_LIMIT assert result == "x" * 3000 - # Test get_image_preview method - def test_get_image_preview_success( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test successful image preview generation. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Update file to have image extension - upload_file.extension = "jpg" - - db_session_with_containers.commit() - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "test_signature" - - generator, mime_type = FileService(engine).get_image_preview( - file_id=upload_file.id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - assert generator is not None - assert mime_type == upload_file.mime_type - mock_external_service_dependencies["file_helpers"].verify_image_signature.assert_called_once() - - def test_get_image_preview_invalid_signature( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test image preview with invalid signature. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Mock invalid signature - mock_external_service_dependencies["file_helpers"].verify_image_signature.return_value = False - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "invalid_signature" - - with pytest.raises(NotFound, match="File not found or signature is invalid"): - FileService(engine).get_image_preview( - file_id=upload_file.id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - def test_get_image_preview_file_not_found( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test image preview with non-existent file. - """ - fake = Faker() - non_existent_id = str(fake.uuid4()) - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "test_signature" - - with pytest.raises(NotFound, match="File not found or signature is invalid"): - FileService(engine).get_image_preview( - file_id=non_existent_id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - def test_get_image_preview_unsupported_file_type( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test image preview with non-image file type. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Update file to have non-image extension - upload_file.extension = "pdf" - - db_session_with_containers.commit() - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "test_signature" - - with pytest.raises(UnsupportedFileTypeError): - FileService(engine).get_image_preview( - file_id=upload_file.id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - # Test get_file_generator_by_file_id method - def test_get_file_generator_by_file_id_success( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test successful file generator retrieval. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "test_signature" - - generator, file_obj = FileService(engine).get_file_generator_by_file_id( - file_id=upload_file.id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - assert generator is not None - assert file_obj.id == upload_file.id - mock_external_service_dependencies["file_helpers"].verify_file_signature.assert_called_once() - - def test_get_file_generator_by_file_id_invalid_signature( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test file generator retrieval with invalid signature. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Mock invalid signature - mock_external_service_dependencies["file_helpers"].verify_file_signature.return_value = False - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "invalid_signature" - - with pytest.raises(NotFound, match="File not found or signature is invalid"): - FileService(engine).get_file_generator_by_file_id( - file_id=upload_file.id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - def test_get_file_generator_by_file_id_file_not_found( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test file generator retrieval with non-existent file. - """ - fake = Faker() - non_existent_id = str(fake.uuid4()) - - timestamp = "1234567890" - nonce = "test_nonce" - sign = "test_signature" - - with pytest.raises(NotFound, match="File not found or signature is invalid"): - FileService(engine).get_file_generator_by_file_id( - file_id=non_existent_id, - timestamp=timestamp, - nonce=nonce, - sign=sign, - ) - - # Test get_public_image_preview method - def test_get_public_image_preview_success( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test successful public image preview generation. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Update file to have image extension - upload_file.extension = "jpg" - - db_session_with_containers.commit() - - generator, mime_type = FileService(engine).get_public_image_preview(file_id=upload_file.id) - - assert generator is not None - assert mime_type == upload_file.mime_type - mock_external_service_dependencies["storage"].load.assert_called_once() - - def test_get_public_image_preview_file_not_found( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test public image preview with non-existent file. - """ - fake = Faker() - non_existent_id = str(fake.uuid4()) - - with pytest.raises(NotFound, match="File not found or signature is invalid"): - FileService(engine).get_public_image_preview(file_id=non_existent_id) - - def test_get_public_image_preview_unsupported_file_type( - self, db_session_with_containers: Session, engine, mock_external_service_dependencies - ): - """ - Test public image preview with non-image file type. - """ - fake = Faker() - account = self._create_test_account(db_session_with_containers, mock_external_service_dependencies) - upload_file = self._create_test_upload_file( - db_session_with_containers, mock_external_service_dependencies, account - ) - - # Update file to have non-image extension - upload_file.extension = "pdf" - - db_session_with_containers.commit() - - with pytest.raises(UnsupportedFileTypeError): - FileService(engine).get_public_image_preview(file_id=upload_file.id) - # Test edge cases and boundary conditions def test_upload_file_empty_content( self, db_session_with_containers: Session, engine, mock_external_service_dependencies diff --git a/api/tests/unit_tests/.ruff.toml b/api/tests/unit_tests/.ruff.toml index 32c17b823bd..1e417642d6b 100644 --- a/api/tests/unit_tests/.ruff.toml +++ b/api/tests/unit_tests/.ruff.toml @@ -61,7 +61,6 @@ extend-select = ["ANN401", "ARG"] "controllers/console/workspace/test_tool_providers.py" = ["ARG001", "ARG005"] "controllers/console/workspace/test_trigger_providers.py" = ["ARG001"] "controllers/console/workspace/test_workspace.py" = ["ARG005"] -"controllers/files/test_image_preview.py" = ["ARG005"] "controllers/files/test_tool_files.py" = ["ARG002", "ARG005"] "controllers/files/test_upload.py" = ["ARG002", "ARG005"] "controllers/inner_api/plugin/test_plugin.py" = ["ARG002"] diff --git a/api/tests/unit_tests/controllers/files/test_image_preview.py b/api/tests/unit_tests/controllers/files/test_image_preview.py deleted file mode 100644 index b1f0efbcd5c..00000000000 --- a/api/tests/unit_tests/controllers/files/test_image_preview.py +++ /dev/null @@ -1,303 +0,0 @@ -import types -from datetime import UTC, datetime -from inspect import unwrap -from unittest.mock import patch - -import pytest -from werkzeug.exceptions import NotFound - -import controllers.files.image_preview as module -from extensions.storage.storage_type import StorageType -from models.enums import CreatorUserRole -from models.model import UploadFile - - -@pytest.fixture(autouse=True) -def mock_db(): - """ - Replace Flask-SQLAlchemy db with a plain object - to avoid touching Flask app context entirely. - """ - fake_db = types.SimpleNamespace(engine=object()) - module.db = fake_db - - -def _upload_file( - *, mime_type: str = "text/plain", size: int = 10, name: str = "test.txt", extension: str = "txt" -) -> UploadFile: - upload_file = UploadFile( - tenant_id="tenant-1", - storage_type=StorageType.LOCAL, - key="uploads/file-id", - name=name, - size=size, - extension=extension, - mime_type=mime_type, - created_by_role=CreatorUserRole.ACCOUNT, - created_by="account-1", - created_at=datetime.now(UTC), - used=False, - ) - upload_file.id = "file-id" - return upload_file - - -def fake_request(args: dict): - """Return a fake request object (NOT a Flask LocalProxy).""" - return types.SimpleNamespace(args=types.SimpleNamespace(to_dict=lambda flat=True: args)) - - -class TestImagePreviewApi: - @patch.object(module, "FileService") - def test_success(self, mock_file_service): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - } - ) - - generator = iter([b"img"]) - mock_file_service.return_value.get_image_preview.return_value = ( - generator, - "image/png", - ) - - api = module.ImagePreviewApi() - get_fn = unwrap(api.get) - - response = get_fn("file-id") - - assert response.mimetype == "image/png" - - @patch.object(module, "FileService") - def test_unsupported_file_type(self, mock_file_service): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - } - ) - - mock_file_service.return_value.get_image_preview.side_effect = ( - module.services.errors.file.UnsupportedFileTypeError() - ) - - api = module.ImagePreviewApi() - get_fn = unwrap(api.get) - - with pytest.raises(module.UnsupportedFileTypeError): - get_fn("file-id") - - -class TestFilePreviewApi: - @patch.object(module, "enforce_download_for_html") - @patch.object(module, "FileService") - def test_inline_preview_uses_upload_file_mimetype(self, mock_file_service, mock_enforce): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "as_attachment": False, - } - ) - - generator = iter([b"data"]) - upload_file = _upload_file( - mime_type="application/pdf", - size=100, - name="doc.pdf", - extension="pdf", - ) - - mock_file_service.return_value.get_file_generator_by_file_id.return_value = ( - generator, - upload_file, - ) - - api = module.FilePreviewApi() - get_fn = unwrap(api.get) - - response = get_fn("file-id") - - assert response.mimetype == "application/pdf" - assert response.headers["Content-Type"] == "application/pdf" - assert response.headers["Content-Length"] == "100" - assert "Accept-Ranges" not in response.headers - mock_enforce.assert_called_once() - - @pytest.mark.parametrize( - ("mime_type", "name", "extension"), - [ - ("Image/SVG+XML; charset=UTF-8", "image.png", "png"), - ("image/png", "image.SVG", "png"), - ("image/png", "image.png", ".SVG"), - ], - ids=("mime-type", "filename", "extension"), - ) - @patch.object(module, "FileService") - def test_svg_preview_forces_download(self, mock_file_service, mime_type, name, extension): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "as_attachment": False, - } - ) - - generator = iter([b""]) - upload_file = _upload_file( - mime_type=mime_type, - size=11, - name=name, - extension=extension, - ) - - mock_file_service.return_value.get_file_generator_by_file_id.return_value = ( - generator, - upload_file, - ) - - api = module.FilePreviewApi() - get_fn = unwrap(api.get) - - response = get_fn("file-id") - - assert response.headers["Content-Disposition"].startswith("attachment") - assert response.headers["Content-Type"] == "application/octet-stream" - assert response.headers["X-Content-Type-Options"] == "nosniff" - - @patch.object(module, "FileService") - def test_html_preview_still_forces_download(self, mock_file_service): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "as_attachment": False, - } - ) - - generator = iter([b""]) - upload_file = _upload_file( - mime_type="text/html", - size=25, - name="unsafe.html", - extension="html", - ) - - mock_file_service.return_value.get_file_generator_by_file_id.return_value = ( - generator, - upload_file, - ) - - api = module.FilePreviewApi() - get_fn = unwrap(api.get) - - response = get_fn("file-id") - - assert response.headers["Content-Disposition"].startswith("attachment") - assert response.headers["Content-Type"] == "application/octet-stream" - assert response.headers["X-Content-Type-Options"] == "nosniff" - - @patch.object(module, "enforce_download_for_html") - @patch.object(module, "FileService") - def test_as_attachment(self, mock_file_service, mock_enforce): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "as_attachment": True, - } - ) - - generator = iter([b"data"]) - upload_file = _upload_file( - mime_type="application/pdf", - name="doc.pdf", - extension="pdf", - ) - - mock_file_service.return_value.get_file_generator_by_file_id.return_value = ( - generator, - upload_file, - ) - - api = module.FilePreviewApi() - get_fn = unwrap(api.get) - - response = get_fn("file-id") - - assert response.headers["Content-Disposition"].startswith("attachment") - assert response.headers["Content-Type"] == "application/octet-stream" - mock_enforce.assert_called_once() - - @patch.object(module, "FileService") - def test_unsupported_file_type(self, mock_file_service): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "as_attachment": False, - } - ) - - mock_file_service.return_value.get_file_generator_by_file_id.side_effect = ( - module.services.errors.file.UnsupportedFileTypeError() - ) - - api = module.FilePreviewApi() - get_fn = unwrap(api.get) - - with pytest.raises(module.UnsupportedFileTypeError): - get_fn("file-id") - - -class TestWorkspaceWebappLogoApi: - @patch.object(module, "FileService") - @patch.object(module.TenantService, "get_custom_config") - def test_success(self, mock_config, mock_file_service): - mock_config.return_value = {"replace_webapp_logo": "logo-id"} - generator = iter([b"logo"]) - - mock_file_service.return_value.get_public_image_preview.return_value = ( - generator, - "image/png", - ) - - api = module.WorkspaceWebappLogoApi() - get_fn = unwrap(api.get) - - response = get_fn("workspace-id") - - assert response.mimetype == "image/png" - - @patch.object(module.TenantService, "get_custom_config") - def test_logo_not_configured(self, mock_config): - mock_config.return_value = {} - - api = module.WorkspaceWebappLogoApi() - get_fn = unwrap(api.get) - - with pytest.raises(NotFound): - get_fn("workspace-id") - - @patch.object(module, "FileService") - @patch.object(module.TenantService, "get_custom_config") - def test_unsupported_file_type(self, mock_config, mock_file_service): - mock_config.return_value = {"replace_webapp_logo": "logo-id"} - mock_file_service.return_value.get_public_image_preview.side_effect = ( - module.services.errors.file.UnsupportedFileTypeError() - ) - - api = module.WorkspaceWebappLogoApi() - get_fn = unwrap(api.get) - - with pytest.raises(module.UnsupportedFileTypeError): - get_fn("workspace-id") diff --git a/api/tests/unit_tests/controllers/files/test_upload_file_delivery.py b/api/tests/unit_tests/controllers/files/test_upload_file_delivery.py new file mode 100644 index 00000000000..b390c526ac4 --- /dev/null +++ b/api/tests/unit_tests/controllers/files/test_upload_file_delivery.py @@ -0,0 +1,236 @@ +import types +from inspect import unwrap +from unittest.mock import patch + +import pytest +from werkzeug.exceptions import NotFound + +import controllers.files.upload_file_delivery as module +from services.errors.file import UnsupportedFileTypeError as UnsupportedFileTypeServiceError +from services.upload_file_delivery_service import ( + UploadFileDelivery, + UploadFileDeliveryNotFoundError, + UploadFileDeliveryRecord, +) + + +def _delivery( + *, + mime_type: str | None = "text/plain", + size: int = 10, + name: str = "test.txt", + extension: str = "txt", + content: bytes | None = None, +) -> UploadFileDelivery: + return UploadFileDelivery( + content=content if content is not None else iter([b"data"]), + file=UploadFileDeliveryRecord( + key="uploads/file-id", + name=name, + size=size, + extension=extension, + mime_type=mime_type, + ), + ) + + +def _fake_request(args: dict[str, object]): + return types.SimpleNamespace(args=types.SimpleNamespace(to_dict=lambda **_kwargs: args)) + + +class TestImagePreviewApi: + @patch.object(module, "application_services") + def test_success(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig"}) + service = mock_application_services.return_value.upload_file_delivery + service.get_signed_image_preview.return_value = _delivery(mime_type="image/png", extension="png") + + response = unwrap(module.ImagePreviewApi().get)("file-id") + + assert response.mimetype == "image/png" + service.get_signed_image_preview.assert_called_once_with( + file_id="file-id", + timestamp="123", + nonce="abc", + sign="sig", + ) + + @patch.object(module, "application_services") + def test_not_found(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig"}) + service = mock_application_services.return_value.upload_file_delivery + service.get_signed_image_preview.side_effect = UploadFileDeliveryNotFoundError( + "File not found or signature is invalid" + ) + + with pytest.raises(NotFound, match="File not found or signature is invalid"): + unwrap(module.ImagePreviewApi().get)("file-id") + + @patch.object(module, "application_services") + def test_unsupported_file_type(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig"}) + service = mock_application_services.return_value.upload_file_delivery + service.get_signed_image_preview.side_effect = UnsupportedFileTypeServiceError() + + with pytest.raises(module.UnsupportedFileTypeError): + unwrap(module.ImagePreviewApi().get)("file-id") + + +class TestFilePreviewApi: + @patch.object(module, "enforce_download_for_html") + @patch.object(module, "application_services") + def test_inline_preview_uses_file_metadata(self, mock_application_services, mock_enforce): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + service = mock_application_services.return_value.upload_file_delivery + service.get_signed_file_preview.return_value = _delivery( + mime_type="application/pdf", + size=100, + name="doc.pdf", + extension="pdf", + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert response.mimetype == "application/pdf" + assert response.headers["Content-Type"] == "application/pdf" + assert response.headers["Content-Length"] == "100" + assert "Accept-Ranges" not in response.headers + mock_enforce.assert_called_once_with( + response, + mime_type="application/pdf", + filename="doc.pdf", + extension="pdf", + ) + + @patch.object(module, "application_services") + def test_audio_preview_supports_ranges(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.return_value = _delivery( + mime_type="audio/mpeg", + extension="mp3", + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert response.headers["Accept-Ranges"] == "bytes" + + @patch.object(module, "application_services") + def test_zero_size_omits_content_length(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.return_value = _delivery( + size=0 + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert "Content-Length" not in response.headers + + @pytest.mark.parametrize( + ("mime_type", "name", "extension"), + [ + ("Image/SVG+XML; charset=UTF-8", "image.png", "png"), + ("image/png", "image.SVG", "png"), + ("image/png", "image.png", ".SVG"), + ], + ids=("mime-type", "filename", "extension"), + ) + @patch.object(module, "application_services") + def test_svg_preview_forces_download(self, mock_application_services, mime_type, name, extension): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.return_value = _delivery( + mime_type=mime_type, + size=11, + name=name, + extension=extension, + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert response.headers["Content-Disposition"].startswith("attachment") + assert response.headers["Content-Type"] == "application/octet-stream" + assert response.headers["X-Content-Type-Options"] == "nosniff" + + @patch.object(module, "application_services") + def test_html_preview_still_forces_download(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.return_value = _delivery( + mime_type="text/html", + size=25, + name="unsafe.html", + extension="html", + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert response.headers["Content-Disposition"].startswith("attachment") + assert response.headers["Content-Type"] == "application/octet-stream" + assert response.headers["X-Content-Type-Options"] == "nosniff" + + @patch.object(module, "application_services") + def test_as_attachment_encodes_filename(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": True}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.return_value = _delivery( + mime_type="application/pdf", + name="报告.pdf", + extension="pdf", + ) + + response = unwrap(module.FilePreviewApi().get)("file-id") + + assert response.headers["Content-Disposition"] == "attachment; filename*=UTF-8''%E6%8A%A5%E5%91%8A.pdf" + assert response.headers["Content-Type"] == "application/octet-stream" + + @patch.object(module, "application_services") + def test_not_found(self, mock_application_services): + module.request = _fake_request({"timestamp": "123", "nonce": "abc", "sign": "sig", "as_attachment": False}) + mock_application_services.return_value.upload_file_delivery.get_signed_file_preview.side_effect = ( + UploadFileDeliveryNotFoundError("File not found or signature is invalid") + ) + + with pytest.raises(NotFound, match="File not found or signature is invalid"): + unwrap(module.FilePreviewApi().get)("file-id") + + +class TestWorkspaceWebappLogoApi: + @patch.object(module, "application_services") + def test_success(self, mock_application_services): + service = mock_application_services.return_value.upload_file_delivery + service.get_workspace_webapp_logo.return_value = _delivery( + content=b"logo", + mime_type="image/png", + extension="png", + ) + + response = unwrap(module.WorkspaceWebappLogoApi().get)("workspace-id") + + assert response.mimetype == "image/png" + service.get_workspace_webapp_logo.assert_called_once_with(workspace_id="workspace-id") + + @patch.object(module, "application_services") + def test_logo_not_configured(self, mock_application_services): + mock_application_services.return_value.upload_file_delivery.get_workspace_webapp_logo.side_effect = ( + UploadFileDeliveryNotFoundError("webapp logo is not found") + ) + + with pytest.raises(NotFound, match="webapp logo is not found"): + unwrap(module.WorkspaceWebappLogoApi().get)("workspace-id") + + @patch.object(module, "application_services") + def test_workspace_not_found_uses_default_404(self, mock_application_services): + mock_application_services.return_value.upload_file_delivery.get_workspace_webapp_logo.side_effect = ( + UploadFileDeliveryNotFoundError() + ) + + with pytest.raises(NotFound) as error: + unwrap(module.WorkspaceWebappLogoApi().get)("workspace-id") + + assert error.value.description == NotFound.description + + @patch.object(module, "application_services") + def test_unsupported_file_type(self, mock_application_services): + mock_application_services.return_value.upload_file_delivery.get_workspace_webapp_logo.side_effect = ( + UnsupportedFileTypeServiceError() + ) + + with pytest.raises(module.UnsupportedFileTypeError): + unwrap(module.WorkspaceWebappLogoApi().get)("workspace-id") diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py index 66bcd463e02..10fb414d88e 100644 --- a/api/tests/unit_tests/extensions/test_ext_application_services.py +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -34,6 +34,7 @@ from repositories.app_tracing_config_repository import SQLAlchemyAppTracingConfi from repositories.human_input_file_upload_repository import SQLAlchemyHumanInputFileUploadRepository from repositories.message_file_preview_repository import MessageFilePreviewQueryRepository from repositories.sqlalchemy_api_workflow_run_repository import DifyAPISQLAlchemyWorkflowRunRepository +from repositories.upload_file_delivery_repository import UploadFileDeliveryQueryRepository from repositories.workflow_app_log_query_repository import WorkflowAppLogQueryRepository from repositories.workflow_run_archive_repository import WorkflowRunArchiveBundleQueryRepository from services import account_forgot_password_service, recommended_app_catalog_gateway @@ -77,6 +78,7 @@ from services.partner_tenant_binding_service import PartnerTenantBindingService from services.retention.workflow_run.archive_download_task_cache import WorkflowRunArchiveDownloadTaskCache from services.retention.workflow_run.archive_log_service import WorkflowRunArchiveService from services.tag_application_service import TagApplicationService +from services.upload_file_delivery_service import UploadFileDeliveryService from services.webapp_access_query_service import WebAppAccessUnavailableError from services.workflow_app_log_query_service import WorkflowAppLogQueryService from services.workflow_run_service import WorkflowRunService @@ -252,6 +254,22 @@ def test_build_application_services_wires_message_file_previews( assert services.message_file_previews._storage is ext_application_services.storage +def test_build_application_services_wires_upload_file_delivery( + sqlite_session_factory: sessionmaker[Session], +) -> None: + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=MagicMock(spec=RedisClientWrapper), + ) + + assert isinstance(services.upload_file_delivery, UploadFileDeliveryService) + assert isinstance(services.upload_file_delivery._files, UploadFileDeliveryQueryRepository) + assert services.upload_file_delivery._files._session_factory is sqlite_session_factory + assert services.upload_file_delivery._storage is ext_application_services.storage + + def test_build_application_services_wires_workflow_run_archives( sqlite_session_factory: sessionmaker[Session], ) -> None: diff --git a/api/tests/unit_tests/pyrefly.toml b/api/tests/unit_tests/pyrefly.toml index 42abc4e6e29..e5886a9c980 100644 --- a/api/tests/unit_tests/pyrefly.toml +++ b/api/tests/unit_tests/pyrefly.toml @@ -136,7 +136,7 @@ project-excludes = [ "controllers/console/workspace/test_snippets.py", "controllers/console/workspace/test_tool_providers.py", "controllers/console/workspace/test_workspace.py", - "controllers/files/test_image_preview.py", + "controllers/files/test_upload_file_delivery.py", "controllers/files/test_tool_files.py", "controllers/files/test_upload.py", "controllers/inner_api/app/test_dsl.py", diff --git a/api/tests/unit_tests/repositories/test_upload_file_delivery_repository.py b/api/tests/unit_tests/repositories/test_upload_file_delivery_repository.py new file mode 100644 index 00000000000..dcadbccec5c --- /dev/null +++ b/api/tests/unit_tests/repositories/test_upload_file_delivery_repository.py @@ -0,0 +1,129 @@ +from datetime import UTC, datetime + +import pytest +from sqlalchemy.orm import Session, sessionmaker + +from extensions.storage.storage_type import StorageType +from models.account import Tenant +from models.enums import CreatorUserRole +from models.model import UploadFile +from repositories.upload_file_delivery_repository import UploadFileDeliveryQueryRepository +from services.upload_file_delivery_service import UploadFileDeliveryNotFoundError, UploadFileDeliveryRecord + +WORKSPACE_ID = "11111111-1111-1111-1111-111111111111" +OTHER_WORKSPACE_ID = "22222222-2222-2222-2222-222222222222" +FILE_ID = "33333333-3333-3333-3333-333333333333" +OTHER_FILE_ID = "44444444-4444-4444-4444-444444444444" + + +def _upload_file(*, file_id: str = FILE_ID, tenant_id: str = WORKSPACE_ID) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=f"upload_files/{tenant_id}/{file_id}.png", + name="logo.png", + size=42, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="55555555-5555-5555-5555-555555555555", + created_at=datetime.now(UTC), + used=True, + ) + upload_file.id = file_id + return upload_file + + +def _workspace(*, workspace_id: str = WORKSPACE_ID, logo_file_id: str | None = None) -> Tenant: + workspace = Tenant(name=f"Workspace {workspace_id}") + workspace.id = workspace_id + if logo_file_id is not None: + workspace.custom_config_dict = {"replace_webapp_logo": logo_file_id} + return workspace + + +def _repository(session_factory: sessionmaker[Session]) -> UploadFileDeliveryQueryRepository: + return UploadFileDeliveryQueryRepository(session_factory=session_factory) + + +def test_get_by_id_returns_detached_record( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + upload_file = _upload_file() + sqlite_session.add(upload_file) + sqlite_session.commit() + + result = _repository(sqlite_session_factory).get_by_id(file_id=upload_file.id) + + assert result == UploadFileDeliveryRecord( + key=upload_file.key, + name=upload_file.name, + size=upload_file.size, + extension=upload_file.extension, + mime_type=upload_file.mime_type, + ) + + +def test_get_by_id_returns_none_when_file_does_not_exist( + sqlite_session_factory: sessionmaker[Session], +) -> None: + assert _repository(sqlite_session_factory).get_by_id(file_id=FILE_ID) is None + + +def test_get_workspace_logo_returns_workspace_owned_file( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + upload_file = _upload_file() + sqlite_session.add_all([_workspace(logo_file_id=upload_file.id), upload_file]) + sqlite_session.commit() + + result = _repository(sqlite_session_factory).get_workspace_logo(workspace_id=WORKSPACE_ID) + + assert result == UploadFileDeliveryRecord( + key=upload_file.key, + name=upload_file.name, + size=upload_file.size, + extension=upload_file.extension, + mime_type=upload_file.mime_type, + ) + + +def test_get_workspace_logo_rejects_missing_workspace( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with pytest.raises(UploadFileDeliveryNotFoundError): + _repository(sqlite_session_factory).get_workspace_logo(workspace_id=WORKSPACE_ID) + + +def test_get_workspace_logo_rejects_workspace_without_configured_logo( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + sqlite_session.add(_workspace()) + sqlite_session.commit() + + with pytest.raises(UploadFileDeliveryNotFoundError, match="webapp logo is not found"): + _repository(sqlite_session_factory).get_workspace_logo(workspace_id=WORKSPACE_ID) + + +def test_get_workspace_logo_returns_none_when_configured_file_does_not_exist( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + sqlite_session.add(_workspace(logo_file_id=FILE_ID)) + sqlite_session.commit() + + assert _repository(sqlite_session_factory).get_workspace_logo(workspace_id=WORKSPACE_ID) is None + + +def test_get_workspace_logo_rejects_file_owned_by_another_workspace( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + foreign_logo = _upload_file(file_id=OTHER_FILE_ID, tenant_id=OTHER_WORKSPACE_ID) + sqlite_session.add_all([_workspace(logo_file_id=foreign_logo.id), foreign_logo]) + sqlite_session.commit() + + assert _repository(sqlite_session_factory).get_workspace_logo(workspace_id=WORKSPACE_ID) is None diff --git a/api/tests/unit_tests/services/test_file_service.py b/api/tests/unit_tests/services/test_file_service.py index 470235cf610..6b979b24e87 100644 --- a/api/tests/unit_tests/services/test_file_service.py +++ b/api/tests/unit_tests/services/test_file_service.py @@ -398,87 +398,6 @@ class TestFileService: with pytest.raises(UnsupportedFileTypeError): file_service.get_file_preview("file_id", "tenant_id") - 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, - patch("services.file_service.storage") as mock_storage, - ): - mock_verify.return_value = True - mock_storage.load.return_value = iter([b"chunk1"]) - - # Execute - gen, mime = file_service.get_image_preview("file_id", "ts", "nonce", "sign") - - # Assert - assert list(gen) == [b"chunk1"] - assert mime == "image/jpeg" - - def test_get_image_preview_invalid_sig(self, file_service): - with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify: - mock_verify.return_value = False - 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): - 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, 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, db_session: Session): - upload_file = self._persist_upload_file(db_session) - - with ( - patch("services.file_service.file_helpers.verify_file_signature") as mock_verify, - patch("services.file_service.storage") as mock_storage, - ): - mock_verify.return_value = True - mock_storage.load.return_value = iter([b"chunk"]) - - gen, file = file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") - assert list(gen) == [b"chunk"] - 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: - mock_verify.return_value = False - 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): - 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, 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" - gen, mime = file_service.get_public_image_preview("file_id") - assert gen == b"image content" - assert mime == "image/png" - - 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, 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, db_session: Session): self._persist_upload_file(db_session) diff --git a/api/tests/unit_tests/services/test_upload_file_delivery_service.py b/api/tests/unit_tests/services/test_upload_file_delivery_service.py new file mode 100644 index 00000000000..b7ec99e508a --- /dev/null +++ b/api/tests/unit_tests/services/test_upload_file_delivery_service.py @@ -0,0 +1,175 @@ +from unittest.mock import Mock, patch + +import pytest + +from services.errors.file import UnsupportedFileTypeError +from services.upload_file_delivery_service import ( + UploadFileDelivery, + UploadFileDeliveryNotFoundError, + UploadFileDeliveryQuery, + UploadFileDeliveryRecord, + UploadFileDeliveryService, + UploadFileStorage, +) + + +def _record(*, extension: str = "png") -> UploadFileDeliveryRecord: + return UploadFileDeliveryRecord( + key="upload_files/tenant-id/file-id", + name=f"file.{extension}", + size=7, + extension=extension, + mime_type="image/png" if extension == "png" else "text/plain", + ) + + +@pytest.fixture +def files() -> Mock: + return Mock(spec=UploadFileDeliveryQuery) + + +@pytest.fixture +def storage() -> Mock: + return Mock(spec=UploadFileStorage) + + +@pytest.fixture +def service(files: Mock, storage: Mock) -> UploadFileDeliveryService: + return UploadFileDeliveryService(files=files, storage=storage) + + +def test_invalid_image_signature_does_not_query_or_load_file( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + with patch("services.upload_file_delivery_service.file_helpers.verify_image_signature", return_value=False): + with pytest.raises(UploadFileDeliveryNotFoundError): + service.get_signed_image_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="invalid") + + files.get_by_id.assert_not_called() + storage.load_stream.assert_not_called() + storage.load_once.assert_not_called() + + +def test_invalid_file_signature_does_not_query_or_load_file( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + with patch("services.upload_file_delivery_service.file_helpers.verify_file_signature", return_value=False): + with pytest.raises(UploadFileDeliveryNotFoundError): + service.get_signed_file_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="invalid") + + files.get_by_id.assert_not_called() + storage.load_stream.assert_not_called() + storage.load_once.assert_not_called() + + +def test_signed_image_preview_loads_image_stream( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + file = _record() + content = iter((b"content",)) + files.get_by_id.return_value = file + storage.load_stream.return_value = content + + with patch("services.upload_file_delivery_service.file_helpers.verify_image_signature", return_value=True): + result = service.get_signed_image_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="valid") + + assert result == UploadFileDelivery(content=content, file=file) + files.get_by_id.assert_called_once_with(file_id="file-id") + storage.load_stream.assert_called_once_with(file.key) + + +def test_signed_image_preview_rejects_non_image_before_loading( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + files.get_by_id.return_value = _record(extension="txt") + + with patch("services.upload_file_delivery_service.file_helpers.verify_image_signature", return_value=True): + with pytest.raises(UnsupportedFileTypeError): + service.get_signed_image_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="valid") + + storage.load_stream.assert_not_called() + + +def test_signed_file_preview_allows_non_image( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + file = _record(extension="txt") + content = iter((b"content",)) + files.get_by_id.return_value = file + storage.load_stream.return_value = content + + with patch("services.upload_file_delivery_service.file_helpers.verify_file_signature", return_value=True): + result = service.get_signed_file_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="valid") + + assert result == UploadFileDelivery(content=content, file=file) + storage.load_stream.assert_called_once_with(file.key) + + +def test_signed_file_preview_reports_missing_file( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + files.get_by_id.return_value = None + + with patch("services.upload_file_delivery_service.file_helpers.verify_file_signature", return_value=True): + with pytest.raises(UploadFileDeliveryNotFoundError): + service.get_signed_file_preview(file_id="missing", timestamp="1", nonce="nonce", sign="valid") + + storage.load_stream.assert_not_called() + + +def test_workspace_logo_loads_content_once( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + file = _record() + files.get_workspace_logo.return_value = file + storage.load_once.return_value = b"content" + + result = service.get_workspace_webapp_logo(workspace_id="workspace-id") + + assert result == UploadFileDelivery(content=b"content", file=file) + files.get_workspace_logo.assert_called_once_with(workspace_id="workspace-id") + storage.load_once.assert_called_once_with(file.key) + storage.load_stream.assert_not_called() + + +def test_workspace_logo_reports_missing_file( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + files.get_workspace_logo.return_value = None + + with pytest.raises(UploadFileDeliveryNotFoundError): + service.get_workspace_webapp_logo(workspace_id="workspace-id") + + storage.load_once.assert_not_called() + + +def test_storage_error_is_not_converted( + service: UploadFileDeliveryService, + files: Mock, + storage: Mock, +) -> None: + storage_error = OSError("storage unavailable") + files.get_by_id.return_value = _record(extension="txt") + storage.load_stream.side_effect = storage_error + + with patch("services.upload_file_delivery_service.file_helpers.verify_file_signature", return_value=True): + with pytest.raises(OSError) as error_info: + service.get_signed_file_preview(file_id="file-id", timestamp="1", nonce="nonce", sign="valid") + + assert error_info.value is storage_error