fix: updates file size limit to use middleware and add tests for uploads (#4883)
This commit is contained in:
parent
1ec1a704dd
commit
712a43958c
8 changed files with 133 additions and 13 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = ["*"]
|
||||
|
||||
|
|
|
|||
58
src/backend/base/langflow/middleware.py
Normal file
58
src/backend/base/langflow/middleware.py
Normal 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)
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue