282 lines
12 KiB
Python
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
|