diff --git a/tests/test_upload_size_guard.py b/tests/test_upload_size_guard.py new file mode 100644 index 0000000..420ed4c --- /dev/null +++ b/tests/test_upload_size_guard.py @@ -0,0 +1,112 @@ +import asyncio +import importlib +import os +import sys + +from fastapi import HTTPException + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + + +UPLOAD_ENV = ( + "MAX_UPLOAD_MB", + "API_KEY", + "API_KEYS", + "ALLOW_ANON_API", + "TRUSTED_HOSTS", + "TRUST_PROXY_HEADERS", +) + + +def _load_app(monkeypatch): + for key in UPLOAD_ENV: + monkeypatch.delenv(key, raising=False) + monkeypatch.setenv("MAX_UPLOAD_MB", "1") + monkeypatch.setenv("ALLOW_ANON_API", "true") + for module in ("web.app", "web.security"): + sys.modules.pop(module, None) + return importlib.import_module("web.app") + + +def _chunked_asgi(app, body: bytes, chunk_size: int = 1024): + """ASGI wrapper: strips Content-Length and feeds the body in chunks, + like a chunked-transfer-encoding upload.""" + + async def wrapped(scope, receive, send): + if scope.get("type") != "http": + await app(scope, receive, send) + return + scope = {**scope, "client": ("198.51.100.1", 50000)} + headers = [(k, v) for k, v in scope["headers"] if k != b"content-length"] + scope = {**scope, "headers": headers} + + sent = 0 + + async def chunked_receive(): + nonlocal sent + if sent < len(body): + piece = body[sent:sent + chunk_size] + sent += len(piece) + return {"type": "http.request", "body": piece, "more_body": True} + return {"type": "http.request", "body": b"", "more_body": False} + + await app(scope, chunked_receive, send) + + return wrapped + + +def _multipart_body(filename: str, content: bytes) -> bytes: + boundary = "----prism-test-boundary" + body = ( + f"--{boundary}\r\n" + f'Content-Disposition: form-data; name="file"; filename="{filename}"\r\n' + "Content-Type: image/png\r\n" + "\r\n" + ).encode("ascii") + content + f"\r\n--{boundary}--\r\n".encode("ascii") + return body + + +def test_metadata_rejects_oversized_chunked_upload_without_content_length(monkeypatch): + import httpx + + from web import security + + app_mod = _load_app(monkeypatch) + max_bytes = security.MAX_UPLOAD_BYTES + + content = b"x" * (max_bytes + 1024) + body = _multipart_body("big.png", content) + + transport = httpx.ASGITransport(app=_chunked_asgi(app_mod.app, body)) + client = httpx.AsyncClient(transport=transport, base_url="http://testserver") + + async def run(): + async with client: + return await client.post( + "/api/metadata", + headers={ + "content-type": "multipart/form-data; boundary=----prism-test-boundary", + }, + ) + + response = asyncio.new_event_loop().run_until_complete(run()) + + assert response.status_code == 413, f"expected 413, got {response.status_code}: {response.text[:200]}" + + +def test_check_upload_size_rejects_non_numeric_content_length(): + from web import security + + class FakeRequest: + headers = {"content-length": "abc"} + + try: + asyncio.new_event_loop().run_until_complete(security.check_upload_size(FakeRequest())) + except HTTPException as exc: + assert exc.status_code == 400 + except ValueError: + raise AssertionError( + "non-numeric Content-Length raised ValueError (500) instead of HTTP 400" + ) + else: + raise AssertionError("non-numeric Content-Length did not raise HTTPException") diff --git a/web/app.py b/web/app.py index f6ab86c..99a31f0 100644 --- a/web/app.py +++ b/web/app.py @@ -13,7 +13,7 @@ import requests as _requests import logging -from fastapi import FastAPI, WebSocket, WebSocketDisconnect, UploadFile, File, Depends, Request +from fastapi import FastAPI, WebSocket, WebSocketDisconnect, UploadFile, File, Depends, HTTPException, Request from fastapi.responses import FileResponse, JSONResponse, Response from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel @@ -1233,9 +1233,30 @@ async def extract_metadata_endpoint(request: Request, file: UploadFile = File(.. loop = asyncio.get_running_loop() def _spool() -> str: - with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: - shutil.copyfileobj(file.file, tmp) - return tmp.name + from web.security import MAX_UPLOAD_BYTES + tmp_path = None + try: + with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: + tmp_path = tmp.name + written = 0 + while True: + chunk = file.file.read(64 * 1024) + if not chunk: + break + written += tmp.write(chunk) + if written > MAX_UPLOAD_BYTES: + raise HTTPException( + status_code=413, + detail=f"File too large. Max {MAX_UPLOAD_BYTES // (1024*1024)} MB allowed.", + ) + return tmp_path + except Exception: + if tmp_path and os.path.exists(tmp_path): + try: + os.unlink(tmp_path) + except OSError: + pass + raise tmp_path = await loop.run_in_executor(None, _spool) try: diff --git a/web/security.py b/web/security.py index 56a53a2..0c0f500 100644 --- a/web/security.py +++ b/web/security.py @@ -172,11 +172,18 @@ def validate_target(target: str) -> str: async def check_upload_size(request: Request) -> None: content_length = request.headers.get("content-length") - if content_length and int(content_length) > MAX_UPLOAD_BYTES: - raise HTTPException( - status_code=413, - detail=f"File too large. Max {MAX_UPLOAD_BYTES // (1024*1024)} MB allowed.", - ) + if content_length: + try: + length = int(content_length) + except ValueError: + raise HTTPException(status_code=400, detail="Invalid Content-Length header.") + if length < 0: + raise HTTPException(status_code=400, detail="Invalid Content-Length header.") + if length > MAX_UPLOAD_BYTES: + raise HTTPException( + status_code=413, + detail=f"File too large. Max {MAX_UPLOAD_BYTES // (1024*1024)} MB allowed.", + ) def validate_scan_id(scan_id: str) -> str: import uuid