223 lines
8.3 KiB
Python
223 lines
8.3 KiB
Python
# -*- 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}")
|