Files

282 lines
12 KiB
Python

"""PNG encoding, CESV v1 obfuscation and bounded, non-overwriting file output."""
import hashlib
import io
import json
import logging
import ntpath
import os
from pathlib import Path
import secrets
import struct
import tempfile
import uuid
import numpy as np
from PIL import Image
from PIL.PngImagePlugin import PngInfo
VERSION = 1
HEADER = struct.Struct(">4sBBHQ32s")
PNG_SIGNATURE = b"\x89PNG\r\n\x1a\n"
MAX_PIXELS = 67_108_864
MAX_PNG_BYTES = 512 * 1024 * 1024
MAX_METADATA_BYTES = 16 * 1024 * 1024
CHUNK_BYTES = 1024 * 1024
LOG = logging.getLogger(__name__)
class SaveBatchError(RuntimeError):
"""A failed batch whose already-published files remain valid."""
def _noop():
pass
def validate_shape(shape):
if len(shape) != 4:
raise ValueError("IMAGE 必须为 [批次, 高, 宽, 通道] 四维张量。")
count, height, width, channels = shape
if any(not isinstance(value, (int, np.integer)) or value < 1 for value in shape):
raise ValueError("IMAGE 的批次、尺寸和通道数必须为正整数。")
if channels not in (1, 3, 4):
raise ValueError("仅支持灰度、RGB 和 RGBA 图片。")
if height * width > MAX_PIXELS:
raise ValueError(f"单张图片有 {height * width:,} 像素,超过 {MAX_PIXELS:,} 像素上限。")
return tuple(map(int, shape))
def to_pixels(array):
array = np.asarray(array)
if not np.issubdtype(array.dtype, np.floating) or not np.isfinite(array).all():
raise ValueError("IMAGE 必须只包含有限的浮点数。")
# Clamping first also avoids overflow for finite values far outside [0, 1].
scaled = np.clip(array, 0.0, 1.0)
scaled *= 255.0
pixels = scaled.astype(np.uint8)
return pixels[..., 0] if pixels.shape[-1] == 1 else pixels
def prepare_metadata(prompt=None, extra_pnginfo=None, disabled=False):
if disabled:
return None
if extra_pnginfo is not None and not isinstance(extra_pnginfo, dict):
raise ValueError("extra_pnginfo 必须是对象。")
entries = []
if prompt is not None:
entries.append(("prompt", prompt))
if extra_pnginfo is not None:
if prompt is not None and "prompt" in extra_pnginfo:
raise ValueError("extra_pnginfo 不能覆盖 prompt 元数据。")
entries.extend(extra_pnginfo.items())
metadata = PngInfo()
total = 0
encoder = json.JSONEncoder(ensure_ascii=True, allow_nan=False)
for key, value in entries:
if not isinstance(key, str):
raise ValueError("PNG 元数据名称必须是字符串。")
try:
keyword = key.encode("latin-1")
except UnicodeEncodeError as error:
raise ValueError("PNG 元数据名称必须能以 Latin-1 表示。") from error
if not 1 <= len(keyword) <= 79 or any(byte < 32 or 127 <= byte <= 160 for byte in keyword) or key.strip() != key or " " in key:
raise ValueError(f"PNG 元数据名称不合法:{key!r}")
total += len(keyword) + 1
pieces = []
for piece in encoder.iterencode(value):
total += len(piece) # ensure_ascii=True: one byte per character.
if total > MAX_METADATA_BYTES:
raise ValueError(f"PNG 元数据超过 {MAX_METADATA_BYTES:,} 字节上限。")
pieces.append(piece)
metadata.add_text(key, "".join(pieces))
return metadata
class LimitedBytesIO(io.BytesIO):
def write(self, data):
size = self.tell() + len(data)
if size > MAX_PNG_BYTES:
raise ValueError(f"PNG 编码至少需要 {size:,} 字节,超过 {MAX_PNG_BYTES:,} 字节上限。")
return super().write(data)
def encode_png(pixels, metadata):
buffer = LimitedBytesIO()
try:
with Image.fromarray(pixels) as image:
image.save(buffer, format="PNG", pnginfo=metadata, compress_level=4)
return buffer
except BaseException:
buffer.close()
raise
def write_all(stream, data, check_interrupt=_noop):
with memoryview(data) as view:
offset = 0
while offset < len(view):
check_interrupt()
written = stream.write(view[offset:])
if not isinstance(written, int) or not 0 < written <= len(view) - offset:
raise OSError("文件写入未完成(短写或无写入进展)。")
offset += written
def write_cesv(stream, png, check_interrupt=_noop, *, mask=None):
if not len(PNG_SIGNATURE) <= len(png) <= MAX_PNG_BYTES or png[:8] != PNG_SIGNATURE:
raise ValueError("PNG 载荷格式或长度无效。")
if mask is None:
mask = secrets.randbelow(255) + 1
if not isinstance(mask, int) or not 1 <= mask <= 255:
raise ValueError("XOR mask 必须为 1 到 255 的整数。")
digest = hashlib.sha256()
for start in range(0, len(png), CHUNK_BYTES):
check_interrupt()
digest.update(png[start:start + CHUNK_BYTES])
header = HEADER.pack(b"CESV", VERSION, mask, 0, len(png), digest.digest())
write_all(stream, header, check_interrupt)
scratch = np.empty(min(CHUNK_BYTES, len(png)), dtype=np.uint8)
for start in range(0, len(png), CHUNK_BYTES):
check_interrupt()
size = min(CHUNK_BYTES, len(png) - start)
np.bitwise_xor(np.frombuffer(png, dtype=np.uint8, count=size, offset=start), mask, out=scratch[:size])
write_all(stream, scratch[:size], check_interrupt)
def validate_prefix(prefix):
if not isinstance(prefix, str) or not prefix or ntpath.splitdrive(prefix)[0] or prefix.startswith(("/", "\\")):
raise ValueError("文件名前缀必须是 output 目录内的非空相对路径。")
parts = prefix.replace("\\", "/").split("/")
if any(not part or part in (".", "..") for part in parts) or any(ord(char) < 32 for char in prefix):
raise ValueError("文件名前缀含有空名称、控制字符或目录跳转。")
if os.name == "nt" and any(char in '<>:"|?*' for char in prefix):
raise ValueError("文件名前缀含有 Windows 不支持的字符;请先展开日期或节点变量。")
return str(Path(*parts))
def contained_folder(output_dir, folder):
root = Path(output_dir).resolve()
resolved = Path(folder).resolve()
if not resolved.is_relative_to(root):
raise ValueError("保存路径(含符号链接目标)必须位于 output 目录内。")
return resolved
def publish_no_replace(temporary, destination):
if os.name == "nt":
os.rename(temporary, destination) # Windows rename refuses an existing target.
else:
os.link(temporary, destination) # Atomic creation, never POSIX rename/replace.
def write_atomic(png, output_dir, folder, filename, counter, check_interrupt=_noop):
temporary = None
primary_error = None
warning = None
folder = contained_folder(output_dir, folder)
if not filename or Path(filename).name != filename or "/" in filename or "\\" in filename or counter < 1:
raise ValueError("最终文件名或流水号无效。")
try:
with tempfile.NamedTemporaryFile(mode="wb", buffering=0, prefix=".encrypt-save-", suffix=".part", dir=folder, delete=False) as stream:
temporary = Path(stream.name)
write_cesv(stream, png, check_interrupt)
stream.flush()
os.fsync(stream.fileno())
while True:
check_interrupt()
contained_folder(output_dir, folder)
final_name = f"{filename}_{counter:05}_.png.enc"
try:
publish_no_replace(temporary, folder / final_name)
break
except FileExistsError:
counter += 1
except BaseException as error:
primary_error = error
raise
finally:
if temporary is not None:
try:
temporary.unlink()
except FileNotFoundError:
pass # Windows publication already moved the temporary file.
except OSError as cleanup_error:
warning = f"临时文件 {temporary.name} 清理失败:{cleanup_error}"
LOG.warning(warning)
if primary_error is not None and hasattr(primary_error, "add_note"):
primary_error.add_note(warning)
return final_name, counter, warning
def save_batch(images, output_dir, filename_prefix, get_save_path, *, convert=np.asarray,
prompt=None, extra_pnginfo=None, disable_metadata=False, notify=None, check_interrupt=_noop):
count, height, width, _ = validate_shape(images.shape)
result = {"version": VERSION, "batch_id": uuid.uuid4().hex, "status": "running", "total": count,
"saved_count": 0, "files": [], "warnings": []}
index = None
interrupted = False
notification_failed = False
def checkpoint():
nonlocal interrupted
try:
check_interrupt()
except BaseException:
interrupted = True
raise
def emit(event, **data):
nonlocal notification_failed
if notify is None or notification_failed:
return
payload = {"version": VERSION, "batch_id": result["batch_id"], "total": count,
"event": event, "saved_count": result["saved_count"], **data}
try:
notify(payload)
except Exception:
notification_failed = True
result["warnings"].append("进度通知不可用;以最终保存结果和磁盘文件为准。")
LOG.warning("Encrypt Save progress notification failed", exc_info=True)
emit("start")
try:
checkpoint()
metadata = prepare_metadata(prompt, extra_pnginfo, disable_metadata)
prefix = validate_prefix(filename_prefix)
contained_folder(output_dir, Path(output_dir) / Path(prefix).parent)
folder, filename, counter, subfolder, _ = get_save_path(prefix, str(output_dir), width, height)
folder = contained_folder(output_dir, folder)
subfolder = str(subfolder).replace("\\", "/")
for index in range(count):
checkpoint()
pixels = to_pixels(convert(images[index]))
buffer = encode_png(pixels, metadata)
del pixels
with buffer:
with buffer.getbuffer() as png:
final_name, counter, warning = write_atomic(png, output_dir, folder, filename.replace("%batch_num%", str(index)), counter, checkpoint)
file = {"index": index, "filename": final_name, "subfolder": subfolder,
"type": "output", "bytes": HEADER.size + len(png)}
result["files"].append(file)
result["saved_count"] += 1
if warning:
result["warnings"].append(warning)
emit("saved", file=file, warnings=list(result["warnings"]))
counter += 1
result["status"] = "complete"
emit("complete", status="complete", files=list(result["files"]), warnings=list(result["warnings"]))
return result
except BaseException as error:
cancelled = interrupted or not isinstance(error, Exception)
status = "cancelled" if cancelled else "failed"
position = f"第 {index + 1} 张保存失败" if index is not None else "保存准备失败"
message = f"{'保存已取消' if cancelled else position},已保留 {result['saved_count']} / {count} 张:{error}"
emit("failed", status=status, failed_index=index, error=message,
files=list(result["files"]), warnings=list(result["warnings"]))
if cancelled:
if hasattr(error, "add_note"):
error.add_note(message)
raise
raise SaveBatchError(message) from error