B-42: 导入功能加固——逐工作表确认、流式上传与诊断证据
This commit is contained in:
@@ -0,0 +1,275 @@
|
||||
"""Streaming multipart/form-data parser for the stdlib HTTP server.
|
||||
|
||||
The upload body is consumed in chunks from the request stream and never
|
||||
assembled into memory as a whole. File parts are streamed straight to a temp
|
||||
file on disk with an incremental size cap; text fields are capped at a small
|
||||
field size. File bytes are preserved exactly (no trailing-byte trimming), the
|
||||
first file part is the only one accepted, and both the extension and the file
|
||||
signature are validated by the caller.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
import re
|
||||
from typing import BinaryIO
|
||||
|
||||
MAX_FIELD_BYTES = 16 * 1024
|
||||
MAX_HEADER_BYTES = 32 * 1024
|
||||
CHUNK_SIZE = 64 * 1024
|
||||
# .xlsx files are ZIP containers (PK\x03\x04 local header or an empty-archive
|
||||
# PK\x05\x06); .xls files are OLE2 compound documents.
|
||||
XLSX_MAGIC = (b"PK\x03\x04", b"PK\x05\x06")
|
||||
XLS_MAGIC = (b"\xd0\xcf\x11\xe0\xa1\xb1\x1a\xe1",)
|
||||
|
||||
|
||||
class MultipartError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UploadedFile:
|
||||
filename: str
|
||||
path: Path
|
||||
size: int
|
||||
|
||||
|
||||
def valid_file_signature(path: Path, suffix: str) -> bool:
|
||||
"""True when the file's leading bytes match its declared extension."""
|
||||
with path.open("rb") as handle:
|
||||
head = handle.read(8)
|
||||
if suffix == ".xlsx":
|
||||
return any(head.startswith(magic) for magic in XLSX_MAGIC)
|
||||
if suffix == ".xls":
|
||||
return any(head.startswith(magic) for magic in XLS_MAGIC)
|
||||
return False
|
||||
|
||||
|
||||
class _MultipartStream:
|
||||
def __init__(
|
||||
self,
|
||||
rfile: BinaryIO,
|
||||
content_length: int,
|
||||
content_type: str,
|
||||
*,
|
||||
max_body_bytes: int,
|
||||
) -> None:
|
||||
lowered = content_type.lower()
|
||||
if not lowered.startswith("multipart/form-data"):
|
||||
raise MultipartError("上传请求必须是 multipart/form-data。")
|
||||
match = re.search(r"boundary=(?:\"([^\"]+)\"|([^;]+))", content_type)
|
||||
if not match:
|
||||
raise MultipartError("上传请求缺少文件边界。")
|
||||
boundary = (match.group(1) or match.group(2)).strip().encode("utf-8")
|
||||
if not boundary or len(boundary) > 200:
|
||||
raise MultipartError("上传请求的文件边界无效。")
|
||||
if content_length <= 0:
|
||||
raise MultipartError("上传请求为空。")
|
||||
if content_length > max_body_bytes:
|
||||
raise MultipartError(f"上传请求超过 {max_body_bytes // (1024 * 1024)} MB 限制。")
|
||||
self.rfile = rfile
|
||||
self.remaining = content_length
|
||||
self.total_read = 0
|
||||
self.buffer = bytearray()
|
||||
self.boundary = boundary
|
||||
self._fill()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Low-level stream helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _fill(self) -> None:
|
||||
if self.total_read >= self.remaining:
|
||||
return
|
||||
want = min(CHUNK_SIZE, self.remaining - self.total_read)
|
||||
data = self.rfile.read(want)
|
||||
if not data:
|
||||
# Client closed early; clamp remaining so every helper sees EOF.
|
||||
self.remaining = self.total_read
|
||||
return
|
||||
self.total_read += len(data)
|
||||
self.buffer.extend(data)
|
||||
|
||||
def _take(self, count: int) -> bytes:
|
||||
while len(self.buffer) < count:
|
||||
if self.total_read >= self.remaining:
|
||||
raise MultipartError("multipart 请求体不完整。")
|
||||
self._fill()
|
||||
out = bytes(self.buffer[:count])
|
||||
del self.buffer[:count]
|
||||
return out
|
||||
|
||||
def _read_line(self) -> bytes:
|
||||
while True:
|
||||
index = self.buffer.find(b"\n")
|
||||
if index >= 0:
|
||||
line = bytes(self.buffer[: index + 1])
|
||||
del self.buffer[: index + 1]
|
||||
return line
|
||||
if self.total_read >= self.remaining:
|
||||
if not self.buffer:
|
||||
raise MultipartError("multipart 请求体不完整。")
|
||||
line = bytes(self.buffer)
|
||||
self.buffer.clear()
|
||||
return line
|
||||
self._fill()
|
||||
|
||||
def _skip_preamble(self) -> None:
|
||||
marker = b"--" + self.boundary
|
||||
while True:
|
||||
index = self.buffer.find(marker)
|
||||
if index >= 0:
|
||||
del self.buffer[: index + len(marker)]
|
||||
return
|
||||
keep = len(marker) - 1
|
||||
if len(self.buffer) > keep:
|
||||
del self.buffer[: len(self.buffer) - keep]
|
||||
if self.total_read >= self.remaining:
|
||||
raise MultipartError("上传请求中没有找到文件。")
|
||||
self._fill()
|
||||
|
||||
def _iter_content(self, boundary: bytes):
|
||||
"""Yield content bytes up to (excluding) the ``\\r\\n--boundary`` marker."""
|
||||
marker = b"\r\n--" + boundary
|
||||
keep = len(marker) - 1
|
||||
while True:
|
||||
index = self.buffer.find(marker)
|
||||
if index >= 0:
|
||||
content = bytes(self.buffer[:index])
|
||||
del self.buffer[:index]
|
||||
if content:
|
||||
yield content
|
||||
return
|
||||
if len(self.buffer) > keep:
|
||||
safe = len(self.buffer) - keep
|
||||
content = bytes(self.buffer[:safe])
|
||||
del self.buffer[:safe]
|
||||
if content:
|
||||
yield content
|
||||
if self.total_read >= self.remaining:
|
||||
raise MultipartError("multipart 请求体缺少结束边界。")
|
||||
self._fill()
|
||||
|
||||
def _after_content_boundary(self) -> str:
|
||||
"""Consume the content-ending marker; return 'part' or 'end'."""
|
||||
marker = b"\r\n--" + self.boundary
|
||||
del self.buffer[: len(marker)]
|
||||
indicator = self._take(2)
|
||||
if indicator == b"--":
|
||||
if self.buffer.startswith(b"\r\n"):
|
||||
del self.buffer[:2]
|
||||
elif self.buffer.startswith(b"\n"):
|
||||
del self.buffer[:1]
|
||||
return "end"
|
||||
if indicator != b"\r\n":
|
||||
raise MultipartError("上传请求格式无效。")
|
||||
return "part"
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Part parsing
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _read_headers(self) -> bytes:
|
||||
total = 0
|
||||
lines: list[bytes] = []
|
||||
while True:
|
||||
line = self._read_line()
|
||||
total += len(line)
|
||||
if total > MAX_HEADER_BYTES:
|
||||
raise MultipartError("multipart 头部过长。")
|
||||
if line in (b"\r\n", b"\n"):
|
||||
return b"".join(lines)
|
||||
lines.append(line)
|
||||
|
||||
@staticmethod
|
||||
def _parse_disposition(header: bytes) -> tuple[str | None, str | None]:
|
||||
name = None
|
||||
filename = None
|
||||
name_match = re.search(br'name="([^"]*)"', header)
|
||||
if name_match:
|
||||
name = name_match.group(1).decode("utf-8", errors="replace")
|
||||
filename_match = re.search(br'filename="([^"]*)"', header)
|
||||
if filename_match:
|
||||
filename = filename_match.group(1).decode("utf-8", errors="replace")
|
||||
return name, filename
|
||||
|
||||
def parse(self, work_dir: Path) -> tuple[dict[str, str], UploadedFile | None]:
|
||||
import tempfile
|
||||
|
||||
self._skip_preamble()
|
||||
next_bytes = self._take(2)
|
||||
if next_bytes == b"--":
|
||||
# Empty multipart body: opening boundary immediately closes.
|
||||
return {}, None
|
||||
if next_bytes != b"\r\n":
|
||||
raise MultipartError("上传请求格式无效。")
|
||||
|
||||
fields: dict[str, str] = {}
|
||||
uploaded: UploadedFile | None = None
|
||||
while True:
|
||||
header = self._read_headers()
|
||||
name, filename = self._parse_disposition(header)
|
||||
if name is None:
|
||||
raise MultipartError("上传请求缺少字段名。")
|
||||
|
||||
if filename is not None:
|
||||
if uploaded is not None:
|
||||
raise MultipartError("一次只能上传一个文件。")
|
||||
with tempfile.NamedTemporaryFile(
|
||||
dir=str(work_dir), prefix="upload-", delete=False
|
||||
) as sink:
|
||||
temp_path = Path(sink.name)
|
||||
size = 0
|
||||
try:
|
||||
for chunk in self._iter_content(self.boundary):
|
||||
size += len(chunk)
|
||||
sink.write(chunk)
|
||||
except Exception:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
uploaded = UploadedFile(
|
||||
filename=filename, path=temp_path, size=size
|
||||
)
|
||||
if self._after_content_boundary() == "end":
|
||||
break
|
||||
else:
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
for chunk in self._iter_content(self.boundary):
|
||||
total += len(chunk)
|
||||
if total > MAX_FIELD_BYTES:
|
||||
raise MultipartError("表单字段超过大小限制。")
|
||||
chunks.append(chunk)
|
||||
fields[name] = b"".join(chunks).decode(
|
||||
"utf-8", errors="replace"
|
||||
).strip()
|
||||
if self._after_content_boundary() == "end":
|
||||
break
|
||||
return fields, uploaded
|
||||
|
||||
|
||||
def parse_upload(
|
||||
rfile: BinaryIO,
|
||||
content_length: int,
|
||||
content_type: str,
|
||||
work_dir: str | Path,
|
||||
*,
|
||||
max_file_bytes: int,
|
||||
) -> tuple[dict[str, str], UploadedFile | None]:
|
||||
"""Stream a multipart upload; returns ``(fields, uploaded_file)``.
|
||||
|
||||
Raises :class:`MultipartError` for malformed or oversized bodies. The
|
||||
caller owns the uploaded temp file and must remove it when done.
|
||||
"""
|
||||
work = Path(work_dir)
|
||||
work.mkdir(parents=True, exist_ok=True)
|
||||
stream = _MultipartStream(
|
||||
rfile, content_length, content_type, max_body_bytes=max_file_bytes + MAX_FIELD_BYTES
|
||||
)
|
||||
fields, uploaded = stream.parse(work)
|
||||
if uploaded is None:
|
||||
raise MultipartError("上传请求中没有找到文件。")
|
||||
if uploaded.size > max_file_bytes:
|
||||
raise MultipartError(f"文件超过 {max_file_bytes // (1024 * 1024)} MB 限制。")
|
||||
return fields, uploaded
|
||||
Reference in New Issue
Block a user