Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions backend/open_webui/routers/files.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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 ""
Expand Down Expand Up @@ -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(
Expand Down
40 changes: 26 additions & 14 deletions backend/open_webui/routers/knowledge.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down
36 changes: 36 additions & 0 deletions backend/open_webui/utils/file_limits.py
Original file line number Diff line number Diff line change
@@ -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
17 changes: 17 additions & 0 deletions backend/open_webui/utils/knowledge_files.py
Original file line number Diff line number Diff line change
@@ -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)
12 changes: 12 additions & 0 deletions backend/open_webui/utils/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)

Expand Down
2 changes: 1 addition & 1 deletion src/lib/components/chat/Artifacts.svelte
Original file line number Diff line number Diff line change
Expand Up @@ -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'
Expand Down
2 changes: 1 addition & 1 deletion src/lib/components/chat/Messages/Markdown/HTMLToken.svelte
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
81 changes: 81 additions & 0 deletions tests/test_file_limit_regressions.py
Original file line number Diff line number Diff line change
@@ -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()
29 changes: 29 additions & 0 deletions tests/test_iframe_sandbox_regressions.py
Original file line number Diff line number Diff line change
@@ -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()
Loading