mirror of
https://github.com/langgenius/dify.git
synced 2026-09-09 05:41:00 +08:00
refactor(api): extract plugin file upload application service (#41808)
This commit is contained in:
parent
c157e59d58
commit
146b193ef7
@ -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
|
||||
|
||||
@ -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",
|
||||
]
|
||||
|
||||
113
api/controllers/files/plugin_file_upload.py
Normal file
113
api/controllers/files/plugin_file_upload.py
Normal 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
|
||||
@ -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()
|
||||
@ -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(
|
||||
|
||||
@ -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),
|
||||
|
||||
45
api/repositories/plugin_file_upload_repository.py
Normal file
45
api/repositories/plugin_file_upload_repository.py
Normal 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
|
||||
54
api/services/plugin_file_upload_gateway.py
Normal file
54
api/services/plugin_file_upload_gateway.py
Normal 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,
|
||||
)
|
||||
116
api/services/plugin_file_upload_service.py
Normal file
116
api/services/plugin_file_upload_service.py
Normal 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,
|
||||
)
|
||||
@ -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"]
|
||||
|
||||
@ -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()
|
||||
@ -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)
|
||||
@ -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
|
||||
)
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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
|
||||
@ -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"
|
||||
198
api/tests/unit_tests/services/test_plugin_file_upload_service.py
Normal file
198
api/tests/unit_tests/services/test_plugin_file_upload_service.py
Normal 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()
|
||||
Loading…
Reference in New Issue
Block a user