diff --git a/src/backend/base/langflow/__main__.py b/src/backend/base/langflow/__main__.py index fb8c01512..695e5b95c 100644 --- a/src/backend/base/langflow/__main__.py +++ b/src/backend/base/langflow/__main__.py @@ -138,6 +138,11 @@ def run( help="Defines the number of retries for the health check.", envvar="LANGFLOW_HEALTH_CHECK_MAX_RETRIES", ), + max_file_size_upload: int = typer.Option( + 100, + help="Defines the maximum file size for the upload in MB.", + envvar="LANGFLOW_MAX_FILE_SIZE_UPLOAD", + ), ): """ Run Langflow. @@ -158,6 +163,7 @@ def run( auto_saving=auto_saving, auto_saving_interval=auto_saving_interval, health_check_max_retries=health_check_max_retries, + max_file_size_upload=max_file_size_upload, ) # create path object if path is provided static_files_dir: Path | None = Path(path) if path else None diff --git a/src/backend/base/langflow/api/v1/files.py b/src/backend/base/langflow/api/v1/files.py index 6304dfcc3..b72f00999 100644 --- a/src/backend/base/langflow/api/v1/files.py +++ b/src/backend/base/langflow/api/v1/files.py @@ -43,6 +43,12 @@ async def upload_file( storage_service: StorageService = Depends(get_storage_service), ): try: + max_file_size_upload = get_storage_service().settings_service.settings.max_file_size_upload + 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." + ) + flow_id_str = str(flow_id) file_content = await file.read() timestamp = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") diff --git a/src/backend/base/langflow/api/v1/schemas.py b/src/backend/base/langflow/api/v1/schemas.py index 959042b34..0b2ec1ad9 100644 --- a/src/backend/base/langflow/api/v1/schemas.py +++ b/src/backend/base/langflow/api/v1/schemas.py @@ -16,6 +16,7 @@ from langflow.services.database.models.base import orjson_dumps from langflow.services.database.models.flow import FlowCreate, FlowRead from langflow.services.database.models.user import UserRead from langflow.services.tracing.schema import Log +from langflow.utils.util_strings import truncate_long_strings class BuildStatus(Enum): @@ -281,6 +282,12 @@ class VertexBuildResponse(BaseModel): timestamp: datetime | None = Field(default_factory=lambda: datetime.now(timezone.utc)) """Timestamp of the build.""" + @field_serializer("data") + def serialize_data(self, data: ResultDataResponse) -> dict: + data_dict = data.model_dump() if isinstance(data, BaseModel) else data + truncated_data = truncate_long_strings(data_dict) + return truncated_data + class VerticesBuiltResponse(BaseModel): vertices: list[VertexBuildResponse] @@ -341,3 +348,4 @@ class ConfigResponse(BaseModel): auto_saving: bool auto_saving_interval: int health_check_max_retries: int + max_file_size_upload: int diff --git a/src/backend/base/langflow/services/database/models/transactions/model.py b/src/backend/base/langflow/services/database/models/transactions/model.py index effcf36cb..493aec48f 100644 --- a/src/backend/base/langflow/services/database/models/transactions/model.py +++ b/src/backend/base/langflow/services/database/models/transactions/model.py @@ -2,12 +2,14 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING from uuid import UUID, uuid4 -from pydantic import field_validator +from pydantic import field_serializer, field_validator from sqlmodel import JSON, Column, Field, Relationship, SQLModel if TYPE_CHECKING: from langflow.services.database.models.flow.model import Flow +from langflow.utils.util_strings import truncate_long_strings + class TransactionBase(SQLModel): timestamp: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) @@ -32,6 +34,11 @@ class TransactionBase(SQLModel): value = UUID(value) return value + @field_serializer("outputs") + def serialize_outputs(self, data) -> dict: + truncated_data = truncate_long_strings(data) + return truncated_data + class TransactionTable(TransactionBase, table=True): # type: ignore __tablename__ = "transaction" diff --git a/src/backend/base/langflow/services/database/models/vertex_builds/model.py b/src/backend/base/langflow/services/database/models/vertex_builds/model.py index ab535dbc2..69341d828 100644 --- a/src/backend/base/langflow/services/database/models/vertex_builds/model.py +++ b/src/backend/base/langflow/services/database/models/vertex_builds/model.py @@ -8,6 +8,8 @@ from sqlmodel import JSON, Column, Field, Relationship, SQLModel if TYPE_CHECKING: from langflow.services.database.models.flow.model import Flow +from langflow.utils.util_strings import truncate_long_strings + class VertexBuildBase(SQLModel): timestamp: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) @@ -38,6 +40,16 @@ class VertexBuildBase(SQLModel): value = value.replace(tzinfo=timezone.utc) return value + @field_serializer("data") + def serialize_data(self, data: dict) -> dict: + truncated_data = truncate_long_strings(data) + return truncated_data + + @field_serializer("artifacts") + def serialize_artifacts(self, data) -> dict: + truncated_data = truncate_long_strings(data) + return truncated_data + class VertexBuildTable(VertexBuildBase, table=True): # type: ignore __tablename__ = "vertex_build" diff --git a/src/backend/base/langflow/services/settings/base.py b/src/backend/base/langflow/services/settings/base.py index b62790312..3f8413036 100644 --- a/src/backend/base/langflow/services/settings/base.py +++ b/src/backend/base/langflow/services/settings/base.py @@ -153,6 +153,8 @@ class Settings(BaseSettings): """The interval in ms at which Langflow will auto save flows.""" health_check_max_retries: int = 5 """The maximum number of retries for the health check.""" + max_file_size_upload: int = 100 + """The maximum file size for the upload in MB.""" @field_validator("dev") @classmethod diff --git a/src/backend/base/langflow/utils/constants.py b/src/backend/base/langflow/utils/constants.py index 5bccdbfef..802aec715 100644 --- a/src/backend/base/langflow/utils/constants.py +++ b/src/backend/base/langflow/utils/constants.py @@ -183,3 +183,5 @@ MESSAGE_SENDER_AI = "Machine" MESSAGE_SENDER_USER = "User" MESSAGE_SENDER_NAME_AI = "AI" MESSAGE_SENDER_NAME_USER = "User" + +MAX_TEXT_LENGTH = 99999 diff --git a/src/backend/base/langflow/utils/util.py b/src/backend/base/langflow/utils/util.py index a4cf3f6f4..c343859b7 100644 --- a/src/backend/base/langflow/utils/util.py +++ b/src/backend/base/langflow/utils/util.py @@ -431,6 +431,7 @@ def update_settings( auto_saving: bool = True, auto_saving_interval: int = 1000, health_check_max_retries: int = 5, + max_file_size_upload: int = 100, ): """Update the settings from a config file.""" from langflow.services.utils import initialize_settings_service @@ -463,6 +464,9 @@ def update_settings( if health_check_max_retries is not None: logger.debug(f"Setting health_check_max_retries to {health_check_max_retries}") settings_service.settings.update_settings(health_check_max_retries=health_check_max_retries) + if max_file_size_upload is not None: + logger.debug(f"Setting max_file_size_upload to {max_file_size_upload}") + settings_service.settings.update_settings(max_file_size_upload=max_file_size_upload) def is_class_method(func, cls): diff --git a/src/backend/base/langflow/utils/util_strings.py b/src/backend/base/langflow/utils/util_strings.py new file mode 100644 index 000000000..4383bc872 --- /dev/null +++ b/src/backend/base/langflow/utils/util_strings.py @@ -0,0 +1,28 @@ +from langflow.utils import constants + + +def truncate_long_strings(data, max_length=None): + """ + Recursively traverse the dictionary or list and truncate strings longer than max_length. + """ + + if max_length is None: + max_length = constants.MAX_TEXT_LENGTH + + if max_length < 0 or not isinstance(data, dict | list): + return data + + if isinstance(data, dict): + for key, value in data.items(): + if isinstance(value, str) and len(value) > max_length: + data[key] = value[:max_length] + "..." + elif isinstance(value, (dict | list)): + truncate_long_strings(value, max_length) + elif isinstance(data, list): + for index, item in enumerate(data): + if isinstance(item, str) and len(item) > max_length: + data[index] = item[:max_length] + "..." + elif isinstance(item, (dict | list)): + truncate_long_strings(item, max_length) + + return data diff --git a/src/backend/tests/unit/utils/test_truncate_long_strings_on_objects.py b/src/backend/tests/unit/utils/test_truncate_long_strings_on_objects.py new file mode 100644 index 000000000..3e1b1df32 --- /dev/null +++ b/src/backend/tests/unit/utils/test_truncate_long_strings_on_objects.py @@ -0,0 +1,98 @@ +from langflow.utils.util_strings import truncate_long_strings +from langflow.utils.constants import MAX_TEXT_LENGTH +import pytest + + +@pytest.mark.parametrize( + "input_data, max_length, expected", + [ + # Test case 1: Simple string truncation + ({"key": "a" * 100}, 10, {"key": "a" * 10 + "..."}), + # Test case 2: Nested dictionary + ({"outer": {"inner": "b" * 100}}, 5, {"outer": {"inner": "b" * 5 + "..."}}), + # Test case 3: List of strings + (["short", "a" * 100, "also short"], 7, ["short", "a" * 7 + "...", "also sh" + "..."]), + # Test case 4: Mixed nested structure + ( + {"key1": ["a" * 100, {"nested": "b" * 100}], "key2": "c" * 100}, + 8, + {"key1": ["a" * 8 + "...", {"nested": "b" * 8 + "..."}], "key2": "c" * 8 + "..."}, + ), + # Test case 5: Empty structures + ({}, 10, {}), + ([], 10, []), + # Test case 6: Strings at exact max_length + ({"exact": "a" * 10}, 10, {"exact": "a" * 10}), + # Test case 7: Non-string values + ({"num": 12345, "bool": True, "none": None}, 5, {"num": 12345, "bool": True, "none": None}), + # Test case 8: Unicode characters + ({"unicode": "こんにちは世界"}, 3, {"unicode": "こんに..."}), + # Test case 9: Very large structure + ( + {"key" + str(i): "value" * i for i in range(1000)}, + 10, + {"key" + str(i): ("value" * i)[:10] + "..." if len("value" * i) > 10 else "value" * i for i in range(1000)}, + ), + ], +) +def test_truncate_long_strings(input_data, max_length, expected): + result = truncate_long_strings(input_data, max_length) + assert result == expected + + +def test_truncate_long_strings_default_max_length(): + long_string = "a" * (MAX_TEXT_LENGTH + 1) + input_data = {"key": long_string} + result = truncate_long_strings(input_data) + assert len(result["key"]) == MAX_TEXT_LENGTH + 3 # +3 for the "..." + + +def test_truncate_long_strings_no_modification(): + input_data = {"short": "short string", "nested": {"also_short": "another short string"}} + result = truncate_long_strings(input_data, 100) + assert result == input_data + + +# Test for type preservation +def test_truncate_long_strings_type_preservation(): + input_data = {"str": "a" * 100, "list": ["b" * 100], "dict": {"nested": "c" * 100}} + result = truncate_long_strings(input_data, 10) + assert isinstance(result, dict) + assert isinstance(result["str"], str) + assert isinstance(result["list"], list) + assert isinstance(result["dict"], dict) + + +# Test for in-place modification +def test_truncate_long_strings_in_place_modification(): + input_data = {"key": "a" * 100} + result = truncate_long_strings(input_data, 10) + assert result is input_data # Check if the same object is returned + + +# Test for invalid input +def test_truncate_long_strings_invalid_input(): + input_string = "not a dict or list" + result = truncate_long_strings(input_string, 10) + assert result == input_string # The function should return the input unchanged + + +# Updated test for negative max_length +def test_truncate_long_strings_negative_max_length(): + input_data = {"key": "value"} + result = truncate_long_strings(input_data, -1) + assert result == input_data # Assuming the function ignores negative max_length + + +# Additional test for zero max_length +def test_truncate_long_strings_zero_max_length(): + input_data = {"key": "value"} + result = truncate_long_strings(input_data, 0) + assert result == {"key": "..."} # Assuming the function truncates to just "..." + + +# Test for very small positive max_length +def test_truncate_long_strings_small_max_length(): + input_data = {"key": "value"} + result = truncate_long_strings(input_data, 1) + assert result == {"key": "v..."} # Assuming the function keeps at least one character diff --git a/src/frontend/src/CustomNodes/GenericNode/components/outputModal/components/switchOutputView/index.tsx b/src/frontend/src/CustomNodes/GenericNode/components/outputModal/components/switchOutputView/index.tsx index e683fe5e0..4cfb192ca 100644 --- a/src/frontend/src/CustomNodes/GenericNode/components/outputModal/components/switchOutputView/index.tsx +++ b/src/frontend/src/CustomNodes/GenericNode/components/outputModal/components/switchOutputView/index.tsx @@ -1,4 +1,6 @@ +import { MAX_TEXT_LENGTH } from "@/constants/constants"; import { LogsLogType, OutputLogType } from "@/types/api"; +import { useMemo } from "react"; import DataOutputComponent from "../../../../../../components/dataOutputComponent"; import ForwardedIconComponent from "../../../../../../components/genericIconComponent"; import { @@ -16,6 +18,7 @@ interface SwitchOutputViewProps { outputName: string; type: "Outputs" | "Logs"; } + const SwitchOutputView: React.FC = ({ nodeId, outputName, @@ -35,6 +38,35 @@ const SwitchOutputView: React.FC = ({ if (resultMessage?.raw) { resultMessage = resultMessage.raw; } + + const resultMessageMemoized = useMemo(() => { + if ( + typeof resultMessage === "string" && + resultMessage.length > MAX_TEXT_LENGTH + ) { + resultMessage = `${resultMessage.substring(0, MAX_TEXT_LENGTH)}...`; + } + + if (Array.isArray(resultMessage)) { + resultMessage = resultMessage.map((item) => { + if (item && typeof item.data === "object") { + const truncatedData = Object.fromEntries( + Object.entries(item.data).map(([key, value]) => { + if (typeof value === "string" && value.length > MAX_TEXT_LENGTH) { + return [key, `${value.substring(0, MAX_TEXT_LENGTH)}...`]; + } + return [key, value]; + }), + ); + return { ...item, data: truncatedData }; + } + return item; + }); + } + + return resultMessage; + }, [resultMessage]); + return type === "Outputs" ? ( <> @@ -42,23 +74,23 @@ const SwitchOutputView: React.FC = ({ - + ).every((item) => item.data) - ? (resultMessage as Array).map((item) => item.data) - : resultMessage - : Object.keys(resultMessage).length > 0 - ? [resultMessage] + Array.isArray(resultMessageMemoized) + ? (resultMessageMemoized as Array).every((item) => item.data) + ? (resultMessageMemoized as Array).map((item) => item.data) + : resultMessageMemoized + : Object.keys(resultMessageMemoized).length > 0 + ? [resultMessageMemoized] : [] } pagination={true} diff --git a/src/frontend/src/components/inputFileComponent/index.tsx b/src/frontend/src/components/inputFileComponent/index.tsx index ae1108b7f..55e14c5d1 100644 --- a/src/frontend/src/components/inputFileComponent/index.tsx +++ b/src/frontend/src/components/inputFileComponent/index.tsx @@ -1,6 +1,6 @@ -import { maxSizeFilesInBytes } from "@/constants/constants"; import { usePostUploadFile } from "@/controllers/API/queries/files/use-post-upload-file"; import { createFileUpload } from "@/helpers/create-file-upload"; +import { useUtilityStore } from "@/stores/utilityStore"; import { useEffect } from "react"; import { CONSOLE_ERROR_MSG, @@ -23,7 +23,7 @@ export default function InputFileComponent({ }: FileComponentType): JSX.Element { const currentFlowId = useFlowsManagerStore((state) => state.currentFlowId); const setErrorData = useAlertStore((state) => state.setErrorData); - + const maxFileSizeUpload = useUtilityStore((state) => state.maxFileSizeUpload); // Clear component state useEffect(() => { if (disabled && value !== "") { @@ -47,9 +47,9 @@ export default function InputFileComponent({ createFileUpload({ multiple: false, accept: fileTypes?.join(",") }).then( (files) => { const file = files[0]; - if (file.size > maxSizeFilesInBytes) { + if (file.size > maxFileSizeUpload) { setErrorData({ - title: INVALID_FILE_SIZE_ALERT(10), + title: INVALID_FILE_SIZE_ALERT(maxFileSizeUpload / 1024 / 1024), }); return; } @@ -68,8 +68,12 @@ export default function InputFileComponent({ // sets the value to the user handleOnNewValue({ value: file.name, file_path }); }, - onError: () => { + onError: (error) => { console.error(CONSOLE_ERROR_MSG); + setErrorData({ + title: "Error uploading file", + list: [error.response?.data?.detail], + }); }, }, ); diff --git a/src/frontend/src/constants/constants.ts b/src/frontend/src/constants/constants.ts index f357b1e00..a8e95ecbf 100644 --- a/src/frontend/src/constants/constants.ts +++ b/src/frontend/src/constants/constants.ts @@ -916,4 +916,4 @@ export const COLOR_OPTIONS = { red: "var(--note-red)", }; -export const maxSizeFilesInBytes = 10 * 1024 * 1024; // 10MB in bytes +export const MAX_TEXT_LENGTH = 99999; diff --git a/src/frontend/src/controllers/API/queries/config/use-get-config.ts b/src/frontend/src/controllers/API/queries/config/use-get-config.ts index c5f9fed5b..dad80c811 100644 --- a/src/frontend/src/controllers/API/queries/config/use-get-config.ts +++ b/src/frontend/src/controllers/API/queries/config/use-get-config.ts @@ -1,4 +1,5 @@ import useFlowsManagerStore from "@/stores/flowsManagerStore"; +import { useUtilityStore } from "@/stores/utilityStore"; import axios from "axios"; import { useQueryFunctionType } from "../../../../types/api"; import { api } from "../../api"; @@ -10,6 +11,7 @@ export interface ConfigResponse { auto_saving: boolean; auto_saving_interval: number; health_check_max_retries: number; + max_file_size_upload: number; } export const useGetConfig: useQueryFunctionType = ( @@ -22,6 +24,9 @@ export const useGetConfig: useQueryFunctionType = ( const setHealthCheckMaxRetries = useFlowsManagerStore( (state) => state.setHealthCheckMaxRetries, ); + const setMaxFileSizeUpload = useUtilityStore( + (state) => state.setMaxFileSizeUpload, + ); const { query } = UseRequestProcessor(); @@ -37,6 +42,7 @@ export const useGetConfig: useQueryFunctionType = ( setAutoSaving(data.auto_saving); setAutoSavingInterval(data.auto_saving_interval); setHealthCheckMaxRetries(data.health_check_max_retries); + setMaxFileSizeUpload(data.max_file_size_upload); } return data; }; diff --git a/src/frontend/src/modals/IOModal/components/IOFieldView/components/FileInput/index.tsx b/src/frontend/src/modals/IOModal/components/IOFieldView/components/FileInput/index.tsx index 2a362a7f1..f076a2e8f 100644 --- a/src/frontend/src/modals/IOModal/components/IOFieldView/components/FileInput/index.tsx +++ b/src/frontend/src/modals/IOModal/components/IOFieldView/components/FileInput/index.tsx @@ -1,7 +1,10 @@ import { Button } from "../../../../../../components/ui/button"; +import { INVALID_FILE_SIZE_ALERT } from "@/constants/alerts_constants"; import { usePostUploadFile } from "@/controllers/API/queries/files/use-post-upload-file"; import { createFileUpload } from "@/helpers/create-file-upload"; +import useAlertStore from "@/stores/alertStore"; +import { useUtilityStore } from "@/stores/utilityStore"; import { useEffect, useState } from "react"; import IconComponent from "../../../../../../components/genericIconComponent"; import { @@ -18,6 +21,8 @@ export default function IOFileInput({ field, updateValue }: IOFileInputProps) { const [isDragging, setIsDragging] = useState(false); const [filePath, setFilePath] = useState(""); const [image, setImage] = useState(null); + const setErrorData = useAlertStore((state) => state.setErrorData); + const maxFileSizeUpload = useUtilityStore((state) => state.maxFileSizeUpload); useEffect(() => { if (filePath) { @@ -74,6 +79,13 @@ export default function IOFileInput({ field, updateValue }: IOFileInputProps) { const upload = async (file) => { if (file) { + if (file.size > maxFileSizeUpload) { + setErrorData({ + title: INVALID_FILE_SIZE_ALERT(maxFileSizeUpload / 1024 / 1024), + }); + return; + } + // Check if a file was selected const fileReader = new FileReader(); fileReader.onload = (event) => { @@ -93,7 +105,11 @@ export default function IOFileInput({ field, updateValue }: IOFileInputProps) { const { file_path } = data; setFilePath(file_path); }, - onError: () => { + onError: (error) => { + setErrorData({ + title: "Error uploading file", + list: [error.response?.data?.detail], + }); console.error("Error occurred while uploading file"); }, }, diff --git a/src/frontend/src/modals/IOModal/components/chatView/chatInput/index.tsx b/src/frontend/src/modals/IOModal/components/chatView/chatInput/index.tsx index 24fd80a06..e66d80205 100644 --- a/src/frontend/src/modals/IOModal/components/chatView/chatInput/index.tsx +++ b/src/frontend/src/modals/IOModal/components/chatView/chatInput/index.tsx @@ -1,5 +1,7 @@ +import { INVALID_FILE_SIZE_ALERT } from "@/constants/alerts_constants"; import { usePostUploadFile } from "@/controllers/API/queries/files/use-post-upload-file"; import useAlertStore from "@/stores/alertStore"; +import { useUtilityStore } from "@/stores/utilityStore"; import { useEffect, useRef, useState } from "react"; import ShortUniqueId from "short-unique-id"; import { @@ -36,6 +38,7 @@ export default function ChatInput({ const [inputFocus, setInputFocus] = useState(false); const fileInputRef = useRef(null); const setErrorData = useAlertStore((state) => state.setErrorData); + const maxFileSizeUpload = useUtilityStore((state) => state.maxFileSizeUpload); useFocusOnUnlock(lockChat, inputRef); useAutoResizeTextArea(chatValue, inputRef); @@ -62,10 +65,16 @@ export default function ChatInput({ const fileInput = event.target as HTMLInputElement; file = fileInput.files?.[0] ?? null; } - if (file) { const fileExtension = file.name.split(".").pop()?.toLowerCase(); + if (file.size > maxFileSizeUpload) { + setErrorData({ + title: INVALID_FILE_SIZE_ALERT(maxFileSizeUpload / 1024 / 1024), + }); + return; + } + if ( !fileExtension || !ALLOWED_IMAGE_INPUT_EXTENSIONS.includes(fileExtension) @@ -99,7 +108,7 @@ export default function ChatInput({ return newFiles; }); }, - onError: () => { + onError: (error) => { setFiles((prev) => { const newFiles = [...prev]; const updatedIndex = newFiles.findIndex((file) => file.id === id); @@ -107,6 +116,10 @@ export default function ChatInput({ newFiles[updatedIndex].error = true; return newFiles; }); + setErrorData({ + title: "Error uploading file", + list: [error.response?.data?.detail], + }); }, }, ); diff --git a/src/frontend/src/modals/IOModal/components/chatView/index.tsx b/src/frontend/src/modals/IOModal/components/chatView/index.tsx index d78246579..28a39f7c4 100644 --- a/src/frontend/src/modals/IOModal/components/chatView/index.tsx +++ b/src/frontend/src/modals/IOModal/components/chatView/index.tsx @@ -1,6 +1,8 @@ +import { INVALID_FILE_SIZE_ALERT } from "@/constants/alerts_constants"; import { useDeleteBuilds } from "@/controllers/API/queries/_builds"; import { usePostUploadFile } from "@/controllers/API/queries/files/use-post-upload-file"; import { track } from "@/customization/utils/analytics"; +import { useUtilityStore } from "@/stores/utilityStore"; import { useEffect, useRef, useState } from "react"; import ShortUniqueId from "short-unique-id"; import IconComponent from "../../../../components/genericIconComponent"; @@ -42,6 +44,7 @@ export default function ChatView({ const updateFlowPool = useFlowStore((state) => state.updateFlowPool); const [id, setId] = useState(""); const { mutate: mutateDeleteFlowPool } = useDeleteBuilds(); + const maxFileSizeUpload = useUtilityStore((state) => state.maxFileSizeUpload); //build chat history useEffect(() => { @@ -173,6 +176,13 @@ export default function ChatView({ if (files) { const file = files?.[0]; const fileExtension = file.name.split(".").pop()?.toLowerCase(); + if (file.size > maxFileSizeUpload) { + setErrorData({ + title: INVALID_FILE_SIZE_ALERT(maxFileSizeUpload / 1024 / 1024), + }); + return; + } + if ( !fileExtension || !ALLOWED_IMAGE_INPUT_EXTENSIONS.includes(fileExtension) @@ -208,7 +218,7 @@ export default function ChatView({ return newFiles; }); }, - onError: () => { + onError: (error) => { setFiles((prev) => { const newFiles = [...prev]; const updatedIndex = newFiles.findIndex((file) => file.id === id); @@ -216,6 +226,10 @@ export default function ChatView({ newFiles[updatedIndex].error = true; return newFiles; }); + setErrorData({ + title: "Error uploading file", + list: [error.response?.data?.detail], + }); }, }, ); diff --git a/src/frontend/src/stores/utilityStore.ts b/src/frontend/src/stores/utilityStore.ts index 32d82dd8c..c42afe18b 100644 --- a/src/frontend/src/stores/utilityStore.ts +++ b/src/frontend/src/stores/utilityStore.ts @@ -18,4 +18,7 @@ export const useUtilityStore = create((set, get) => ({ playgroundScrollBehaves: "instant", setPlaygroundScrollBehaves: (behaves: ScrollBehavior) => set({ playgroundScrollBehaves: behaves }), + maxFileSizeUpload: 100 * 1024 * 1024, // 100MB in bytes + setMaxFileSizeUpload: (maxFileSizeUpload: number) => + set({ maxFileSizeUpload: maxFileSizeUpload * 1024 * 1024 }), })); diff --git a/src/frontend/src/types/zustand/utility/index.ts b/src/frontend/src/types/zustand/utility/index.ts index a14f8bdf8..17a2e69cb 100644 --- a/src/frontend/src/types/zustand/utility/index.ts +++ b/src/frontend/src/types/zustand/utility/index.ts @@ -5,4 +5,6 @@ export type UtilityStoreType = { setHealthCheckTimeout: (timeout: string | null) => void; playgroundScrollBehaves: ScrollBehavior; setPlaygroundScrollBehaves: (behaves: ScrollBehavior) => void; + maxFileSizeUpload: number; + setMaxFileSizeUpload: (maxFileSizeUpload: number) => void; }; diff --git a/src/frontend/tests/extended/features/limit-file-size-upload.spec.ts b/src/frontend/tests/extended/features/limit-file-size-upload.spec.ts new file mode 100644 index 000000000..5d4dc8bee --- /dev/null +++ b/src/frontend/tests/extended/features/limit-file-size-upload.spec.ts @@ -0,0 +1,128 @@ +import { expect, test } from "@playwright/test"; +import * as dotenv from "dotenv"; +import { readFileSync } from "fs"; +import path from "path"; + +test("user should not be able to upload a file larger than the limit", async ({ + page, +}) => { + const maxFileSizeUpload = 0.001; + await page.route("**/api/v1/config", (route) => { + route.fulfill({ + status: 200, + contentType: "application/json", + body: JSON.stringify({ + max_file_size_upload: maxFileSizeUpload, + }), + headers: { + "content-type": "application/json", + ...route.request().headers(), + }, + }); + }); + test.skip( + !process?.env?.OPENAI_API_KEY, + "OPENAI_API_KEY required to run this test", + ); + + if (!process.env.CI) { + dotenv.config({ path: path.resolve(__dirname, "../../.env") }); + } + + await page.goto("/"); + + await page.waitForTimeout(1000); + + let modalCount = 0; + try { + const modalTitleElement = await page?.getByTestId("modal-title"); + if (modalTitleElement) { + modalCount = await modalTitleElement.count(); + } + } catch (error) { + modalCount = 0; + } + + while (modalCount === 0) { + await page.getByText("New Project", { exact: true }).click(); + await page.waitForTimeout(3000); + modalCount = await page.getByTestId("modal-title")?.count(); + } + + await page.getByRole("heading", { name: "Basic Prompting" }).click(); + await page.waitForSelector('[title="fit view"]', { + timeout: 100000, + }); + + await page.getByTitle("fit view").click(); + await page.getByTitle("zoom out").click(); + await page.getByTitle("zoom out").click(); + await page.getByTitle("zoom out").click(); + + let outdatedComponents = await page.getByTestId("icon-AlertTriangle").count(); + + while (outdatedComponents > 0) { + await page.getByTestId("icon-AlertTriangle").first().click(); + await page.waitForTimeout(1000); + outdatedComponents = await page.getByTestId("icon-AlertTriangle").count(); + } + + await page + .getByTestId("popover-anchor-input-api_key") + .fill(process.env.OPENAI_API_KEY ?? ""); + + await page.getByTestId("dropdown_str_model_name").click(); + await page.getByTestId("gpt-4o-1-option").click(); + + await page.waitForSelector("text=Chat Input", { timeout: 30000 }); + + await page.getByText("Chat Input", { exact: true }).click(); + await page.getByTestId("more-options-modal").click(); + await page.getByTestId("edit-button-modal").click(); + await page.getByText("Close").last().click(); + + await page.getByText("Playground", { exact: true }).click(); + + // Read the image file as a binary string + const filePath = "tests/assets/chain.png"; + const fileContent = readFileSync(filePath, "base64"); + + // Create the DataTransfer and File objects within the browser context + const dataTransfer = await page.evaluateHandle( + ({ fileContent }) => { + const dt = new DataTransfer(); + const byteCharacters = atob(fileContent); + const byteNumbers = new Array(byteCharacters.length); + for (let i = 0; i < byteCharacters.length; i++) { + byteNumbers[i] = byteCharacters.charCodeAt(i); + } + const byteArray = new Uint8Array(byteNumbers); + const file = new File([byteArray], "chain.png", { type: "image/png" }); + dt.items.add(file); + return dt; + }, + { fileContent }, + ); + + await page.waitForSelector('[data-testid="input-chat-playground"]', { + timeout: 100000, + }); + + // Locate the target element + const element = await page.getByTestId("input-chat-playground"); + + // Dispatch the drop event on the target element + await element.dispatchEvent("drop", { dataTransfer }); + + await page.waitForTimeout(1000); + + await page.waitForSelector("text=The file size is too large", { + timeout: 10000, + }); + + await expect( + page.getByText( + `The file size is too large. Please select a file smaller than ${maxFileSizeUpload}MB`, + ), + ).toBeVisible(); +}); diff --git a/src/frontend/vite.config.mts b/src/frontend/vite.config.mts index 9183dc0d8..5296c259e 100644 --- a/src/frontend/vite.config.mts +++ b/src/frontend/vite.config.mts @@ -36,7 +36,9 @@ export default defineConfig(({ mode }) => { }, define: { "process.env.BACKEND_URL": JSON.stringify(env.BACKEND_URL), - "process.env.ACCESS_TOKEN_EXPIRE_SECONDS": JSON.stringify(env.ACCESS_TOKEN_EXPIRE_SECONDS), + "process.env.ACCESS_TOKEN_EXPIRE_SECONDS": JSON.stringify( + env.ACCESS_TOKEN_EXPIRE_SECONDS, + ), "process.env.CI": JSON.stringify(env.CI), }, plugins: [react(), svgr(), tsconfigPaths()], @@ -47,4 +49,4 @@ export default defineConfig(({ mode }) => { }, }, }; -}); \ No newline at end of file +});