Skip to content
Merged
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
112 changes: 112 additions & 0 deletions tests/test_upload_size_guard.py
Original file line number Diff line number Diff line change
@@ -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")
29 changes: 25 additions & 4 deletions web/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
17 changes: 12 additions & 5 deletions web/security.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading