refactor(api): extract plugin file upload application service (#41808)

This commit is contained in:
非法操作 2026-09-08 03:42:32 +00:00 committed by GitHub
parent c157e59d58
commit 146b193ef7
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
18 changed files with 1093 additions and 600 deletions

View File

@ -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

View File

@ -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",
]

View File

@ -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

View File

@ -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()

View File

@ -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(

View File

@ -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),

View File

@ -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

View File

@ -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,
)

View File

@ -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,
)

View File

@ -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"]

View File

@ -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()

View File

@ -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)

View File

@ -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
)

View File

@ -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:

View File

@ -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",

View File

@ -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

View File

@ -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"

View File

@ -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()