Files

494 lines
14 KiB
Python
Executable File

"""视频拼接引擎纯逻辑模块.
所有函数均为纯函数,不调用 FFmpeg、不操作文件。
便于单元测试,也方便被其他模块复用。
"""
from __future__ import annotations
from pathlib import Path
from typing import Any, Optional
# ── 帧率解析 ──────────────────────────────────────────────────────────────────
def parse_fps(fps_value: Any) -> float:
"""解析帧率字符串/数值.
支持格式:
- 数字: 30 → 30.0
- 分数: "30/1" → 30.0, "24000/1001" → 23.976...
- 字符串数字: "30" → 30.0
Args:
fps_value: 帧率值(字符串、数字等)
Returns:
帧率(fps),失败返回 30.0
"""
if fps_value is None:
return 30.0
try:
fps_str = str(fps_value).strip()
if not fps_str:
return 30.0
if "/" in fps_str:
num_str, den_str = fps_str.split("/", 1)
num = float(num_str)
den = float(den_str)
if den == 0:
return 30.0
return num / den
return float(fps_str)
except (ValueError, TypeError, ZeroDivisionError):
return 30.0
def format_fps_filter(fps: float) -> str:
"""格式化 fps 滤镜参数.
Args:
fps: 帧率
Returns:
fps 滤镜字符串
"""
# 接近整数时用整数形式
if abs(fps - round(fps)) < 0.001:
return f"fps={int(fps)}"
return f"fps={fps:.3f}"
# ── 输出参数计算 ──────────────────────────────────────────────────────────────
def resolve_output_params(
config_width: int,
config_height: int,
config_fps: float,
first_video_info: Optional[dict] = None,
default_width: int = 1080,
default_height: int = 1920,
default_fps: float = 30.0,
) -> tuple[int, int, float]:
"""计算输出视频参数.
优先级:
1. config 中显式指定的(非 0 值)
2. 第一段视频的探测参数
3. 默认值
Args:
config_width: 配置的宽度(0 表示未指定)
config_height: 配置的高度(0 表示未指定)
config_fps: 配置的帧率(0 表示未指定)
first_video_info: 第一段视频的探测信息字典
default_width: 默认宽度
default_height: 默认高度
default_fps: 默认帧率
Returns:
(宽度, 高度, 帧率)
"""
width = config_width
height = config_height
fps = config_fps
info = first_video_info or {}
# 宽度:用配置 → 探测 → 默认
if width == 0:
width = int(info.get("width", default_width))
# 高度
if height == 0:
height = int(info.get("height", default_height))
# 帧率
if fps == 0:
fps_str = info.get("r_frame_rate", f"{int(default_fps)}/1")
fps = parse_fps(fps_str)
# 确保都是有效值
width = max(1, width)
height = max(1, height)
fps = max(1.0, fps)
return width, height, fps
def calculate_scaled_size(
src_w: int,
src_h: int,
target_w: int,
target_h: int,
) -> tuple[int, int, int, int]:
"""计算等比缩放后的尺寸和填充偏移.
保持宽高比,不足的部分用黑边填充。
Args:
src_w: 原始宽度
src_h: 原始高度
target_w: 目标宽度
target_h: 目标高度
Returns:
(缩放后宽度, 缩放后高度, X偏移, Y偏移)
"""
if src_w <= 0 or src_h <= 0:
return (target_w, target_h, 0, 0)
src_ratio = src_w / src_h
target_ratio = target_w / target_h
if abs(src_ratio - target_ratio) < 0.001:
# 比例相同,直接缩放
return (target_w, target_h, 0, 0)
elif src_ratio > target_ratio:
# 源更宽,以宽度为准,上下填充
scaled_w = target_w
scaled_h = int(target_w / src_ratio)
offset_x = 0
offset_y = (target_h - scaled_h) // 2
return (scaled_w, scaled_h, offset_x, offset_y)
else:
# 源更高,以高度为准,左右填充
scaled_h = target_h
scaled_w = int(target_h * src_ratio)
offset_x = (target_w - scaled_w) // 2
offset_y = 0
return (scaled_w, scaled_h, offset_x, offset_y)
# ── stream copy 判断 ─────────────────────────────────────────────────────────
def can_use_stream_copy(
segments: list[dict],
target_width: int,
target_height: int,
target_fps: float,
force_reencode: bool = False,
) -> bool:
"""判断是否可以使用 stream copy(无损拼接).
stream copy 条件:
1. force_reencode 为 False
2. 所有视频段的编码格式、分辨率、帧率均相同
3. 目标参数与源参数一致(不需要转码)
Args:
segments: 视频段列表,每个元素包含 codec_name/width/height/fps
target_width: 目标宽度
target_height: 目标高度
target_fps: 目标帧率
force_reencode: 是否强制重编码
Returns:
是否可以用 stream copy
"""
if force_reencode:
return False
if not segments:
return False
# 用第一段作为基准
first = segments[0]
base_codec = first.get("codec_name", "")
base_width = int(first.get("width", 0))
base_height = int(first.get("height", 0))
base_fps = parse_fps(first.get("r_frame_rate", "30/1"))
# 目标参数必须与基准一致
if target_width != base_width or target_height != base_height:
return False
if abs(target_fps - base_fps) > 0.01:
return False
# 所有段必须参数一致
for seg in segments[1:]:
if seg.get("codec_name", "") != base_codec:
return False
if int(seg.get("width", 0)) != base_width:
return False
if int(seg.get("height", 0)) != base_height:
return False
seg_fps = parse_fps(seg.get("r_frame_rate", "30/1"))
if abs(seg_fps - base_fps) > 0.01:
return False
return True
# ── 文件列表生成(demuxer 模式) ───────────────────────────────────────────────
def generate_concat_file_list(
video_paths: list[str],
) -> str:
"""生成 concat demuxer 模式的文件列表内容.
格式:
file '/path/to/video1.mp4'
file '/path/to/video2.mp4'
Args:
video_paths: 视频文件路径列表
Returns:
文件列表文本内容
"""
lines = []
for path in video_paths:
# 转义单引号
escaped = path.replace("'", "'\\''")
lines.append(f"file '{escaped}'")
return "\n".join(lines) + "\n"
# ── 滤镜链构建 ────────────────────────────────────────────────────────────────
def build_scale_pad_filter(
target_w: int,
target_h: int,
src_w: int = 0,
src_h: int = 0,
) -> str:
"""构建 scale + pad 滤镜(等比缩放+黑边填充).
Args:
target_w: 目标宽度
target_h: 目标高度
src_w: 源宽度(0 表示未知,用 iw/ih)
src_h: 源高度(0 表示未知)
Returns:
滤镜字符串
"""
# 使用 FFmpeg 表达式,动态计算
return (
f"scale={target_w}:{target_h}:force_original_aspect_ratio=decrease,"
f"pad={target_w}:{target_h}:(ow-iw)/2:(oh-ih)/2:black"
)
def build_fps_filter(fps: float) -> str:
"""构建 fps 滤镜.
Args:
fps: 目标帧率
Returns:
fps 滤镜字符串
"""
return format_fps_filter(fps)
def build_setpts_filter() -> str:
"""构建 setpts 滤镜(重置时间戳).
Returns:
setpts 滤镜字符串
"""
return "setpts=PTS-STARTPTS"
def build_concat_filter(
num_inputs: int,
has_audio: bool = True,
) -> str:
"""构建 concat 滤镜.
Args:
num_inputs: 输入数量
has_audio: 是否包含音频轨
Returns:
concat 滤镜字符串(包含输入标签)
"""
if num_inputs <= 0:
return ""
n = num_inputs
v = 1 # 视频轨数
a = 1 if has_audio else 0 # 音频轨数
# 构建输入标签
input_labels = "".join(f"[{i}:v][{i}:a]" if has_audio else f"[{i}:v]" for i in range(n))
output_label = "[concat_v]" + ("[concat_a]" if has_audio else "")
return f"{input_labels}concat=n={n}:v={v}:a={a}{output_label}"
def build_single_segment_filter_chain(
target_width: int,
target_height: int,
target_fps: float,
segment_index: int,
has_audio: bool = True,
) -> str:
"""构建单段视频的预处理滤镜链.
每段视频需要:缩放填充 → 帧率统一 → 重置时间戳
Args:
target_width: 目标宽度
target_height: 目标高度
target_fps: 目标帧率
segment_index: 段索引(用于生成标签)
has_audio: 是否包含音频
Returns:
滤镜字符串
"""
scale_pad = build_scale_pad_filter(target_width, target_height)
fps = build_fps_filter(target_fps)
setpts = build_setpts_filter()
input_v = f"[{segment_index}:v]"
output_v = f"[v{segment_index}]"
video_chain = f"{input_v}{scale_pad},{fps},{setpts}{output_v}"
if has_audio:
input_a = f"[{segment_index}:a]"
output_a = f"[a{segment_index}]"
# 音频也需要重置时间戳
audio_chain = f"{input_a}asetpts=PTS-STARTPTS{output_a}"
return f"{video_chain};{audio_chain}"
return video_chain
# ── 配置验证 ──────────────────────────────────────────────────────────────────
def validate_concat_config(config: dict) -> tuple[bool, list[str]]:
"""验证拼接配置.
Args:
config: 配置字典
Returns:
(是否合法, 错误信息列表)
"""
errors: list[str] = []
segments = config.get("segments", [])
if not segments:
errors.append("至少需要一个视频段")
return (False, errors)
if len(segments) < 1:
errors.append("视频段数量不能少于 1")
# 检查每个段
for i, seg in enumerate(segments):
video_path = seg.get("video_path", "")
if not video_path:
errors.append(f"第 {i+1} 段缺少 video_path")
# 输出参数
output_width = config.get("output_width", 0)
output_height = config.get("output_height", 0)
if output_width < 0:
errors.append("output_width 不能为负数")
if output_height < 0:
errors.append("output_height 不能为负数")
output_fps = config.get("output_fps", 0)
if output_fps < 0:
errors.append("output_fps 不能为负数")
return (len(errors) == 0, errors)
# ── 路径验证 ──────────────────────────────────────────────────────────────────
def validate_video_path(video_path: str, work_dir: str | Path) -> tuple[bool, str]:
"""验证视频路径是否安全.
检查:
1. 路径不为空
2. 路径不包含 .. 回溯
3. 路径在 work_dir 内(安全边界)
Args:
video_path: 视频文件路径
work_dir: 工作目录
Returns:
(是否合法, 错误信息)
"""
if not video_path:
return (False, "视频路径不能为空")
path_str = str(video_path)
work_str = str(work_dir)
# 检查路径遍历
if ".." in Path(path_str).parts:
return (False, "视频路径不能包含 .. 回溯")
# 绝对路径才做边界检查;相对路径默认相对于 work_dir
if not Path(path_str).is_absolute():
return (True, "")
# 绝对路径检查是否在工作目录内
try:
video_abs = Path(path_str).resolve()
work_abs = Path(work_str).resolve()
if work_abs.is_absolute() and not str(video_abs).startswith(str(work_abs)):
return (False, "视频路径必须在工作目录内")
except (OSError, ValueError):
pass # 解析失败时跳过边界检查
return (True, "")
# ── 工具函数 ──────────────────────────────────────────────────────────────────
def estimate_total_duration(segments: list[dict]) -> float:
"""估算总时长.
Args:
segments: 视频段列表,每个元素包含 duration 字段
Returns:
总时长(秒)
"""
total = 0.0
for seg in segments:
dur = seg.get("duration", 0)
try:
total += float(dur)
except (ValueError, TypeError):
pass
return total
def count_valid_segments(segments: list[dict]) -> int:
"""统计有效视频段数量(有 video_path 的).
Args:
segments: 视频段列表
Returns:
有效段数量
"""
count = 0
for seg in segments:
if seg.get("video_path"):
count += 1
return count