387 lines
17 KiB
Python
387 lines
17 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""批量抠图自检:路径校验 / 命名规则 / 任务流水线 / HTTP 接口契约。
|
||
|
||
**不加载任何真实权重**(用假引擎直接产出结果图),所以跑起来只要几秒。
|
||
|
||
用法:
|
||
python\\python.exe selftest_batch.py # 全量
|
||
python selftest_batch.py --keep # 保留临时目录,便于看产物
|
||
python selftest_batch.py --logic-only # 只跑不依赖 torch 的纯逻辑部分
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import json
|
||
import shutil
|
||
import sys
|
||
import tempfile
|
||
import threading
|
||
import time
|
||
import urllib.error
|
||
import urllib.request
|
||
from pathlib import Path
|
||
from typing import List, Optional
|
||
|
||
PROJECT_ROOT = Path(__file__).resolve().parent
|
||
if str(PROJECT_ROOT) not in sys.path:
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
|
||
_FAILS: List[str] = []
|
||
|
||
|
||
def check(name: str, cond: bool, extra: object = "") -> None:
|
||
print(f" {'[通过]' if cond else '[失败]'} {name}" + ("" if cond else f" ← {extra}"))
|
||
if not cond:
|
||
_FAILS.append(name)
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# 1) 纯逻辑:路径校验 + 命名规则(不需要 torch)
|
||
# --------------------------------------------------------------------------- #
|
||
def part_logic(tmp: Path) -> None:
|
||
from birefnet_web import batch
|
||
|
||
src, out = tmp / "src", tmp / "out"
|
||
out.mkdir(parents=True)
|
||
(src / "sub").mkdir(parents=True)
|
||
(out / "RMBG_already_1.png").write_bytes(b"x")
|
||
(src / "photo.jpg").write_bytes(b"x")
|
||
(src / "a_pretty_long_original_filename_here.png").write_bytes(b"x")
|
||
(src / "note.txt").write_bytes(b"x")
|
||
(src / "RMBG_old_1.png").write_bytes(b"x")
|
||
(src / "sub" / "deep.webp").write_bytes(b"x")
|
||
|
||
print("\n[1] 路径校验")
|
||
missing = tmp / "nope"
|
||
for label, fn in (("输入", batch.check_input_dir), ("输出", batch.check_output_dir)):
|
||
try:
|
||
fn(str(missing), PROJECT_ROOT)
|
||
check(f"{label}目录不存在时抛错", False, "没有抛 BatchPathError")
|
||
except batch.BatchPathError as exc:
|
||
check(f"{label}目录不存在时抛错", True)
|
||
check(f"{label}报错时未创建目录", not missing.exists())
|
||
print(f" → {exc}")
|
||
try:
|
||
batch.check_input_dir(str(src / "photo.jpg"), PROJECT_ROOT)
|
||
check("输入是文件时抛错", False, "没有抛 BatchPathError")
|
||
except batch.BatchPathError as exc:
|
||
check("输入是文件时抛错", True)
|
||
print(f" → {exc}")
|
||
try:
|
||
batch.check_output_dir(str(PROJECT_ROOT), PROJECT_ROOT, allow_system_dir=False)
|
||
check("项目目录可写时通过", True)
|
||
except batch.BatchPathError as exc:
|
||
check("项目目录可写时通过", False, exc)
|
||
check("相对路径基于项目根解析", batch.resolve_user_path("outputs", PROJECT_ROOT) == PROJECT_ROOT / "outputs")
|
||
check("成对引号被去掉", batch.resolve_user_path(' "outputs" ', PROJECT_ROOT) == PROJECT_ROOT / "outputs")
|
||
|
||
print("\n[2] 目录扫描")
|
||
names = sorted(p.name for p in batch.scan_images(src))
|
||
check("非递归只扫顶层图片", names == ["a_pretty_long_original_filename_here.png", "photo.jpg"], names)
|
||
deep = batch.scan_images(src, recursive=True)
|
||
check("递归包含子目录", len(deep) == 3, [p.name for p in deep])
|
||
check("跳过非图片文件", all(p.suffix != ".txt" for p in deep))
|
||
check("跳过本工具产物 RMBG_*", all(not p.name.startswith("RMBG_") for p in deep))
|
||
check("跳过输出目录", all(out not in p.parents for p in batch.scan_images(src, recursive=True, skip_dirs=[out])))
|
||
|
||
print("\n[3] 命名规则")
|
||
long_src = src / "a_pretty_long_original_filename_here.png"
|
||
check("RMBG_ + 原主干 + _ + 时间戳 + 原后缀",
|
||
batch.build_output_name(src / "photo.jpg", 1700000000, "color") == "RMBG_photo_1700000000.jpg")
|
||
check("主干超 20 字符截断",
|
||
batch.build_output_name(long_src, 1700000000, "transparent") == "RMBG_a_pretty_long_origin_1700000000.png",
|
||
batch.build_output_name(long_src, 1700000000, "transparent"))
|
||
check("中文按字符截断(不是按字节)",
|
||
batch.build_output_name(Path("一二三四五六七八九十一二三四五六七八九十一二三.png"), 7, "transparent")
|
||
== "RMBG_一二三四五六七八九十一二三四五六七八九十_7.png")
|
||
check("透明 + jpg 退化 .png(jpg 装不下 alpha)",
|
||
batch.build_output_name(src / "photo.jpg", 1, "transparent") == "RMBG_photo_1.png")
|
||
check("纯色 + jpg 保持 .jpg",
|
||
batch.build_output_name(src / "photo.jpg", 1, "color") == "RMBG_photo_1.jpg")
|
||
check("非法字符被清洗", batch.build_output_name(Path("a?c*d.jpg"), 1, "color") == "RMBG_a_c_d_1.jpg")
|
||
check("同名文件避让为 _1", batch.unique_path(out, "RMBG_already_1.png").name == "RMBG_already_1_1.png")
|
||
check("无同名时原样返回", batch.unique_path(out, "RMBG_new_1.png").name == "RMBG_new_1.png")
|
||
|
||
print("\n[4] 最终尺寸 · 最长边换算")
|
||
from birefnet_web import fsutil
|
||
|
||
fit = fsutil.fit_longest_side
|
||
check("0 = 保持原图尺寸", fit((2000, 3000), 0) == (2000, 3000))
|
||
check("负数同样按不缩放处理", fit((2000, 3000), -1) == (2000, 3000))
|
||
check("需求例 2:竖图 2000×3000 → 1440 得 960×1440", fit((2000, 3000), 1440) == (960, 1440), fit((2000, 3000), 1440))
|
||
check("需求例 3:横图 3000×2000 → 1440 得 1440×960", fit((3000, 2000), 1440) == (1440, 960), fit((3000, 2000), 1440))
|
||
check("最长边已等于目标值时不动", fit((1440, 960), 1440) == (1440, 960))
|
||
check("小图按比例放大(800×600 → 1600 得 1600×1200)", fit((800, 600), 1600) == (1600, 1200), fit((800, 600), 1600))
|
||
check("正方形等比缩放", fit((1000, 1000), 500) == (500, 500))
|
||
check("极端比例也不会算出 0 像素", fit((10000, 3), 10) == (10, 1), fit((10000, 3), 10))
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# 2) 任务流水线(假引擎,不加载权重)
|
||
# --------------------------------------------------------------------------- #
|
||
class FakeModel:
|
||
key = "fake:safetensors"
|
||
name = "Fake"
|
||
|
||
def to_dict(self) -> dict:
|
||
return {"key": self.key, "name": self.name, "file": "fake.safetensors", "arch": "v1",
|
||
"size_mb": 1.0, "backbone": "swin_v1_tiny", "path": "fake"}
|
||
|
||
|
||
class FakeRegistry:
|
||
model_dirs: List[Path] = []
|
||
|
||
def list(self) -> List[FakeModel]:
|
||
return [FakeModel()]
|
||
|
||
def get(self, key: str) -> FakeModel:
|
||
if key != FakeModel.key:
|
||
raise KeyError(key)
|
||
return FakeModel()
|
||
|
||
|
||
class FakeEngine:
|
||
"""只按参数产出一张图,用来验证落盘与命名,不碰 torch。"""
|
||
|
||
def __init__(self) -> None:
|
||
self.calls = 0
|
||
|
||
def unload(self) -> None:
|
||
pass
|
||
|
||
def remove_background(self, image, options, progress=None): # noqa: ANN001
|
||
self.calls += 1
|
||
if progress:
|
||
progress("推理", 0.5)
|
||
size = image.size
|
||
transparent = options.get("background") == "transparent"
|
||
rgba = image.convert("RGBA")
|
||
rgba.putalpha(128)
|
||
cutout = rgba if transparent else rgba.convert("RGB")
|
||
return {
|
||
"cutout": cutout,
|
||
"mask": image.convert("L") if options.get("output_mask") else None,
|
||
"width": size[0], "height": size[1], "coverage": 0.5, "elapsed": 0.01,
|
||
"input_size": (1024, 1024), "warning": None, "device": "cpu", "dtype": "float32",
|
||
}
|
||
|
||
|
||
def _seed_images(root: Path) -> None:
|
||
from PIL import Image
|
||
|
||
root.mkdir(parents=True, exist_ok=True)
|
||
Image.new("RGB", (64, 48), (200, 30, 30)).save(root / "photo.jpg")
|
||
Image.new("RGB", (32, 32), (30, 200, 30)).save(root / "a_pretty_long_original_filename_here.png")
|
||
(root / "note.txt").write_text("not an image", encoding="utf-8")
|
||
|
||
|
||
def _wait(task, timeout: float = 30.0) -> bool:
|
||
deadline = time.time() + timeout
|
||
while time.time() < deadline:
|
||
if task.finished_at is not None:
|
||
return True
|
||
time.sleep(0.05)
|
||
return False
|
||
|
||
|
||
def part_pipeline(tmp: Path) -> None:
|
||
from birefnet_web.server import TaskManager
|
||
|
||
print("\n[4] 批量任务流水线(假引擎)")
|
||
src, out = tmp / "pipe_src", tmp / "pipe_out"
|
||
_seed_images(src)
|
||
# 两个不同子目录里的同名文件 → 检验同一秒内的重名避让
|
||
(src / "a").mkdir()
|
||
(src / "b").mkdir()
|
||
from PIL import Image
|
||
|
||
Image.new("RGB", (16, 16), (0, 0, 255)).save(src / "a" / "pic.png")
|
||
Image.new("RGB", (16, 16), (255, 0, 255)).save(src / "b" / "pic.png")
|
||
out.mkdir()
|
||
|
||
engine = FakeEngine()
|
||
manager = TaskManager(engine, tmp / "default_out", max_tasks=5)
|
||
options = {"background": "transparent", "output_mask": False}
|
||
|
||
task = manager.submit_batch(src, out, options, recursive=True)
|
||
check("任务模式标记为 batch", task.mode == "batch")
|
||
check("原图未被拷贝进任务目录(就地读取)",
|
||
all(Path(j.input_path).parent != tmp / "default_out" for j in task.images))
|
||
check("递归扫描命中 4 张", len(task.images) == 4, len(task.images))
|
||
check("等待处理结束", _wait(task))
|
||
check("任务状态为 done", task.state == "done", task.state)
|
||
|
||
produced = sorted(p.name for p in out.iterdir())
|
||
print(f" 产物:{produced}")
|
||
check("产物数量 = 图片数量", len(produced) == 4, produced)
|
||
check("全部带 RMBG_ 前缀", all(n.startswith("RMBG_") for n in produced))
|
||
check("jpg + 透明 → .png", any(n.startswith("RMBG_photo_") and n.endswith(".png") for n in produced))
|
||
check("长主干截断到 20 字符",
|
||
any(n.startswith("RMBG_a_pretty_long_origin_") for n in produced))
|
||
check("同名文件避让(无覆盖)",
|
||
sum(1 for n in produced if n.startswith("RMBG_pic_")) == 2, produced)
|
||
check("每张都记录了 output_name",
|
||
all(j.output_name and (out / j.output_name).is_file() for j in task.images))
|
||
|
||
# 纯色背景 + 同时输出遮罩
|
||
out2 = tmp / "pipe_out_color"
|
||
out2.mkdir()
|
||
task2 = manager.submit_batch(src / "a", out2, {"background": "color", "output_mask": True})
|
||
check("子目录只有 1 张", len(task2.images) == 1)
|
||
check("等待第二个任务结束", _wait(task2))
|
||
names2 = sorted(p.name for p in out2.iterdir())
|
||
print(f" 产物:{names2}")
|
||
check("纯色 + png 保持 .png", any(n.startswith("RMBG_pic_") and n.endswith(".png") and "_mask" not in n for n in names2))
|
||
check("同时输出遮罩 _mask.png", any(n.endswith("_mask.png") for n in names2), names2)
|
||
|
||
# 输入目录不存在 → 由 HTTP 层拦下,这里只验目录空的情况
|
||
try:
|
||
manager.submit_batch(tmp / "empty_in", out, options)
|
||
check("空目录抛 ValueError", False, "没有抛")
|
||
except ValueError as exc:
|
||
check("空目录抛 ValueError", True)
|
||
print(f" → {exc}")
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
# 3) HTTP 契约(前端就是照这个调的)
|
||
# --------------------------------------------------------------------------- #
|
||
def _post(url: str, payload: dict) -> tuple:
|
||
body = json.dumps(payload).encode("utf-8")
|
||
req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
|
||
try:
|
||
with urllib.request.urlopen(req, timeout=10) as res: # noqa: S310
|
||
return res.status, json.loads(res.read().decode("utf-8"))
|
||
except urllib.error.HTTPError as exc:
|
||
return exc.code, json.loads(exc.read().decode("utf-8"))
|
||
|
||
|
||
def _get(url: str) -> tuple:
|
||
try:
|
||
with urllib.request.urlopen(url, timeout=10) as res: # noqa: S310
|
||
return res.status, json.loads(res.read().decode("utf-8"))
|
||
except urllib.error.HTTPError as exc:
|
||
return exc.code, json.loads(exc.read().decode("utf-8"))
|
||
|
||
|
||
def part_http(tmp: Path) -> None:
|
||
from birefnet_web.server import TaskManager, WebUIServer
|
||
|
||
print("\n[5] HTTP 接口 /api/batch")
|
||
engine = FakeEngine()
|
||
out_root = tmp / "http_default_out"
|
||
tasks = TaskManager(engine, out_root, max_tasks=5)
|
||
server = WebUIServer(
|
||
("127.0.0.1", 0),
|
||
registry=FakeRegistry(), # type: ignore[arg-type]
|
||
engine=engine, # type: ignore[arg-type]
|
||
tasks=tasks,
|
||
web_dir=PROJECT_ROOT / "web",
|
||
output_root=out_root,
|
||
node_dir=PROJECT_ROOT / "vendor",
|
||
logger=lambda msg: None,
|
||
project_root=PROJECT_ROOT,
|
||
)
|
||
port = server.server_address[1]
|
||
threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True).start()
|
||
base = f"http://127.0.0.1:{port}"
|
||
try:
|
||
src = tmp / "http_src"
|
||
_seed_images(src)
|
||
out = tmp / "http_out"
|
||
out.mkdir()
|
||
|
||
status, data = _post(f"{base}/api/batch", {"input_dir": str(tmp / "missing")})
|
||
check("输入目录不存在 → 400", status == 400, (status, data))
|
||
check("错误信息直接点出不存在", "不存在" in str(data.get("error")), data)
|
||
check("仍然没有创建该目录", not (tmp / "missing").exists())
|
||
|
||
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(tmp / "missing_out")})
|
||
check("输出目录不存在 → 400", status == 400, (status, data))
|
||
check("仍然没有创建输出目录", not (tmp / "missing_out").exists())
|
||
|
||
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(src)})
|
||
check("输入 == 输出时不会「一张都扫不到」", status == 200 and data.get("count") == 2, (status, data))
|
||
if status == 200:
|
||
_wait(tasks.get(data["task_id"]))
|
||
for f in sorted(src.glob("RMBG_*")):
|
||
f.unlink()
|
||
|
||
status, data = _post(f"{base}/api/batch",
|
||
{"input_dir": str(src), "output_dir": str(out), "options": {"background": "transparent"}})
|
||
check("正常提交 → 200", status == 200 and data.get("count") == 2, (status, data))
|
||
task_id = data.get("task_id")
|
||
check("回传输出目录绝对路径", data.get("output_dir") == str(out), data.get("output_dir"))
|
||
|
||
if task_id:
|
||
_wait(tasks.get(task_id))
|
||
status, detail = _get(f"{base}/api/tasks/{task_id}?tail=1")
|
||
task = detail.get("task", {})
|
||
check("任务详情带 mode/input_dir/output_dir",
|
||
task.get("mode") == "batch" and task.get("input_dir") == str(src) and task.get("output_dir") == str(out),
|
||
task.get("mode"))
|
||
check("tail=1 只回传 1 张(载荷可控)", len(task.get("images", [])) == 1, len(task.get("images", [])))
|
||
check("images_total = 2", task.get("images_total") == 2, task.get("images_total"))
|
||
check("images_from 指向窗口起点", task.get("images_from") == 1, task.get("images_from"))
|
||
check("批量任务未被塞进单图结果网格的判据(mode=batch)", task.get("mode") == "batch")
|
||
opts = task.get("options") or {}
|
||
check("未指定的参数落到新默认值(不输出遮罩 / 最终尺寸不限)",
|
||
opts.get("output_mask") is False and opts.get("final_longest_side") == 0,
|
||
{k: opts.get(k) for k in ("output_mask", "final_longest_side")})
|
||
|
||
status, listing = _get(f"{base}/api/tasks")
|
||
modes = [t.get("mode") for t in listing.get("tasks", [])]
|
||
check("任务列表里能区分两种模式", set(modes) <= {"batch", "upload"} and "batch" in modes, modes)
|
||
|
||
status, bad = _post(f"{base}/api/batch", {"input_dir": str(src), "recursive": True})
|
||
check("留空输出目录 → 落到项目 outputs", status == 200 and bad.get("output_dir") == str(out_root),
|
||
bad.get("output_dir"))
|
||
if status == 200:
|
||
_wait(tasks.get(bad["task_id"]))
|
||
finally:
|
||
server.shutdown()
|
||
server.server_close()
|
||
|
||
|
||
# --------------------------------------------------------------------------- #
|
||
def main(argv: Optional[List[str]] = None) -> int:
|
||
parser = argparse.ArgumentParser(description="批量抠图自检")
|
||
parser.add_argument("--keep", action="store_true", help="保留临时目录")
|
||
parser.add_argument("--logic-only", action="store_true", help="只跑不依赖 torch 的纯逻辑部分")
|
||
args = parser.parse_args(argv)
|
||
|
||
tmp = Path(tempfile.mkdtemp(prefix="birefnet_batch_"))
|
||
print("=" * 68)
|
||
print("批量抠图自检")
|
||
print(f"临时目录 {tmp}")
|
||
print("=" * 68)
|
||
try:
|
||
part_logic(tmp)
|
||
if args.logic_only:
|
||
print("\n[跳过] 任务流水线与 HTTP 契约(--logic-only)")
|
||
else:
|
||
try:
|
||
part_pipeline(tmp)
|
||
part_http(tmp)
|
||
except ImportError as exc:
|
||
print(f"\n[跳过] 任务流水线与 HTTP 契约(缺少依赖:{exc})")
|
||
finally:
|
||
if args.keep:
|
||
print(f"\n临时目录已保留:{tmp}")
|
||
else:
|
||
shutil.rmtree(tmp, ignore_errors=True)
|
||
|
||
print("=" * 68)
|
||
if _FAILS:
|
||
print(f"自检未通过:{len(_FAILS)} 项失败")
|
||
for name in _FAILS:
|
||
print(f" · {name}")
|
||
return 1
|
||
print("自检全部通过 ✔")
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main())
|