fix: updates file size limit to use middleware and add tests for uploads (#4883)

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-11-28 09:25:26 -03:00 • committed by GitHub
commit 712a43958c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 133 additions and 13 deletions

View file

@ -12,7 +12,7 @@ from fastapi.responses import StreamingResponse
from langflow.api.utils import AsyncDbSession, CurrentActiveUser
from langflow.api.v1.schemas import UploadFileResponse
from langflow.services.database.models.flow import Flow
from langflow.services.deps import get_settings_service, get_storage_service
from langflow.services.deps import get_storage_service
from langflow.services.storage.service import StorageService
from langflow.services.storage.utils import build_content_type_from_extension
@ -46,16 +46,6 @@ async def upload_file(
session: AsyncDbSession,
storage_service: Annotated[StorageService, Depends(get_storage_service)],
) -> UploadFileResponse:
try:
max_file_size_upload = get_settings_service().settings.max_file_size_upload
except Exception as e:
raise HTTPException(status_code=500, detail=str(e)) from e
if file.size > max_file_size_upload * 1024 * 1024:
raise HTTPException(
status_code=413, detail=f"File size is larger than the maximum file size {max_file_size_upload}MB."
)
try:
flow_id_str = str(flow_id)
flow = await session.get(Flow, flow_id_str)

View file

@ -30,7 +30,7 @@ class CloudflareWorkersAIEmbeddingsComponent(LCModelComponent):
display_name="Model Name",
info="List of supported models https://developers.cloudflare.com/workers-ai/models/#text-embeddings",
required=True,
value="@cf/baai/bge-base-en-v1.5"
value="@cf/baai/bge-base-en-v1.5",
),
BoolInput(
name="strip_new_lines",
@ -75,6 +75,6 @@ class CloudflareWorkersAIEmbeddingsComponent(LCModelComponent):
strip_new_lines=self.strip_new_lines,
)
except Exception as e:
raise ValueError(f"Could not connect to CloudflareWorkersAIEmbeddings API: {str(e)}") from e
raise ValueError(f"Could not connect to CloudflareWorkersAIEmbeddings API: {e!s}") from e
return embeddings

View file

@ -28,6 +28,7 @@ from langflow.initial_setup.setup import (
from langflow.interface.types import get_and_cache_all_types_dict
from langflow.interface.utils import setup_llm_caching
from langflow.logging.logger import configure
from langflow.middleware import ContentSizeLimitMiddleware
from langflow.services.deps import get_settings_service, get_telemetry_service
from langflow.services.utils import initialize_services, teardown_services
@ -132,6 +133,10 @@ def create_app():
configure()
lifespan = get_lifespan(version=__version__)
app = FastAPI(lifespan=lifespan, title="Langflow", version=__version__)
app.add_middleware(
ContentSizeLimitMiddleware,
)
setup_sentry(app)
origins = ["*"]

View file

@ -0,0 +1,58 @@
from fastapi import HTTPException
from loguru import logger
from langflow.services.deps import get_settings_service
class MaxFileSizeException(HTTPException):
def __init__(self, detail: str = "File size is larger than the maximum file size {}MB"):
super().__init__(status_code=413, detail=detail)
# Adapted from https://github.com/steinnes/content-size-limit-asgi/blob/master/content_size_limit_asgi/middleware.py#L26
class ContentSizeLimitMiddleware:
"""Content size limiting middleware for ASGI applications.
Args:
app (ASGI application): ASGI application
max_content_size (optional): the maximum content size allowed in bytes, None for no limit
exception_cls (optional): the class of exception to raise (ContentSizeExceeded is the default)
"""
def __init__(
self,
app,
):
self.app = app
self.logger = logger
def receive_wrapper(self, receive):
received = 0
async def inner():
max_file_size_upload = get_settings_service().settings.max_file_size_upload
nonlocal received
message = await receive()
if message["type"] != "http.request" or max_file_size_upload is None:
return message
body_len = len(message.get("body", b""))
received += body_len
if received > max_file_size_upload * 1024 * 1024:
# max_content_size is in bytes, convert to MB
received_in_mb = round(received / (1024 * 1024), 3)
msg = (
f"Content size limit exceeded. Maximum allowed is {max_file_size_upload}MB"
f" and got {received_in_mb}MB."
)
raise MaxFileSizeException(msg)
return message
return inner
async def __call__(self, scope, receive, send):
if scope["type"] != "http":
await self.app(scope, receive, send)
return
wrapper = self.receive_wrapper(receive)
await self.app(scope, wrapper, send)

View file

@ -78,6 +78,7 @@ dev-dependencies = [
"asgi-lifespan>=2.1.0",
"pytest-codspeed>=3.0.0",
"pytest-github-actions-annotate-failures>=0.2.0",
"types-aiofiles>=24.1.0.20240626",
]
[build-system]