212 lines
9.3 KiB
Python
212 lines
9.3 KiB
Python
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import shutil
|
|
import tempfile
|
|
import zipfile
|
|
from dataclasses import dataclass
|
|
from pathlib import Path, PurePosixPath
|
|
from uuid import UUID
|
|
|
|
from fastapi import HTTPException
|
|
|
|
MAX_ARCHIVE_BYTES = 256 * 1024 * 1024
|
|
MAX_ENTRIES = 10_000
|
|
MAX_COMPRESSION_RATIO = 200
|
|
MAX_METADATA_ENTRY_BYTES = 16 * 1024 * 1024
|
|
_CHUNK_SIZE = 1024 * 1024
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class StagedBlob:
|
|
archive_path: str
|
|
staging_path: Path
|
|
size: int
|
|
sha256: str
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ParsedArchive:
|
|
backup_id: UUID
|
|
archive_sha256: str
|
|
entities: dict[str, list[dict]]
|
|
files: dict[str, StagedBlob]
|
|
staging_dir: Path
|
|
|
|
|
|
|
|
def backup_error(code: str, message: str, status_code: int = 422) -> HTTPException:
|
|
return HTTPException(status_code, {"code": code, "message": message})
|
|
|
|
|
|
|
|
def canonical_json(value) -> bytes:
|
|
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode()
|
|
|
|
|
|
|
|
def build_archive_to_path(
|
|
output_path: Path, backup_id: UUID, entities: dict[str, list[dict]], files: dict[str, Path]
|
|
) -> None:
|
|
entity_entries = {f"data/{name}.json": canonical_json(rows) for name, rows in entities.items()}
|
|
checksums = {name: hashlib.sha256(data).hexdigest() for name, data in entity_entries.items()}
|
|
for name, path in files.items():
|
|
digest = hashlib.sha256()
|
|
with path.open("rb") as source:
|
|
while chunk := source.read(_CHUNK_SIZE):
|
|
digest.update(chunk)
|
|
checksums[name] = digest.hexdigest()
|
|
manifest = {
|
|
"format": "dodo-backup", "version": 2, "backup_id": str(backup_id),
|
|
"entities": {name: len(rows) for name, rows in entities.items()}, "checksums": checksums,
|
|
}
|
|
with zipfile.ZipFile(output_path, "w", compression=zipfile.ZIP_DEFLATED, allowZip64=True) as archive:
|
|
archive.writestr("manifest.json", canonical_json(manifest))
|
|
for name, data in entity_entries.items():
|
|
archive.writestr(name, data)
|
|
for name, path in files.items():
|
|
archive.write(path, name)
|
|
|
|
|
|
|
|
def build_archive(backup_id: UUID, entities: dict[str, list[dict]], files: dict[str, bytes]) -> bytes:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
root = Path(directory)
|
|
file_paths = {}
|
|
for index, (name, data) in enumerate(files.items()):
|
|
path = root / str(index)
|
|
path.write_bytes(data)
|
|
file_paths[name] = path
|
|
output = root / "backup.zip"
|
|
build_archive_to_path(output, backup_id, entities, file_paths)
|
|
return output.read_bytes()
|
|
|
|
|
|
|
|
def _safe_name(name: str) -> bool:
|
|
path = PurePosixPath(name)
|
|
return bool(name) and not name.startswith("/") and "\\" not in name and ".." not in path.parts
|
|
|
|
|
|
|
|
def _read_metadata(archive: zipfile.ZipFile, info: zipfile.ZipInfo) -> bytes:
|
|
if info.file_size > MAX_METADATA_ENTRY_BYTES:
|
|
raise backup_error("backup_size_invalid", "备份元数据过大")
|
|
with archive.open(info) as source:
|
|
data = source.read(MAX_METADATA_ENTRY_BYTES + 1)
|
|
if len(data) > MAX_METADATA_ENTRY_BYTES:
|
|
raise backup_error("backup_size_invalid", "备份元数据过大")
|
|
return data
|
|
|
|
|
|
|
|
def _stream_blob(archive: zipfile.ZipFile, info: zipfile.ZipInfo, destination: Path) -> tuple[int, str]:
|
|
digest = hashlib.sha256()
|
|
size = 0
|
|
with archive.open(info) as source, destination.open("xb") as output:
|
|
while chunk := source.read(_CHUNK_SIZE):
|
|
size += len(chunk)
|
|
digest.update(chunk)
|
|
output.write(chunk)
|
|
return size, digest.hexdigest()
|
|
|
|
|
|
|
|
def parse_archive_path(path: Path, *, max_archive_bytes: int = MAX_ARCHIVE_BYTES) -> ParsedArchive:
|
|
size = path.stat().st_size
|
|
if size <= 0 or size > max_archive_bytes:
|
|
raise backup_error("backup_size_invalid", "备份文件大小无效")
|
|
digest = hashlib.sha256()
|
|
with path.open("rb") as source:
|
|
while chunk := source.read(_CHUNK_SIZE):
|
|
digest.update(chunk)
|
|
staging_dir = Path(tempfile.mkdtemp(prefix="dodo-backup-blobs-"))
|
|
try:
|
|
with zipfile.ZipFile(path) as archive:
|
|
infos = archive.infolist()
|
|
names = [item.filename for item in infos]
|
|
if len(infos) > MAX_ENTRIES:
|
|
raise backup_error("backup_too_many_entries", "备份条目过多")
|
|
if len(names) != len(set(names)):
|
|
raise backup_error("backup_duplicate_entry", "备份包含重复条目")
|
|
if any(not _safe_name(item.filename) or item.is_dir() or item.flag_bits & 1 for item in infos):
|
|
raise backup_error("backup_unsafe_path", "备份包含不安全路径")
|
|
if "manifest.json" not in names:
|
|
raise backup_error("backup_manifest_missing", "备份缺少 manifest")
|
|
if sum(item.file_size for item in infos) > max_archive_bytes:
|
|
raise backup_error("backup_size_invalid", "备份解压后过大")
|
|
if any(item.file_size and item.compress_size == 0 for item in infos):
|
|
raise backup_error("backup_compression_invalid", "备份压缩比异常")
|
|
if any(item.compress_size and item.file_size / item.compress_size > MAX_COMPRESSION_RATIO for item in infos):
|
|
raise backup_error("backup_compression_invalid", "备份压缩比异常")
|
|
by_name = {item.filename: item for item in infos}
|
|
manifest = json.loads(_read_metadata(archive, by_name["manifest.json"]))
|
|
if manifest.get("format") != "dodo-backup" or manifest.get("version") != 2:
|
|
raise backup_error("backup_version_unsupported", "不支持的备份版本")
|
|
backup_id = UUID(manifest["backup_id"])
|
|
checksums = manifest["checksums"]
|
|
declared_entities = manifest["entities"]
|
|
if not isinstance(checksums, dict) or not isinstance(declared_entities, dict):
|
|
raise TypeError
|
|
content_names = set(names) - {"manifest.json"}
|
|
if set(checksums) != content_names:
|
|
raise backup_error("backup_manifest_mismatch", "manifest 与 ZIP 条目不一致")
|
|
# Calendar subscriptions are an additive backup-v2 entity. Accept archives
|
|
# produced by older helpers that checksum the new file but omit its count.
|
|
optional_entities = {"calendar_subscriptions"}
|
|
undeclared_optional = {
|
|
f"data/{name}.json" for name in optional_entities - set(declared_entities)
|
|
}
|
|
if set(declared_entities) != {
|
|
name[5:-5]
|
|
for name in content_names - undeclared_optional
|
|
if name.startswith("data/") and name.endswith(".json")
|
|
}:
|
|
raise backup_error("backup_manifest_mismatch", "manifest 实体清单不一致")
|
|
entities: dict[str, list[dict]] = {}
|
|
files: dict[str, StagedBlob] = {}
|
|
for index, name in enumerate(sorted(content_names)):
|
|
info = by_name[name]
|
|
if name.startswith("data/") and name.endswith(".json"):
|
|
data = _read_metadata(archive, info)
|
|
actual_digest = hashlib.sha256(data).hexdigest()
|
|
value = json.loads(data)
|
|
if not isinstance(value, list) or any(not isinstance(row, dict) for row in value):
|
|
raise ValueError
|
|
entities[name[5:-5]] = value
|
|
else:
|
|
blob_path = staging_dir / str(index)
|
|
blob_size, actual_digest = _stream_blob(archive, info, blob_path)
|
|
files[name] = StagedBlob(name, blob_path, blob_size, actual_digest)
|
|
if not isinstance(checksums[name], str) or actual_digest != checksums[name]:
|
|
raise backup_error("backup_checksum_mismatch", "备份校验和不匹配")
|
|
if set(declared_entities) != set(entities):
|
|
required_declared = set(entities) - optional_entities
|
|
if set(declared_entities) != required_declared or any(
|
|
entities.get(name) for name in optional_entities - set(declared_entities)
|
|
):
|
|
raise backup_error("backup_manifest_mismatch", "manifest 实体清单不一致")
|
|
if any(type(count) is not int or count < 0 or count != len(entities[name]) for name, count in declared_entities.items()):
|
|
raise backup_error("backup_manifest_mismatch", "manifest 实体数量不一致")
|
|
except HTTPException:
|
|
shutil.rmtree(staging_dir, ignore_errors=True)
|
|
raise
|
|
except (zipfile.BadZipFile, OSError) as exc:
|
|
shutil.rmtree(staging_dir, ignore_errors=True)
|
|
raise backup_error("backup_invalid_zip", "无效的 ZIP 备份") from exc
|
|
except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
|
|
shutil.rmtree(staging_dir, ignore_errors=True)
|
|
raise backup_error("backup_manifest_invalid", "无效的备份 manifest 或数据") from exc
|
|
return ParsedArchive(backup_id, digest.hexdigest(), entities, files, staging_dir)
|
|
|
|
|
|
|
|
def parse_archive(content: bytes) -> ParsedArchive:
|
|
if not content or len(content) > MAX_ARCHIVE_BYTES:
|
|
raise backup_error("backup_size_invalid", "备份文件大小无效")
|
|
with tempfile.NamedTemporaryFile() as staged:
|
|
staged.write(content)
|
|
staged.flush()
|
|
return parse_archive_path(Path(staged.name))
|