Files
BiRefNet_WebUI/birefnet_web/batch.py
T
2026-10-08 09:46:47 +08:00

223 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""目录批量抠图:路径校验 / 目录扫描 / 输出文件命名。
三条硬约束(对应需求):
1. **校验阶段只读不写**:输入目录、输出目录都必须**已经存在**,本模块
绝不 ``mkdir``。路径写错时直接报错,而不是顺手把目录树建出来 ——
这样一次误输入(或网页上的恶意路径)不会在磁盘上留下任何东西。
2. **命名固定**:``RMBG_<原主干>_<Unix 时间戳>.<后缀>``,原主干超过
:data:`STEM_MAXLEN` 个字符时截断。主干统一走
:func:`birefnet_web.imageops.safe_stem` 清洗,杜绝 ``..`` 之类的穿越。
3. **后缀自适应**:透明输出需要 alpha 通道,``.jpg`` / ``.bmp`` 装不下,
此时退化为 ``.png``(强行按原后缀保存会被 PIL 静默压成白底图)。
本模块不依赖 torch,可以脱离推理环境单独测试。
"""
from __future__ import annotations
import os
from pathlib import Path
from typing import Iterable, List, Optional, Sequence
from .fsutil import ALPHA_SUFFIXES, IMAGE_SUFFIXES, safe_stem
#: 输出文件名统一前缀(扫描输入目录时也用来自我过滤,避免把产物再抠一遍)
BATCH_PREFIX = "RMBG_"
#: 原文件名主干的最大长度(超过即截断)
STEM_MAXLEN = 20
#: 重名避让时最多尝试的序号
_MAX_DEDUP = 999
class BatchPathError(ValueError):
"""路径不合规:不存在 / 不是目录 / 不可写 / 命中系统目录黑名单。"""
# --------------------------------------------------------------------------- #
# 路径解析与校验
# --------------------------------------------------------------------------- #
def _critical_dirs() -> List[Path]:
"""返回当前系统上「不该往里写文件」的目录(仅用于提示级别的高危拦截)。"""
roots: List[Path] = []
names = ("SystemRoot", "windir", "ProgramFiles", "ProgramFiles(x86)", "ProgramData")
for name in names:
raw = os.environ.get(name)
if not raw:
continue
p = Path(raw)
if p.is_absolute():
roots.append(p)
if os.name != "nt":
roots.extend(Path(p) for p in ("/bin", "/sbin", "/etc", "/usr", "/boot", "/dev", "/proc", "/sys"))
return roots
def resolve_user_path(raw: object, base: Path) -> Path:
"""把用户在界面上敲的一行路径规范化成绝对路径。
Args:
raw: 原始输入(可能带首尾空格、成对引号、``~`` 或 ``%VAR%``)。
base: 相对路径的基准目录(本项目的 ``PROJECT_ROOT``)。
Returns:
规范化后的绝对路径(**不校验存在性**,也不做任何创建)。
Raises:
BatchPathError: 输入为空。
"""
text = str(raw or "").strip().strip('"').strip("'").strip()
if not text:
raise BatchPathError("路径不能为空")
text = os.path.expandvars(os.path.expanduser(text))
path = Path(text)
if not path.is_absolute():
path = base / path
return Path(os.path.normpath(str(path)))
def check_input_dir(raw: object, base: Path) -> Path:
"""校验「图片所在目录」,必须已存在且是目录。
Raises:
BatchPathError: 路径为空 / 不存在 / 不是目录。
"""
path = resolve_user_path(raw, base)
if not path.exists():
raise BatchPathError(f"输入目录不存在:{path}(本工具不会自动创建目录,请先确认路径)")
if not path.is_dir():
raise BatchPathError(f"输入路径不是目录:{path}")
return path
def check_output_dir(raw: object, base: Path, *, allow_system_dir: bool = False) -> Path:
"""校验「结果输出目录」,必须已存在、是目录且可写。
Args:
raw: 用户输入或默认输出目录。
base: 相对路径基准。
allow_system_dir: 为 True 时跳过系统目录黑名单(测试用)。
Raises:
BatchPathError: 路径为空 / 不存在 / 不是目录 / 命中黑名单 / 不可写。
"""
path = resolve_user_path(raw, base)
if not path.exists():
raise BatchPathError(f"输出目录不存在:{path}(本工具不会自动创建目录,请先手动建好)")
if not path.is_dir():
raise BatchPathError(f"输出路径不是目录:{path}")
if not allow_system_dir:
if path == Path(path.anchor): # C:\ 或 /
raise BatchPathError(f"拒绝把结果写进磁盘根目录:{path}")
for critical in _critical_dirs():
if path == critical or critical in path.parents:
raise BatchPathError(f"拒绝把结果写进系统目录:{path}")
if not os.access(path, os.W_OK):
raise BatchPathError(f"输出目录不可写(权限不足):{path}")
return path
# --------------------------------------------------------------------------- #
# 目录扫描
# --------------------------------------------------------------------------- #
def scan_images(
root: Path,
*,
recursive: bool = False,
skip_dirs: Sequence[Path] = (),
) -> List[Path]:
"""扫描目录下的图片文件。
Args:
root: 已校验过的图片目录。
recursive: 是否包含子目录。
skip_dirs: 需要跳过的目录(通常是输出目录,避免读到自己的产物)。
Returns:
排序后的图片路径列表(先按目录、再按文件名,保证同一批次的处理顺序稳定)。
"""
skipped = [Path(d).resolve() for d in skip_dirs]
walker: Iterable[Path] = root.rglob("*") if recursive else root.glob("*")
found: List[Path] = []
for path in walker:
try:
if not path.is_file():
continue
except OSError: # 权限/坏链接
continue
if path.suffix.lower() not in IMAGE_SUFFIXES:
continue
# 本工具自己的产物(RMBG_ 前缀)不参与批量,避免反复抠同一张图
if path.name.startswith(BATCH_PREFIX):
continue
if skipped:
parent = path.parent.resolve()
if any(parent == s or s in parent.parents or parent in s.parents for s in skipped):
continue
found.append(path)
return sorted(found, key=lambda p: (str(p.parent).lower(), p.name.lower()))
# --------------------------------------------------------------------------- #
# 输出命名
# --------------------------------------------------------------------------- #
def choose_suffix(orig_suffix: str, background: str) -> str:
"""决定输出文件后缀:透明模式下如果原后缀装不下 alpha,就退化为 ``.png``。
Args:
orig_suffix: 原文件后缀(含点,如 ``.jpg``)。
background: 背景模式,``transparent`` 或其它(纯色)。
Returns:
合法的输出后缀字符串。
"""
suffix = (orig_suffix or "").lower()
if background == "transparent":
return suffix if suffix in ALPHA_SUFFIXES else ".png"
return suffix if suffix in IMAGE_SUFFIXES else ".png"
def build_output_name(source: Path, timestamp: int, background: str = "transparent") -> str:
"""按 ``RMBG_<主干(≤20)>_<Unix 时间戳>.<后缀>`` 组装输出文件名。
Args:
source: 原始图片路径。
timestamp: Unix 时间戳(秒),由调用方在**处理时**取,便于追溯。
background: 背景模式,决定后缀是否允许退化为 ``.png``。
Returns:
输出文件名(不含目录)。
"""
stem = safe_stem(source.name, fallback="image", maxlen=STEM_MAXLEN)
return f"{BATCH_PREFIX}{stem}_{int(timestamp)}{choose_suffix(source.suffix, background)}"
def unique_path(directory: Path, name: str) -> Path:
"""在同一秒内出现同名输出时,用 ``_1`` / ``_2`` 递增避让,绝不覆盖已有文件。
Args:
directory: 输出目录(已存在)。
name: 目标文件名。
Returns:
可安全写入的完整路径。
Raises:
BatchPathError: 避让序号耗尽(同名文件超过 :data:`_MAX_DEDUP` 个)。
"""
target = directory / name
if not target.exists():
return target
for i in range(1, _MAX_DEDUP + 1):
candidate = directory / f"{target.stem}_{i}{target.suffix}"
if not candidate.exists():
return candidate
raise BatchPathError(f"输出目录中同名文件过多,请检查:{name}")