diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index b9bb15c7b400..6c61f382d56b 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -36,6 +36,7 @@ from open_webui.routers.audio import transcribe from open_webui.storage.provider import Storage from open_webui.utils.auth import get_admin_user, get_verified_user +from open_webui.utils.file_limits import file_size_exceeds_limit, get_upload_file_size from pydantic import BaseModel log = logging.getLogger(__name__) @@ -107,6 +108,15 @@ def upload_file( unsanitized_filename = file.filename filename = os.path.basename(unsanitized_filename) + if not internal: + max_file_size = request.app.state.config.FILE_MAX_SIZE + file_size = get_upload_file_size(file) + if file_size_exceeds_limit(file_size, max_file_size): + raise HTTPException( + status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, + detail=ERROR_MESSAGES.FILE_TOO_LARGE(f"{max_file_size} MB"), + ) + file_extension = os.path.splitext(filename)[1] # Remove the leading dot from the file extension file_extension = file_extension[1:] if file_extension else "" @@ -204,6 +214,8 @@ def upload_file( detail=ERROR_MESSAGES.DEFAULT("Error uploading file"), ) + except HTTPException: + raise except Exception as e: log.exception(e) raise HTTPException( diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index e6e55f4d3880..99e9b57d6b17 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -22,6 +22,11 @@ from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_verified_user from open_webui.utils.access_control import has_access, has_permission +from open_webui.utils.knowledge_files import ( + file_id_in_knowledge, + get_knowledge_file_ids, + user_owns_file_or_is_admin, +) from open_webui.env import SRC_LOG_LEVELS @@ -361,6 +366,11 @@ def add_file_to_knowledge_by_id( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.NOT_FOUND, ) + if not user_owns_file_or_is_admin(user, file): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) if not file.data: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -440,6 +450,11 @@ def update_file_from_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) + if not file_id_in_knowledge(knowledge, form_data.file_id): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT("file_id"), + ) file = Files.get_file_by_id(form_data.file_id) if not file: raise HTTPException( @@ -510,6 +525,11 @@ def remove_file_from_knowledge_by_id( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) + if not file_id_in_knowledge(knowledge, form_data.file_id): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT("file_id"), + ) file = Files.get_file_by_id(form_data.file_id) if not file: raise HTTPException( @@ -527,22 +547,9 @@ def remove_file_from_knowledge_by_id( log.debug(e) pass - try: - # Remove the file's collection from vector database - file_collection = f"file-{form_data.file_id}" - if VECTOR_DB_CLIENT.has_collection(collection_name=file_collection): - VECTOR_DB_CLIENT.delete_collection(collection_name=file_collection) - except Exception as e: - log.debug("This was most likely caused by bypassing embedding processing") - log.debug(e) - pass - - # Delete file from database - Files.delete_file_by_id(form_data.file_id) - if knowledge: data = knowledge.data or {} - file_ids = data.get("file_ids", []) + file_ids = get_knowledge_file_ids(knowledge) if form_data.file_id in file_ids: file_ids.remove(form_data.file_id) @@ -714,6 +721,11 @@ def add_files_to_knowledge_batch( status_code=status.HTTP_400_BAD_REQUEST, detail=f"File {form.file_id} not found", ) + if not user_owns_file_or_is_admin(user, file): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) files.append(file) # Process files diff --git a/backend/open_webui/utils/file_limits.py b/backend/open_webui/utils/file_limits.py new file mode 100644 index 000000000000..9443b8d58c8e --- /dev/null +++ b/backend/open_webui/utils/file_limits.py @@ -0,0 +1,36 @@ +import os +from typing import Any, Optional + + +BYTES_PER_MB = 1024 * 1024 + + +def get_upload_file_size(file: Any) -> Optional[int]: + size = getattr(file, "size", None) + if isinstance(size, int): + return size + + stream = getattr(file, "file", None) + if stream is None: + return None + + try: + current_position = stream.tell() + stream.seek(0, os.SEEK_END) + size = stream.tell() + stream.seek(current_position) + return size + except (AttributeError, OSError): + return None + + +def file_size_exceeds_limit(size: Optional[int], max_size_mb: Optional[int]) -> bool: + return ( + size is not None + and max_size_mb is not None + and size > max_size_mb * BYTES_PER_MB + ) + + +def file_count_exceeds_limit(files: Any, max_count: Optional[int]) -> bool: + return isinstance(files, list) and max_count is not None and len(files) > max_count diff --git a/backend/open_webui/utils/knowledge_files.py b/backend/open_webui/utils/knowledge_files.py new file mode 100644 index 000000000000..3cd6158e6945 --- /dev/null +++ b/backend/open_webui/utils/knowledge_files.py @@ -0,0 +1,17 @@ +from typing import Any + + +def get_knowledge_file_ids(knowledge: Any) -> list[str]: + data = getattr(knowledge, "data", None) or {} + file_ids = data.get("file_ids", []) + return file_ids if isinstance(file_ids, list) else [] + + +def user_owns_file_or_is_admin(user: Any, file: Any) -> bool: + return getattr(user, "role", None) == "admin" or getattr( + file, "user_id", None + ) == getattr(user, "id", None) + + +def file_id_in_knowledge(knowledge: Any, file_id: str) -> bool: + return file_id in get_knowledge_file_ids(knowledge) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index b1e69db2640f..089dfb86961b 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -82,6 +82,7 @@ process_filter_functions, ) from open_webui.utils.code_interpreter import execute_code_jupyter +from open_webui.utils.file_limits import file_count_exceeds_limit from open_webui.tasks import create_task @@ -721,6 +722,17 @@ async def process_chat_payload(request, form_data, user, metadata, model): form_data = apply_params_to_form_data(form_data, model) log.debug(f"form_data: {form_data}") + if file_count_exceeds_limit( + metadata.get("files"), request.app.state.config.FILE_MAX_COUNT + ): + raise HTTPException( + status_code=400, + detail=( + "You can only chat with a maximum of " + f"{request.app.state.config.FILE_MAX_COUNT} file(s) at a time." + ), + ) + event_emitter = get_event_emitter(metadata) event_call = get_event_call(metadata) diff --git a/src/lib/components/chat/Artifacts.svelte b/src/lib/components/chat/Artifacts.svelte index a6caa4210649..4dafea55ea84 100644 --- a/src/lib/components/chat/Artifacts.svelte +++ b/src/lib/components/chat/Artifacts.svelte @@ -335,7 +335,7 @@ title="Content" srcdoc={contents[selectedContentIdx].content} class="w-full border-0 h-full rounded-none" - sandbox="allow-scripts allow-downloads{($settings?.iframeSandboxAllowForms ?? false) + sandbox="allow-scripts{($settings?.iframeSandboxAllowForms ?? false) ? ' allow-forms' : ''}{($settings?.iframeSandboxAllowSameOrigin ?? false) ? ' allow-same-origin' diff --git a/src/lib/components/chat/Messages/Markdown/HTMLToken.svelte b/src/lib/components/chat/Messages/Markdown/HTMLToken.svelte index f917badac9e2..a5873e270c8c 100644 --- a/src/lib/components/chat/Messages/Markdown/HTMLToken.svelte +++ b/src/lib/components/chat/Messages/Markdown/HTMLToken.svelte @@ -93,7 +93,7 @@ src={`${WEBUI_BASE_URL}/api/v1/files/${fileId}/content/html`} title="Content" frameborder="0" - sandbox="allow-scripts allow-downloads{($settings?.iframeSandboxAllowForms ?? false) + sandbox="allow-scripts{($settings?.iframeSandboxAllowForms ?? false) ? ' allow-forms' : ''}{($settings?.iframeSandboxAllowSameOrigin ?? false) ? ' allow-same-origin' : ''}" referrerpolicy="strict-origin-when-cross-origin" diff --git a/tests/test_file_limit_regressions.py b/tests/test_file_limit_regressions.py new file mode 100644 index 000000000000..42b3d4a708b4 --- /dev/null +++ b/tests/test_file_limit_regressions.py @@ -0,0 +1,81 @@ +import ast +import importlib.util +from io import BytesIO +from pathlib import Path +from types import SimpleNamespace + + +ROOT = Path(__file__).resolve().parents[1] +HELPER_PATH = ROOT / "backend" / "open_webui" / "utils" / "file_limits.py" +FILES_ROUTER_PATH = ROOT / "backend" / "open_webui" / "routers" / "files.py" +MIDDLEWARE_PATH = ROOT / "backend" / "open_webui" / "utils" / "middleware.py" + + +def load_helper_module(): + spec = importlib.util.spec_from_file_location("file_limits", HELPER_PATH) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def get_function_source(path, function_name): + source = path.read_text() + tree = ast.parse(source) + for node in tree.body: + if ( + isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + and node.name == function_name + ): + return "\n".join(source.splitlines()[node.lineno - 1 : node.end_lineno]) + raise AssertionError(f"Function not found: {function_name}") + + +def test_upload_file_size_uses_size_attribute_when_present(): + helper = load_helper_module() + upload = SimpleNamespace(size=123, file=BytesIO(b"ignored")) + + assert helper.get_upload_file_size(upload) == 123 + + +def test_upload_file_size_falls_back_to_stream_without_changing_position(): + helper = load_helper_module() + stream = BytesIO(b"abcdef") + stream.seek(2) + upload = SimpleNamespace(file=stream) + + assert helper.get_upload_file_size(upload) == 6 + assert stream.tell() == 2 + + +def test_file_size_and_count_limits_are_enforced_only_when_configured(): + helper = load_helper_module() + + assert helper.file_size_exceeds_limit(2 * helper.BYTES_PER_MB + 1, 2) + assert not helper.file_size_exceeds_limit(2 * helper.BYTES_PER_MB, 2) + assert not helper.file_size_exceeds_limit(100, None) + assert helper.file_count_exceeds_limit([{}, {}, {}], 2) + assert not helper.file_count_exceeds_limit([{}, {}], 2) + assert not helper.file_count_exceeds_limit([{}, {}, {}], None) + + +def test_upload_route_checks_size_before_storage_write(): + upload_source = get_function_source(FILES_ROUTER_PATH, "upload_file") + + size_check = upload_source.index("file_size_exceeds_limit(file_size, max_file_size)") + storage_write = upload_source.index("Storage.upload_file") + assert size_check < storage_write + assert "HTTP_413_REQUEST_ENTITY_TOO_LARGE" in upload_source + + +def test_chat_payload_checks_client_file_count_before_model_knowledge_append(): + process_source = get_function_source(MIDDLEWARE_PATH, "process_chat_payload") + + count_check = process_source.index("file_count_exceeds_limit(") + model_knowledge = process_source.index("model_knowledge =") + assert count_check < model_knowledge + + +if __name__ == "__main__": + for name, fn in sorted(globals().items()): + if name.startswith("test_") and callable(fn): + fn() diff --git a/tests/test_iframe_sandbox_regressions.py b/tests/test_iframe_sandbox_regressions.py new file mode 100644 index 000000000000..670616057482 --- /dev/null +++ b/tests/test_iframe_sandbox_regressions.py @@ -0,0 +1,29 @@ +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +SCRIPTED_IFRAME_FILES = [ + ROOT / "src" / "lib" / "components" / "chat" / "Artifacts.svelte", + ( + ROOT + / "src" + / "lib" + / "components" + / "chat" + / "Messages" + / "Markdown" + / "HTMLToken.svelte" + ), +] + + +def test_script_enabled_chat_iframes_do_not_allow_downloads_by_default(): + for path in SCRIPTED_IFRAME_FILES: + source = path.read_text() + assert "allow-downloads" not in source + + +if __name__ == "__main__": + for name, fn in sorted(globals().items()): + if name.startswith("test_") and callable(fn): + fn() diff --git a/tests/test_knowledge_file_regressions.py b/tests/test_knowledge_file_regressions.py new file mode 100644 index 000000000000..faa6873145cf --- /dev/null +++ b/tests/test_knowledge_file_regressions.py @@ -0,0 +1,86 @@ +import ast +import importlib.util +from pathlib import Path +from types import SimpleNamespace + + +ROOT = Path(__file__).resolve().parents[1] +HELPER_PATH = ROOT / "backend" / "open_webui" / "utils" / "knowledge_files.py" +ROUTER_PATH = ROOT / "backend" / "open_webui" / "routers" / "knowledge.py" + + +def load_helper_module(): + spec = importlib.util.spec_from_file_location("knowledge_files", HELPER_PATH) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def get_function_source(function_name): + source = ROUTER_PATH.read_text() + tree = ast.parse(source) + for node in tree.body: + if ( + isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + and node.name == function_name + ): + return "\n".join(source.splitlines()[node.lineno - 1 : node.end_lineno]) + raise AssertionError(f"Function not found: {function_name}") + + +def test_file_access_helper_only_allows_owner_or_admin(): + helper = load_helper_module() + owner = SimpleNamespace(id="user-1", role="user") + other = SimpleNamespace(id="user-2", role="user") + admin = SimpleNamespace(id="admin", role="admin") + file = SimpleNamespace(user_id="user-1") + + assert helper.user_owns_file_or_is_admin(owner, file) + assert helper.user_owns_file_or_is_admin(admin, file) + assert not helper.user_owns_file_or_is_admin(other, file) + + +def test_file_id_membership_helper_rejects_missing_and_malformed_lists(): + helper = load_helper_module() + + assert helper.file_id_in_knowledge( + SimpleNamespace(data={"file_ids": ["file-a"]}), "file-a" + ) + assert not helper.file_id_in_knowledge( + SimpleNamespace(data={"file_ids": ["file-a"]}), "file-b" + ) + assert not helper.file_id_in_knowledge(SimpleNamespace(data={}), "file-a") + assert not helper.file_id_in_knowledge( + SimpleNamespace(data={"file_ids": "file-a"}), "file-a" + ) + + +def test_knowledge_remove_does_not_delete_global_file_record(): + remove_source = get_function_source("remove_file_from_knowledge_by_id") + + assert "Files.delete_file_by_id" not in remove_source + + +def test_knowledge_file_routes_validate_ownership_and_membership_before_processing(): + add_source = get_function_source("add_file_to_knowledge_by_id") + batch_source = get_function_source("add_files_to_knowledge_batch") + update_source = get_function_source("update_file_from_knowledge_by_id") + remove_source = get_function_source("remove_file_from_knowledge_by_id") + + assert "user_owns_file_or_is_admin(user, file)" in add_source + assert "user_owns_file_or_is_admin(user, file)" in batch_source + + for source in (update_source, remove_source): + membership_check = source.index( + "file_id_in_knowledge(knowledge, form_data.file_id)" + ) + file_lookup = source.index("Files.get_file_by_id(form_data.file_id)") + vector_delete = source.index("VECTOR_DB_CLIENT.delete") + assert membership_check < file_lookup + assert membership_check < vector_delete + + +if __name__ == "__main__": + for name, fn in sorted(globals().items()): + if name.startswith("test_") and callable(fn): + fn()