diff --git a/api/.importlinter b/api/.importlinter index 0c91b3c82f8..d9796d56504 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -484,6 +484,20 @@ forbidden_modules = sqlalchemy werkzeug +[importlinter:contract:plugin-file-upload-service-boundary] +name = Plugin file upload application service is framework and persistence neutral +type = forbidden +source_modules = + services.plugin_file_upload_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + [importlinter:contract:account-activation-service-boundary] name = Account activation application service is framework and persistence neutral type = forbidden diff --git a/api/controllers/files/__init__.py b/api/controllers/files/__init__.py index 657fccdbb81..6b4b409d33c 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, tool_files, upload, upload_file_delivery +from . import appdeploy_files, plugin_file_upload, tool_files, upload_file_delivery api.add_namespace(files_ns) @@ -23,7 +23,7 @@ __all__ = [ "appdeploy_files", "bp", "files_ns", + "plugin_file_upload", "tool_files", - "upload", "upload_file_delivery", ] diff --git a/api/controllers/files/plugin_file_upload.py b/api/controllers/files/plugin_file_upload.py new file mode 100644 index 00000000000..6ef5ba3d3f2 --- /dev/null +++ b/api/controllers/files/plugin_file_upload.py @@ -0,0 +1,113 @@ +"""Signed plugin file upload endpoint.""" + +from typing import Literal + +from flask import request +from flask_restx import Resource +from flask_restx.api import HTTPStatus +from pydantic import BaseModel, Field + +from controllers.common.errors import ( + FilenameNotExistsError, + FileTooLargeError, + NoFileUploadedError, + UnsupportedFileTypeError, +) +from controllers.common.schema import ( + JsonResponseWithStatus, + query_params_from_model, + register_response_schema_models, + register_schema_models, +) +from controllers.console.wraps import setup_required +from controllers.files import files_ns +from extensions.ext_application_services import application_services +from fields.file_fields import FileResponse +from libs.exception import BaseHTTPException +from libs.helper import dump_response +from services.errors.file import FileTooLargeError as ServiceFileTooLargeError +from services.errors.file import UnsupportedFileTypeError as ServiceUnsupportedFileTypeError +from services.plugin_file_upload_service import PluginFileUploadAccessDeniedError + + +class PluginUploadQuery(BaseModel): + timestamp: str = Field(..., description="Unix timestamp for signature verification") + nonce: str = Field(..., description="Random nonce for signature verification") + sign: str = Field(..., description="HMAC signature") + tenant_id: str = Field(..., description="Tenant identifier") + user_id: str = Field(..., description="User identifier") + user_from: Literal["account", "end-user"] | None = Field(default=None, description="User identity type") + conversation_id: str | None = Field(default=None, description="Conversation identifier") + max_size: int | None = Field(default=None, ge=0, description="Signed maximum file size in bytes") + + +class InvalidPluginFileUploadError(BaseHTTPException): + error_code = "invalid_plugin_file_upload" + description = "The plugin file upload request is invalid or expired." + code = HTTPStatus.FORBIDDEN + + +_PLUGIN_UPLOAD_PARAMS = { + **query_params_from_model(PluginUploadQuery), + "file": { + "description": "File to upload for plugin usage.", + "in": "formData", + "type": "file", + "required": True, + }, +} + + +register_schema_models(files_ns, PluginUploadQuery) +register_response_schema_models(files_ns, FileResponse) + + +@files_ns.route("/upload/for-plugin") +class PluginUploadFileApi(Resource): + @setup_required + @files_ns.doc("upload_plugin_file") + @files_ns.doc( + description="Upload a file for plugin usage with signature verification", + consumes=["multipart/form-data"], + params=_PLUGIN_UPLOAD_PARAMS, + responses={ + 201: "File uploaded successfully", + 400: "Invalid query parameters, no file was uploaded, or the file has no name", + 403: "The signed upload request is invalid or expired", + 413: "File too large", + 415: "Unsupported file type", + }, + ) + @files_ns.response(HTTPStatus.CREATED, "File uploaded", files_ns.models[FileResponse.__name__]) + def post(self) -> JsonResponseWithStatus: + args = PluginUploadQuery.model_validate(request.args.to_dict(flat=True)) + file = request.files.get("file") + if file is None: + raise NoFileUploadedError() + if not file.filename: + raise FilenameNotExistsError() + if not file.mimetype: + raise UnsupportedFileTypeError() + + try: + result = application_services().plugin_file_uploads.upload( + stream=file.stream, + filename=file.filename, + mimetype=file.mimetype, + tenant_id=args.tenant_id, + user_id=args.user_id, + user_from=args.user_from, + conversation_id=args.conversation_id, + timestamp=args.timestamp, + nonce=args.nonce, + sign=args.sign, + max_size=args.max_size, + ) + except PluginFileUploadAccessDeniedError as error: + raise InvalidPluginFileUploadError() from error + except ServiceFileTooLargeError as error: + raise FileTooLargeError(error.description) from error + except ServiceUnsupportedFileTypeError as error: + raise UnsupportedFileTypeError() from error + + return dump_response(FileResponse, result), HTTPStatus.CREATED diff --git a/api/controllers/files/upload.py b/api/controllers/files/upload.py deleted file mode 100644 index 1da0d720be3..00000000000 --- a/api/controllers/files/upload.py +++ /dev/null @@ -1,162 +0,0 @@ -from typing import Literal - -from flask import request -from flask_restx import Resource -from flask_restx.api import HTTPStatus -from pydantic import BaseModel, Field -from werkzeug.exceptions import Forbidden - -import services -from core.db.session_factory import session_factory -from core.tools.signature import sign_tool_file, verify_plugin_file_signature -from core.tools.tool_file_manager import ToolFileManager, resolve_extension -from core.workflow.file_reference import build_file_reference -from fields.file_fields import FileResponse -from services.account_service import TenantService - -from ..common.errors import ( - FileTooLargeError, - UnsupportedFileTypeError, -) -from ..common.schema import register_schema_models -from ..console.wraps import setup_required -from ..files import files_ns -from ..inner_api.plugin.wraps import get_user - - -class PluginUploadQuery(BaseModel): - timestamp: str = Field(..., description="Unix timestamp for signature verification") - nonce: str = Field(..., description="Random nonce for signature verification") - sign: str = Field(..., description="HMAC signature") - tenant_id: str = Field(..., description="Tenant identifier") - user_id: str | None = Field(default=None, description="User identifier") - user_from: Literal["account", "end-user"] | None = Field(default=None, description="User identity type") - conversation_id: str | None = Field(default=None, description="Conversation identifier") - max_size: int | None = Field(default=None, ge=0, description="Signed maximum file size in bytes") - - -register_schema_models(files_ns, PluginUploadQuery) - - -register_schema_models(files_ns, FileResponse) - - -@files_ns.route("/upload/for-plugin") -class PluginUploadFileApi(Resource): - @setup_required - @files_ns.expect(files_ns.models[PluginUploadQuery.__name__]) - @files_ns.doc("upload_plugin_file") - @files_ns.doc(description="Upload a file for plugin usage with signature verification") - @files_ns.doc( - responses={ - 201: "File uploaded successfully", - 400: "Invalid request parameters", - 403: "Forbidden - Invalid signature or missing parameters", - 413: "File too large", - 415: "Unsupported file type", - } - ) - @files_ns.response(HTTPStatus.CREATED, "File uploaded", files_ns.models[FileResponse.__name__]) - def post(self): - """Upload a file for plugin usage. - - Accepts a file upload with signature verification for security. - The file must be accompanied by valid timestamp, nonce, and signature parameters. - - Returns: - dict: File metadata including ID, canonical ``reference`` for - output-file reconstruction, URLs, and properties - int: HTTP status code (201 for success) - - Raises: - Forbidden: Invalid signature or missing required parameters - FileTooLargeError: File exceeds size limit - UnsupportedFileTypeError: File type not supported - """ - args = PluginUploadQuery.model_validate(request.args.to_dict(flat=True)) - - file = request.files.get("file") - if file is None: - raise Forbidden("File is required.") - - timestamp = args.timestamp - nonce = args.nonce - sign = args.sign - tenant_id = args.tenant_id - if args.user_from == "account": - if args.user_id is None: - raise Forbidden("Invalid request.") - with session_factory.create_session() as session: - is_tenant_member = TenantService.account_belongs_to_tenant( - args.user_id, - tenant_id, - session=session, - ) - if not is_tenant_member: - raise Forbidden("Invalid request.") - owner_id = args.user_id - else: - owner_id = get_user(tenant_id, args.user_id).id - - filename = file.filename - mimetype = file.mimetype - - if not filename or not mimetype: - raise Forbidden("Invalid request.") - - if not verify_plugin_file_signature( - filename=filename, - mimetype=mimetype, - tenant_id=tenant_id, - user_id=owner_id, - conversation_id=args.conversation_id, - user_from=args.user_from, - timestamp=timestamp, - nonce=nonce, - sign=sign, - max_size=args.max_size, - ): - raise Forbidden("Invalid request.") - - try: - if args.max_size is None: - file_binary = file.stream.read() - else: - file_binary = file.stream.read(args.max_size + 1) - if len(file_binary) > args.max_size: - raise FileTooLargeError("File size exceeds the signed upload limit.") - - tool_file = ToolFileManager().create_file_by_raw( - user_id=owner_id, - tenant_id=tenant_id, - file_binary=file_binary, - mimetype=mimetype, - filename=filename, - conversation_id=args.conversation_id, - ) - - extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype) - preview_url = sign_tool_file(tool_file_id=tool_file.id, extension=extension, for_external=True) - - # Create a dictionary with all the necessary attributes - result = FileResponse( - id=tool_file.id, - reference=build_file_reference(record_id=tool_file.id), - name=tool_file.name, - size=tool_file.size, - extension=extension, - mime_type=mimetype, - preview_url=preview_url, - source_url=tool_file.original_url, - original_url=tool_file.original_url, - user_id=tool_file.user_id, - tenant_id=tool_file.tenant_id, - conversation_id=tool_file.conversation_id, - file_key=tool_file.file_key, - ) - - return result.model_dump(mode="json"), 201 - except services.errors.file.FileTooLargeError as file_too_large_error: - raise FileTooLargeError(file_too_large_error.description) - except services.errors.file.UnsupportedFileTypeError: - raise UnsupportedFileTypeError() diff --git a/api/core/tools/signature.py b/api/core/tools/signature.py index fc5a7642ac5..f6954d67e5a 100644 --- a/api/core/tools/signature.py +++ b/api/core/tools/signature.py @@ -164,8 +164,13 @@ def verify_plugin_file_signature( if sign != recalculated_encoded_sign: return False + try: + signed_at = int(timestamp) + except ValueError: + return False + current_time = int(time.time()) - return current_time - int(timestamp) <= dify_config.FILES_ACCESS_TIMEOUT + return current_time - signed_at <= dify_config.FILES_ACCESS_TIMEOUT def _plugin_upload_signature_payload( diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 666313a2a30..07aa6edc21e 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -52,6 +52,7 @@ from repositories.installation_state_repository import InstallationStateReposito from repositories.message_file_preview_repository import MessageFilePreviewQueryRepository from repositories.oauth_access_token_repository import SQLAlchemyOAuthAccessTokenRepository from repositories.oauth_server_repository import RedisOAuthServerTokenRepository, SQLAlchemyOAuthServerRepository +from repositories.plugin_file_upload_repository import SQLAlchemyPluginFileUploadOwnerRepository from repositories.recommended_app_catalog_repository import DatabaseRecommendedAppCatalogRepository from repositories.sqlalchemy_api_workflow_run_repository import DifyAPISQLAlchemyWorkflowRunRepository from repositories.step_by_step_tour_repository import SQLAlchemyStepByStepTourStateRepository @@ -167,6 +168,8 @@ from services.notification_service import NotificationService from services.notion_data_source_gateway import NotionDataSourceGateway from services.oauth_server_service import OAUTH_ACCESS_TOKEN_EXPIRES_IN, OAuthServerService from services.partner_tenant_binding_service import PartnerTenantBindingService +from services.plugin_file_upload_gateway import ToolFilePluginUploadGateway +from services.plugin_file_upload_service import PluginFileUploadService from services.recommended_app_catalog_gateway import ( BuiltinRecommendedAppCatalogGateway, RecommendedAppCatalogRouter, @@ -269,6 +272,7 @@ class ApplicationServices: files: FileService human_input_file_uploads: HumanInputFileUploadService message_file_previews: MessageFilePreviewService + plugin_file_uploads: PluginFileUploadService tool_file_downloads: ToolFileDownloadService upload_file_delivery: UploadFileDeliveryService oauth_server: OAuthServerService @@ -674,6 +678,10 @@ def build_application_services( files=MessageFilePreviewQueryRepository(session_factory=database_client), storage=storage, ), + plugin_file_uploads=PluginFileUploadService( + owners=SQLAlchemyPluginFileUploadOwnerRepository(session_factory=database_client), + files=ToolFilePluginUploadGateway(tool_files=ToolFileManager()), + ), tool_file_downloads=ToolFileDownloadService(tool_files=ToolFileManager()), upload_file_delivery=UploadFileDeliveryService( files=UploadFileDeliveryQueryRepository(session_factory=database_client), diff --git a/api/repositories/plugin_file_upload_repository.py b/api/repositories/plugin_file_upload_repository.py new file mode 100644 index 00000000000..2b39aa8af89 --- /dev/null +++ b/api/repositories/plugin_file_upload_repository.py @@ -0,0 +1,45 @@ +"""Persistence queries for signed plugin file upload owners.""" + +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import TenantAccountJoin +from models.model import EndUser +from services.plugin_file_upload_service import PluginFileUploadOwnerQuery, PluginUploadUserFrom + + +class SQLAlchemyPluginFileUploadOwnerRepository(PluginFileUploadOwnerQuery): + def __init__(self, *, session_factory: sessionmaker[Session]) -> None: + self._session_factory = session_factory + + @override + def owner_exists( + self, + *, + tenant_id: str, + user_id: str, + user_from: PluginUploadUserFrom, + ) -> bool: + if user_from == "account": + statement = ( + select(TenantAccountJoin.id) + .where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == user_id, + ) + .limit(1) + ) + else: + statement = ( + select(EndUser.id) + .where( + EndUser.tenant_id == tenant_id, + EndUser.id == user_id, + ) + .limit(1) + ) + + with self._session_factory() as session: + return session.scalar(statement) is not None diff --git a/api/services/plugin_file_upload_gateway.py b/api/services/plugin_file_upload_gateway.py new file mode 100644 index 00000000000..4cfa13747aa --- /dev/null +++ b/api/services/plugin_file_upload_gateway.py @@ -0,0 +1,54 @@ +"""ToolFile adapter for signed plugin uploads.""" + +from typing import override + +from core.tools.signature import sign_tool_file +from core.tools.tool_file_manager import ToolFileManager, resolve_extension +from core.workflow.file_reference import build_file_reference +from services.plugin_file_upload_service import PluginFileUploadFiles, PluginFileUploadResult + + +class ToolFilePluginUploadGateway(PluginFileUploadFiles): + def __init__(self, *, tool_files: ToolFileManager) -> None: + self._tool_files = tool_files + + @override + def store( + self, + *, + user_id: str, + tenant_id: str, + conversation_id: str | None, + content: bytes, + mimetype: str, + filename: str, + ) -> PluginFileUploadResult: + tool_file = self._tool_files.create_file_by_raw( + user_id=user_id, + tenant_id=tenant_id, + conversation_id=conversation_id, + file_binary=content, + mimetype=mimetype, + filename=filename, + ) + extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype) + + return PluginFileUploadResult( + id=tool_file.id, + reference=build_file_reference(record_id=tool_file.id), + name=tool_file.name, + size=tool_file.size, + extension=extension, + mime_type=mimetype, + preview_url=sign_tool_file( + tool_file_id=tool_file.id, + extension=extension, + for_external=True, + ), + source_url=tool_file.original_url, + original_url=tool_file.original_url, + user_id=tool_file.user_id, + tenant_id=tool_file.tenant_id, + conversation_id=tool_file.conversation_id, + file_key=tool_file.file_key, + ) diff --git a/api/services/plugin_file_upload_service.py b/api/services/plugin_file_upload_service.py new file mode 100644 index 00000000000..67470f737e4 --- /dev/null +++ b/api/services/plugin_file_upload_service.py @@ -0,0 +1,116 @@ +"""Application service for signed plugin file uploads.""" + +from dataclasses import dataclass +from typing import IO, Literal, Protocol + +from core.tools.signature import verify_plugin_file_signature +from services.errors.file import FileTooLargeError + +PluginUploadUserFrom = Literal["account", "end-user"] | None + + +class PluginFileUploadAccessDeniedError(PermissionError): + pass + + +@dataclass(frozen=True, slots=True) +class PluginFileUploadResult: + id: str + reference: str + name: str + size: int + extension: str + mime_type: str + preview_url: str + source_url: str | None + original_url: str | None + user_id: str + tenant_id: str + conversation_id: str | None + file_key: str + + +class PluginFileUploadOwnerQuery(Protocol): + def owner_exists( + self, + *, + tenant_id: str, + user_id: str, + user_from: PluginUploadUserFrom, + ) -> bool: ... + + +class PluginFileUploadFiles(Protocol): + def store( + self, + *, + user_id: str, + tenant_id: str, + conversation_id: str | None, + content: bytes, + mimetype: str, + filename: str, + ) -> PluginFileUploadResult: ... + + +class PluginFileUploadService: + def __init__( + self, + *, + owners: PluginFileUploadOwnerQuery, + files: PluginFileUploadFiles, + ) -> None: + self._owners = owners + self._files = files + + def upload( + self, + *, + stream: IO[bytes], + filename: str, + mimetype: str, + tenant_id: str, + user_id: str, + user_from: PluginUploadUserFrom, + conversation_id: str | None, + timestamp: str, + nonce: str, + sign: str, + max_size: int | None, + ) -> PluginFileUploadResult: + if not verify_plugin_file_signature( + filename=filename, + mimetype=mimetype, + tenant_id=tenant_id, + user_id=user_id, + user_from=user_from, + conversation_id=conversation_id, + timestamp=timestamp, + nonce=nonce, + sign=sign, + max_size=max_size, + ): + raise PluginFileUploadAccessDeniedError + + if not self._owners.owner_exists( + tenant_id=tenant_id, + user_id=user_id, + user_from=user_from, + ): + raise PluginFileUploadAccessDeniedError + + if max_size is None: + content = stream.read() + else: + content = stream.read(max_size + 1) + if len(content) > max_size: + raise FileTooLargeError("File size exceeds the signed upload limit.") + + return self._files.store( + user_id=user_id, + tenant_id=tenant_id, + conversation_id=conversation_id, + content=content, + mimetype=mimetype, + filename=filename, + ) diff --git a/api/tests/unit_tests/.ruff.toml b/api/tests/unit_tests/.ruff.toml index 1e417642d6b..09056d4e952 100644 --- a/api/tests/unit_tests/.ruff.toml +++ b/api/tests/unit_tests/.ruff.toml @@ -62,7 +62,6 @@ extend-select = ["ANN401", "ARG"] "controllers/console/workspace/test_trigger_providers.py" = ["ARG001"] "controllers/console/workspace/test_workspace.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"] "controllers/inner_api/plugin/test_plugin_wraps.py" = ["ARG001", "ARG002", "ARG003", "TID251"] "controllers/inner_api/test_runtime_credentials.py" = ["ARG001"] diff --git a/api/tests/unit_tests/controllers/files/test_plugin_file_upload.py b/api/tests/unit_tests/controllers/files/test_plugin_file_upload.py new file mode 100644 index 00000000000..704afa30878 --- /dev/null +++ b/api/tests/unit_tests/controllers/files/test_plugin_file_upload.py @@ -0,0 +1,332 @@ +import io +import types +from collections.abc import Callable +from inspect import unwrap +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from pydantic import ValidationError + +import controllers.files.plugin_file_upload as module +from controllers.files import bp as files_blueprint +from core.workflow.file_reference import build_file_reference +from enums import DeploymentEdition +from services.errors.file import FileTooLargeError as ServiceFileTooLargeError +from services.errors.file import UnsupportedFileTypeError as ServiceUnsupportedFileTypeError +from services.plugin_file_upload_service import PluginFileUploadAccessDeniedError, PluginFileUploadResult + + +class DummyFile: + def __init__( + self, + *, + filename: str | None = "report.pdf", + mimetype: str | None = "application/pdf", + content: bytes = b"content", + ) -> None: + self.filename = filename + self.mimetype = mimetype + self.stream = io.BytesIO(content) + + +def _fake_request(args: dict[str, object], *, file: DummyFile | None = None) -> types.SimpleNamespace: + return types.SimpleNamespace( + args=types.SimpleNamespace(to_dict=lambda **_kwargs: args), + files={"file": file} if file is not None else {}, + ) + + +def _result() -> PluginFileUploadResult: + return PluginFileUploadResult( + id="file-id", + reference=build_file_reference(record_id="file-id"), + name="report.pdf", + size=7, + extension=".pdf", + mime_type="application/pdf", + preview_url="https://files.example.com/files/tools/file-id.pdf?signed", + source_url=None, + original_url=None, + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="tools/tenant-id/file.pdf", + ) + + +def _valid_args() -> dict[str, object]: + return { + "timestamp": "123", + "nonce": "nonce", + "sign": "signature", + "tenant_id": "tenant-id", + "user_id": "user-id", + "user_from": "end-user", + "conversation_id": "conversation-id", + "max_size": "1024", + } + + +@pytest.fixture +def files_app(config_overrides: Callable[..., None]) -> Flask: + config_overrides(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) + app = Flask(__name__) + app.config["TESTING"] = True + app.register_blueprint(files_blueprint) + return app + + +class TestPluginUploadFileApi: + def test_upload_query_requires_the_signed_user_id(self) -> None: + args = _valid_args() + del args["user_id"] + + with pytest.raises(ValidationError): + module.PluginUploadQuery.model_validate(args) + + @patch.object(module, "application_services") + def test_upload_returns_the_existing_plugin_file_contract( + self, + application_services: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + file = DummyFile() + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=file)) + service = application_services.return_value.plugin_file_uploads + service.upload.return_value = _result() + + response, status = unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + assert status == 201 + assert response == { + "id": "file-id", + "reference": build_file_reference(record_id="file-id"), + "name": "report.pdf", + "size": 7, + "extension": ".pdf", + "mime_type": "application/pdf", + "created_by": None, + "created_at": None, + "preview_url": "https://files.example.com/files/tools/file-id.pdf?signed", + "source_url": None, + "original_url": None, + "user_id": "user-id", + "tenant_id": "tenant-id", + "conversation_id": "conversation-id", + "file_key": "tools/tenant-id/file.pdf", + } + service.upload.assert_called_once_with( + stream=file.stream, + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-id", + user_id="user-id", + user_from="end-user", + conversation_id="conversation-id", + timestamp="123", + nonce="nonce", + sign="signature", + max_size=1024, + ) + + def test_missing_file_has_a_specific_client_error(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(module, "request", _fake_request(_valid_args())) + + with pytest.raises(module.NoFileUploadedError): + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + @pytest.mark.parametrize( + ("file", "expected_error"), + [ + pytest.param(DummyFile(filename=""), module.FilenameNotExistsError, id="filename"), + pytest.param(DummyFile(mimetype=""), module.UnsupportedFileTypeError, id="mimetype"), + ], + ) + def test_invalid_file_metadata_has_a_specific_client_error( + self, + monkeypatch: pytest.MonkeyPatch, + file: DummyFile, + expected_error: type[Exception], + ) -> None: + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=file)) + + with pytest.raises(expected_error): + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + @patch.object(module, "application_services") + def test_access_denied_is_reported_without_leaking_identity_details( + self, + application_services: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=DummyFile())) + application_services.return_value.plugin_file_uploads.upload.side_effect = PluginFileUploadAccessDeniedError() + + with pytest.raises(module.InvalidPluginFileUploadError) as error_info: + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + assert error_info.value.code == 403 + assert error_info.value.error_code == "invalid_plugin_file_upload" + + @patch.object(module, "application_services") + def test_signed_size_limit_is_reported_as_413( + self, + application_services: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=DummyFile())) + application_services.return_value.plugin_file_uploads.upload.side_effect = ServiceFileTooLargeError( + "signed limit exceeded" + ) + + with pytest.raises(module.FileTooLargeError) as error_info: + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + assert error_info.value.code == 413 + assert error_info.value.__cause__ is not None + + @patch.object(module, "application_services") + def test_unsupported_file_type_is_reported_as_415( + self, + application_services: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=DummyFile())) + application_services.return_value.plugin_file_uploads.upload.side_effect = ServiceUnsupportedFileTypeError() + + with pytest.raises(module.UnsupportedFileTypeError) as error_info: + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + assert error_info.value.code == 415 + assert error_info.value.__cause__ is not None + + @patch.object(module, "application_services") + def test_unexpected_failure_is_not_relabelled_as_a_client_error( + self, + application_services: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + storage_error = OSError("storage unavailable") + monkeypatch.setattr(module, "request", _fake_request(_valid_args(), file=DummyFile())) + application_services.return_value.plugin_file_uploads.upload.side_effect = storage_error + + with pytest.raises(OSError) as error_info: + unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) + + assert error_info.value is storage_error + + +class TestPluginUploadFileHttpContract: + @patch.object(module, "application_services") + def test_multipart_upload_returns_strict_201_and_consumer_fields( + self, + application_services: MagicMock, + files_app: Flask, + ) -> None: + application_services.return_value.plugin_file_uploads.upload.return_value = _result() + + response = files_app.test_client().post( + "/files/upload/for-plugin", + query_string=_valid_args(), + data={"file": (io.BytesIO(b"content"), "report.pdf")}, + content_type="multipart/form-data", + ) + + assert response.status_code == 201 + assert response.get_json() == { + "id": "file-id", + "reference": build_file_reference(record_id="file-id"), + "name": "report.pdf", + "size": 7, + "extension": ".pdf", + "mime_type": "application/pdf", + "created_by": None, + "created_at": None, + "preview_url": "https://files.example.com/files/tools/file-id.pdf?signed", + "source_url": None, + "original_url": None, + "user_id": "user-id", + "tenant_id": "tenant-id", + "conversation_id": "conversation-id", + "file_key": "tools/tenant-id/file.pdf", + } + + @patch.object(module, "application_services") + def test_invalid_signature_returns_a_structured_403( + self, + application_services: MagicMock, + files_app: Flask, + ) -> None: + application_services.return_value.plugin_file_uploads.upload.side_effect = PluginFileUploadAccessDeniedError() + + response = files_app.test_client().post( + "/files/upload/for-plugin", + query_string=_valid_args(), + data={"file": (io.BytesIO(b"content"), "report.pdf")}, + content_type="multipart/form-data", + ) + + assert response.status_code == 403 + assert response.get_json() == { + "code": "invalid_plugin_file_upload", + "message": "The plugin file upload request is invalid or expired.", + "status": 403, + } + + @patch.object(module, "application_services") + def test_signed_size_limit_returns_a_structured_413( + self, + application_services: MagicMock, + files_app: Flask, + ) -> None: + application_services.return_value.plugin_file_uploads.upload.side_effect = ServiceFileTooLargeError( + "signed limit exceeded" + ) + + response = files_app.test_client().post( + "/files/upload/for-plugin", + query_string=_valid_args(), + data={"file": (io.BytesIO(b"content"), "report.pdf")}, + content_type="multipart/form-data", + ) + + assert response.status_code == 413 + assert response.get_json() == { + "code": "file_too_large", + "message": "signed limit exceeded", + "status": 413, + } + + def test_missing_signed_user_id_returns_400_before_service_call(self, files_app: Flask) -> None: + query = _valid_args() + del query["user_id"] + + with patch.object(module, "application_services") as application_services: + response = files_app.test_client().post( + "/files/upload/for-plugin", + query_string=query, + data={"file": (io.BytesIO(b"content"), "report.pdf")}, + content_type="multipart/form-data", + ) + + assert response.status_code == 400 + assert response.get_json()["code"] == "invalid_param" + application_services.assert_not_called() + + def test_missing_file_returns_a_structured_400(self, files_app: Flask) -> None: + with patch.object(module, "application_services") as application_services: + response = files_app.test_client().post( + "/files/upload/for-plugin", + query_string=_valid_args(), + data={}, + content_type="multipart/form-data", + ) + + assert response.status_code == 400 + assert response.get_json() == { + "code": "no_file_uploaded", + "message": "Please upload your file.", + "status": 400, + } + application_services.assert_not_called() diff --git a/api/tests/unit_tests/controllers/files/test_upload.py b/api/tests/unit_tests/controllers/files/test_upload.py deleted file mode 100644 index 7deb85df876..00000000000 --- a/api/tests/unit_tests/controllers/files/test_upload.py +++ /dev/null @@ -1,433 +0,0 @@ -import io -import types -from contextlib import contextmanager -from inspect import unwrap -from unittest.mock import MagicMock, patch - -import pytest -from sqlalchemy.orm import Session -from werkzeug.exceptions import Forbidden - -import controllers.files.upload as module -from core.workflow.file_reference import build_file_reference -from models import Account, TenantAccountJoin -from models.account import AccountStatus -from models.enums import EndUserType -from models.model import EndUser -from models.tools import ToolFile - - -def fake_request(args: dict, file=None): - return types.SimpleNamespace( - args=types.SimpleNamespace(to_dict=lambda flat=True: args), - files={"file": file} if file else {}, - ) - - -def _persist_account_memberships(session: Session) -> None: - account = Account(name="Tenant member", email="member@example.com", status=AccountStatus.ACTIVE) - account.id = "account-1" - decoy = Account(name="Other tenant member", email="decoy@example.com", status=AccountStatus.ACTIVE) - decoy.id = "account-outside-tenant" - session.add_all( - [ - account, - decoy, - TenantAccountJoin(tenant_id="tenant-1", account_id=account.id), - TenantAccountJoin(tenant_id="tenant-other", account_id=decoy.id), - ] - ) - session.commit() - - -def _end_user(user_id: str = "user-1") -> EndUser: - return EndUser( - id=user_id, - tenant_id="tenant-1", - type=EndUserType.SERVICE_API, - session_id="session-1", - ) - - -class DummyFile: - def __init__(self, filename="test.txt", mimetype="text/plain", content=b"data"): - self.filename = filename - self.mimetype = mimetype - self._content = content - self.stream = io.BytesIO(content) - - def read(self): - return self.stream.read() - - -class RecordingStream(io.BytesIO): - def __init__(self, content: bytes, events: list[str]): - super().__init__(content) - self.events = events - - def read(self, *args, **kwargs): - self.events.append("file-read") - return super().read(*args, **kwargs) - - -def _tool_file(*, name: str = "test.txt", mimetype: str = "text/plain") -> ToolFile: - tool_file = ToolFile( - user_id="user-1", - tenant_id="tenant-1", - conversation_id=None, - file_key="file-key", - mimetype=mimetype, - original_url="http://original", - name=name, - size=10, - ) - tool_file.id = "file-id" - return tool_file - - -class TestPluginUploadFileApi: - @patch.object(module, "verify_plugin_file_signature", return_value=True) - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "sign_tool_file", return_value="signed-url") - @patch.object(module, "ToolFileManager") - def test_success_upload( - self, - mock_tool_file_manager, - mock_sign_tool_file, - mock_get_user, - mock_verify_signature, - ): - dummy_file = DummyFile(filename="report.docx", mimetype="application/octet-stream") - - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - "conversation_id": "conversation-1", - }, - file=dummy_file, - ) - - tool_file_manager_instance = mock_tool_file_manager.return_value - tool_file_manager_instance.create_file_by_raw.return_value = _tool_file( - name="report.docx", - mimetype="application/octet-stream", - ) - - api = module.PluginUploadFileApi() - post_fn = unwrap(api.post) - - result, status_code = post_fn(api) - - assert status_code == 201 - assert result["id"] == "file-id" - assert result["reference"] == build_file_reference(record_id="file-id") - assert result["preview_url"] == "signed-url" - assert result["extension"] == ".docx" - mock_verify_signature.assert_called_once() - assert mock_verify_signature.call_args.kwargs["conversation_id"] == "conversation-1" - tool_file_manager_instance.create_file_by_raw.assert_called_once() - assert tool_file_manager_instance.create_file_by_raw.call_args.kwargs["conversation_id"] == "conversation-1" - mock_sign_tool_file.assert_called_once_with( - tool_file_id="file-id", - extension=".docx", - for_external=True, - ) - - @patch.object(module, "get_user") - @patch.object(module, "ToolFileManager") - @pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True) - def test_account_upload_preserves_signed_account_owner( - self, - mock_tool_file_manager, - mock_get_user, - monkeypatch: pytest.MonkeyPatch, - sqlite_session: Session, - ): - _persist_account_memberships(sqlite_session) - events: list[str] = [] - dummy_file = DummyFile(filename="report.pdf", mimetype="application/pdf", content=b"account-owned") - dummy_file.stream = RecordingStream(b"account-owned", events) - - @contextmanager - def membership_session(): - events.append("membership-session-enter") - try: - yield sqlite_session - finally: - events.append("membership-session-exit") - - monkeypatch.setattr(module.session_factory, "create_session", membership_session) - monkeypatch.setattr( - module, - "request", - fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "account-1", - "user_from": "account", - }, - file=dummy_file, - ), - ) - tool_file_manager = mock_tool_file_manager.return_value - tool_file_manager.create_file_by_raw.side_effect = lambda **_kwargs: ( - events.append("storage-create-file") or _tool_file(name="report.pdf", mimetype="application/pdf") - ) - mock_tool_file_manager.sign_file.return_value = "signed-url" - - with patch.object( - module, - "verify_plugin_file_signature", - side_effect=lambda **_kwargs: events.append("signature-verify") or True, - ) as verify_signature: - api = module.PluginUploadFileApi() - result, status_code = unwrap(api.post)(api) - - assert status_code == 201 - assert result["reference"] == build_file_reference(record_id="file-id") - assert events == [ - "membership-session-enter", - "membership-session-exit", - "signature-verify", - "file-read", - "storage-create-file", - ] - mock_get_user.assert_not_called() - verify_signature.assert_called_once_with( - filename="report.pdf", - mimetype="application/pdf", - tenant_id="tenant-1", - user_id="account-1", - conversation_id=None, - user_from="account", - timestamp="123", - nonce="abc", - sign="sig", - max_size=None, - ) - tool_file_manager.create_file_by_raw.assert_called_once_with( - user_id="account-1", - tenant_id="tenant-1", - file_binary=b"account-owned", - mimetype="application/pdf", - filename="report.pdf", - conversation_id=None, - ) - - @patch.object(module, "verify_plugin_file_signature") - @patch.object(module, "get_user") - @patch.object(module, "ToolFileManager") - @pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True) - def test_account_upload_rejects_owner_outside_tenant( - self, - mock_tool_file_manager, - mock_get_user, - mock_verify_signature, - monkeypatch: pytest.MonkeyPatch, - sqlite_session: Session, - ): - _persist_account_memberships(sqlite_session) - events: list[str] = [] - - @contextmanager - def membership_session(): - events.append("membership-session-enter") - try: - yield sqlite_session - finally: - events.append("membership-session-exit") - - monkeypatch.setattr(module.session_factory, "create_session", membership_session) - monkeypatch.setattr( - module, - "request", - fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "account-outside-tenant", - "user_from": "account", - }, - file=DummyFile(), - ), - ) - - api = module.PluginUploadFileApi() - with pytest.raises(Forbidden): - unwrap(api.post)(api) - - assert events == ["membership-session-enter", "membership-session-exit"] - mock_get_user.assert_not_called() - mock_verify_signature.assert_not_called() - mock_tool_file_manager.assert_not_called() - - def test_missing_file(self): - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - } - ) - - api = module.PluginUploadFileApi() - post_fn = unwrap(api.post) - - with pytest.raises(Forbidden): - post_fn(api) - - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "verify_plugin_file_signature", return_value=False) - def test_invalid_signature(self, mock_verify, mock_get_user): - dummy_file = DummyFile() - - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "bad", - "tenant_id": "tenant-1", - "user_id": "user-1", - }, - file=dummy_file, - ) - - api = module.PluginUploadFileApi() - post_fn = unwrap(api.post) - - with pytest.raises(Forbidden): - post_fn(api) - - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "verify_plugin_file_signature", return_value=True) - @patch.object(module, "ToolFileManager") - def test_file_too_large( - self, - mock_tool_file_manager, - mock_verify, - mock_get_user, - ): - dummy_file = DummyFile() - - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - }, - file=dummy_file, - ) - - mock_tool_file_manager.return_value.create_file_by_raw.side_effect = ( - module.services.errors.file.FileTooLargeError("too large") - ) - - api = module.PluginUploadFileApi() - post_fn = unwrap(api.post) - - with pytest.raises(module.FileTooLargeError): - post_fn(api) - - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "verify_plugin_file_signature", return_value=True) - @patch.object(module, "ToolFileManager") - def test_signed_max_size_bounds_file_read( - self, - mock_tool_file_manager, - mock_verify, - mock_get_user, - ): - dummy_file = DummyFile(content=b"data") - dummy_file.stream = MagicMock() - dummy_file.stream.read.return_value = b"data" - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - "max_size": "4", - }, - file=dummy_file, - ) - mock_tool_file_manager.return_value.create_file_by_raw.return_value = _tool_file() - mock_tool_file_manager.sign_file.return_value = "signed-url" - - unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) - - dummy_file.stream.read.assert_called_once_with(5) - assert mock_verify.call_args.kwargs["max_size"] == 4 - assert mock_tool_file_manager.return_value.create_file_by_raw.call_args.kwargs["file_binary"] == b"data" - - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "verify_plugin_file_signature", return_value=True) - @patch.object(module, "ToolFileManager") - def test_signed_max_size_rejects_oversized_file_before_creation( - self, - mock_tool_file_manager, - mock_verify, - mock_get_user, - ): - dummy_file = DummyFile(content=b"oversized") - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - "max_size": "4", - }, - file=dummy_file, - ) - - with pytest.raises(module.FileTooLargeError): - unwrap(module.PluginUploadFileApi().post)(module.PluginUploadFileApi()) - - mock_tool_file_manager.assert_not_called() - - @patch.object(module, "get_user", return_value=_end_user()) - @patch.object(module, "verify_plugin_file_signature", return_value=True) - @patch.object(module, "ToolFileManager") - def test_unsupported_file_type( - self, - mock_tool_file_manager, - mock_verify, - mock_get_user, - ): - dummy_file = DummyFile() - - module.request = fake_request( - { - "timestamp": "123", - "nonce": "abc", - "sign": "sig", - "tenant_id": "tenant-1", - "user_id": "user-1", - }, - file=dummy_file, - ) - - mock_tool_file_manager.return_value.create_file_by_raw.side_effect = ( - module.services.errors.file.UnsupportedFileTypeError() - ) - - api = module.PluginUploadFileApi() - post_fn = unwrap(api.post) - - with pytest.raises(module.UnsupportedFileTypeError): - post_fn(api) diff --git a/api/tests/unit_tests/core/tools/test_signature.py b/api/tests/unit_tests/core/tools/test_signature.py index 7dc87014935..ed1b043087a 100644 --- a/api/tests/unit_tests/core/tools/test_signature.py +++ b/api/tests/unit_tests/core/tools/test_signature.py @@ -2,6 +2,9 @@ from __future__ import annotations +import base64 +import hashlib +import hmac from collections.abc import Callable from typing import Literal from urllib.parse import parse_qs, urlparse @@ -300,3 +303,23 @@ def test_verify_plugin_file_signature_rejects_invalid_signatures( ) is False ) + + +def test_verify_plugin_file_signature_rejects_malformed_signed_timestamp() -> None: + timestamp = "not-a-timestamp" + nonce = "nonce" + payload = f"upload|report.pdf|application/pdf|tenant-id|user-id||{timestamp}|{nonce}" + sign = base64.urlsafe_b64encode(hmac.new(b"unit-secret", payload.encode(), hashlib.sha256).digest()).decode() + + assert ( + verify_plugin_file_signature( + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-id", + user_id="user-id", + timestamp=timestamp, + nonce=nonce, + sign=sign, + ) + is False + ) 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 c54de8aaec4..df115b908bf 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_statistic_query_repository import AppStatisticQueryReposit from repositories.app_tracing_config_repository import SQLAlchemyAppTracingConfigRepository from repositories.human_input_file_upload_repository import SQLAlchemyHumanInputFileUploadRepository from repositories.message_file_preview_repository import MessageFilePreviewQueryRepository +from repositories.plugin_file_upload_repository import SQLAlchemyPluginFileUploadOwnerRepository 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 @@ -76,6 +77,8 @@ from services.human_input_file_upload_service import HumanInputFileUploadService from services.init_validation_service import InvalidInitializationPasswordError from services.message_file_preview_service import MessageFilePreviewService from services.partner_tenant_binding_service import PartnerTenantBindingService +from services.plugin_file_upload_gateway import ToolFilePluginUploadGateway +from services.plugin_file_upload_service import PluginFileUploadService 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 @@ -256,6 +259,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_plugin_file_upload_boundary( + 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.plugin_file_uploads, PluginFileUploadService) + assert isinstance(services.plugin_file_uploads._owners, SQLAlchemyPluginFileUploadOwnerRepository) + assert services.plugin_file_uploads._owners._session_factory is sqlite_session_factory + assert isinstance(services.plugin_file_uploads._files, ToolFilePluginUploadGateway) + + def test_build_application_services_wires_tool_file_downloads( sqlite_session_factory: sessionmaker[Session], ) -> None: diff --git a/api/tests/unit_tests/pyrefly.toml b/api/tests/unit_tests/pyrefly.toml index e5886a9c980..651f39a79f3 100644 --- a/api/tests/unit_tests/pyrefly.toml +++ b/api/tests/unit_tests/pyrefly.toml @@ -138,7 +138,6 @@ project-excludes = [ "controllers/console/workspace/test_workspace.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", "controllers/inner_api/plugin/test_agent_config.py", "controllers/inner_api/plugin/test_plugin.py", diff --git a/api/tests/unit_tests/repositories/test_plugin_file_upload_repository.py b/api/tests/unit_tests/repositories/test_plugin_file_upload_repository.py new file mode 100644 index 00000000000..2e9fafc01cc --- /dev/null +++ b/api/tests/unit_tests/repositories/test_plugin_file_upload_repository.py @@ -0,0 +1,74 @@ +from sqlalchemy import func, select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Account, AccountStatus, TenantAccountJoin +from models.enums import EndUserType +from models.model import EndUser +from repositories.plugin_file_upload_repository import SQLAlchemyPluginFileUploadOwnerRepository + + +def _account(account_id: str) -> Account: + account = Account( + name=f"Account {account_id}", + email=f"{account_id}@example.com", + status=AccountStatus.ACTIVE, + ) + account.id = account_id + return account + + +def _end_user(*, user_id: str, tenant_id: str, session_id: str) -> EndUser: + return EndUser( + id=user_id, + tenant_id=tenant_id, + type=EndUserType.SERVICE_API, + session_id=session_id, + ) + + +def test_account_owner_must_belong_to_the_signed_tenant( + sqlite_session_factory: sessionmaker[Session], +) -> None: + member = _account("member-id") + other = _account("other-id") + with sqlite_session_factory.begin() as session: + session.add_all( + [ + member, + other, + TenantAccountJoin(tenant_id="tenant-id", account_id=member.id), + TenantAccountJoin(tenant_id="other-tenant-id", account_id=other.id), + ] + ) + + repository = SQLAlchemyPluginFileUploadOwnerRepository(session_factory=sqlite_session_factory) + + assert repository.owner_exists(tenant_id="tenant-id", user_id=member.id, user_from="account") is True + assert repository.owner_exists(tenant_id="tenant-id", user_id=other.id, user_from="account") is False + + +def test_end_user_owner_must_match_id_and_tenant( + sqlite_session_factory: sessionmaker[Session], +) -> None: + owner = _end_user(user_id="owner-id", tenant_id="tenant-id", session_id="shared-session") + other = _end_user(user_id="other-id", tenant_id="other-tenant-id", session_id="shared-session") + with sqlite_session_factory.begin() as session: + session.add_all([owner, other]) + + repository = SQLAlchemyPluginFileUploadOwnerRepository(session_factory=sqlite_session_factory) + + assert repository.owner_exists(tenant_id="tenant-id", user_id=owner.id, user_from=None) is True + assert repository.owner_exists(tenant_id="tenant-id", user_id=owner.id, user_from="end-user") is True + assert repository.owner_exists(tenant_id="tenant-id", user_id=other.id, user_from="end-user") is False + assert repository.owner_exists(tenant_id="tenant-id", user_id="shared-session", user_from=None) is False + + +def test_missing_end_user_is_not_created_during_authorization( + sqlite_session_factory: sessionmaker[Session], +) -> None: + repository = SQLAlchemyPluginFileUploadOwnerRepository(session_factory=sqlite_session_factory) + + assert repository.owner_exists(tenant_id="tenant-id", user_id="missing-id", user_from=None) is False + + with sqlite_session_factory() as session: + assert session.scalar(select(func.count()).select_from(EndUser)) == 0 diff --git a/api/tests/unit_tests/services/test_plugin_file_upload_gateway.py b/api/tests/unit_tests/services/test_plugin_file_upload_gateway.py new file mode 100644 index 00000000000..193cbc8e116 --- /dev/null +++ b/api/tests/unit_tests/services/test_plugin_file_upload_gateway.py @@ -0,0 +1,89 @@ +from unittest.mock import MagicMock, patch + +from core.tools.tool_file_manager import ToolFileManager +from core.workflow.file_reference import build_file_reference +from models.tools import ToolFile +from services.plugin_file_upload_gateway import ToolFilePluginUploadGateway +from services.plugin_file_upload_service import PluginFileUploadResult + + +def _tool_file() -> ToolFile: + file = ToolFile( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="tools/tenant-id/generated.pdf", + mimetype="application/pdf", + original_url=None, + name="report.pdf", + size=7, + ) + file.id = "file-id" + return file + + +def test_store_adapts_tool_file_to_transport_neutral_result() -> None: + tool_files = MagicMock(spec=ToolFileManager) + tool_files.create_file_by_raw.return_value = _tool_file() + gateway = ToolFilePluginUploadGateway(tool_files=tool_files) + + with patch("services.plugin_file_upload_gateway.sign_tool_file", return_value="signed-url") as sign_file: + result = gateway.store( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + content=b"content", + mimetype="application/pdf", + filename="report.pdf", + ) + + assert result == PluginFileUploadResult( + id="file-id", + reference=build_file_reference(record_id="file-id"), + name="report.pdf", + size=7, + extension=".pdf", + mime_type="application/pdf", + preview_url="signed-url", + source_url=None, + original_url=None, + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="tools/tenant-id/generated.pdf", + ) + tool_files.create_file_by_raw.assert_called_once_with( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_binary=b"content", + mimetype="application/pdf", + filename="report.pdf", + ) + sign_file.assert_called_once_with( + tool_file_id="file-id", + extension=".pdf", + for_external=True, + ) + + +def test_filename_extension_wins_over_generic_mimetype() -> None: + tool_files = MagicMock(spec=ToolFileManager) + file = _tool_file() + file.name = "report.docx" + file.mimetype = "application/octet-stream" + tool_files.create_file_by_raw.return_value = file + gateway = ToolFilePluginUploadGateway(tool_files=tool_files) + + with patch("services.plugin_file_upload_gateway.sign_tool_file", return_value="signed-url"): + result = gateway.store( + user_id="user-id", + tenant_id="tenant-id", + conversation_id=None, + content=b"content", + mimetype="application/octet-stream", + filename="report.docx", + ) + + assert result.extension == ".docx" + assert result.mime_type == "application/octet-stream" diff --git a/api/tests/unit_tests/services/test_plugin_file_upload_service.py b/api/tests/unit_tests/services/test_plugin_file_upload_service.py new file mode 100644 index 00000000000..f247c04a11a --- /dev/null +++ b/api/tests/unit_tests/services/test_plugin_file_upload_service.py @@ -0,0 +1,198 @@ +import io +from unittest.mock import Mock, patch + +import pytest + +from services.errors.file import FileTooLargeError +from services.plugin_file_upload_service import ( + PluginFileUploadAccessDeniedError, + PluginFileUploadFiles, + PluginFileUploadOwnerQuery, + PluginFileUploadResult, + PluginFileUploadService, + PluginUploadUserFrom, +) + + +def _result() -> PluginFileUploadResult: + return PluginFileUploadResult( + id="file-id", + reference="reference", + name="report.pdf", + size=4, + extension=".pdf", + mime_type="application/pdf", + preview_url="signed-url", + source_url=None, + original_url=None, + user_id="user-id", + tenant_id="tenant-id", + conversation_id=None, + file_key="file-key", + ) + + +@pytest.fixture +def owners() -> Mock: + query = Mock(spec=PluginFileUploadOwnerQuery) + query.owner_exists.return_value = True + return query + + +@pytest.fixture +def files() -> Mock: + gateway = Mock(spec=PluginFileUploadFiles) + gateway.store.return_value = _result() + return gateway + + +@pytest.fixture +def service(owners: Mock, files: Mock) -> PluginFileUploadService: + return PluginFileUploadService(owners=owners, files=files) + + +def _upload( + service: PluginFileUploadService, + *, + stream: io.BytesIO | Mock | None = None, + user_id: str = "user-id", + user_from: PluginUploadUserFrom = None, + max_size: int | None = None, +) -> PluginFileUploadResult: + return service.upload( + stream=stream or io.BytesIO(b"data"), + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-id", + user_id=user_id, + user_from=user_from, + conversation_id="conversation-id", + timestamp="123", + nonce="nonce", + sign="signature", + max_size=max_size, + ) + + +@pytest.mark.parametrize("user_from", [None, "end-user", "account"]) +def test_valid_ticket_authorizes_owner_then_stores_file( + service: PluginFileUploadService, + owners: Mock, + files: Mock, + user_from: PluginUploadUserFrom, +) -> None: + stream = io.BytesIO(b"data") + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=True) as verify: + result = _upload(service, stream=stream, user_from=user_from) + + assert result == _result() + verify.assert_called_once_with( + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-id", + user_id="user-id", + conversation_id="conversation-id", + user_from=user_from, + timestamp="123", + nonce="nonce", + sign="signature", + max_size=None, + ) + owners.owner_exists.assert_called_once_with( + tenant_id="tenant-id", + user_id="user-id", + user_from=user_from, + ) + files.store.assert_called_once_with( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + content=b"data", + mimetype="application/pdf", + filename="report.pdf", + ) + + +def test_invalid_signature_has_no_database_or_stream_side_effect( + service: PluginFileUploadService, + owners: Mock, + files: Mock, +) -> None: + stream = Mock() + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=False): + with pytest.raises(PluginFileUploadAccessDeniedError): + _upload(service, stream=stream) + + owners.owner_exists.assert_not_called() + stream.read.assert_not_called() + files.store.assert_not_called() + + +def test_unknown_owner_is_rejected_before_reading_or_storing( + service: PluginFileUploadService, + owners: Mock, + files: Mock, +) -> None: + stream = Mock() + owners.owner_exists.return_value = False + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=True): + with pytest.raises(PluginFileUploadAccessDeniedError): + _upload(service, stream=stream, user_from="account") + + stream.read.assert_not_called() + files.store.assert_not_called() + + +def test_signed_size_reads_only_one_byte_beyond_the_limit( + service: PluginFileUploadService, + files: Mock, +) -> None: + stream = Mock() + stream.read.return_value = b"data" + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=True): + _upload(service, stream=stream, max_size=4) + + stream.read.assert_called_once_with(5) + assert files.store.call_args.kwargs["content"] == b"data" + + +def test_zero_signed_size_accepts_an_empty_file( + service: PluginFileUploadService, + files: Mock, +) -> None: + stream = Mock() + stream.read.return_value = b"" + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=True): + _upload(service, stream=stream, max_size=0) + + stream.read.assert_called_once_with(1) + assert files.store.call_args.kwargs["content"] == b"" + + +@pytest.mark.parametrize( + ("max_size", "content"), + [ + pytest.param(4, b"12345", id="positive-limit"), + pytest.param(0, b"1", id="zero-limit"), + ], +) +def test_signed_size_rejects_oversized_content_before_storage( + service: PluginFileUploadService, + files: Mock, + max_size: int, + content: bytes, +) -> None: + stream = Mock() + stream.read.return_value = content + + with patch("services.plugin_file_upload_service.verify_plugin_file_signature", return_value=True): + with pytest.raises(FileTooLargeError, match="signed upload limit"): + _upload(service, stream=stream, max_size=max_size) + + stream.read.assert_called_once_with(max_size + 1) + files.store.assert_not_called()