Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 01204795d0 | |||
| e175f60bef | |||
| ad0a559035 | |||
| c5933e865a | |||
| 90f703e49e | |||
| cafe063859 | |||
| d92ff7aa36 | |||
| cb85b33ee9 |
+292
-466
@@ -1,9 +1,6 @@
|
||||
"""BGM 混音纯逻辑单元测试."""
|
||||
"""bgm_mixer_pure 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from video_processing.bgm_mixer_pure import (
|
||||
from apps.worker.video_processing.bgm_mixer_pure import (
|
||||
BGMPureConfig,
|
||||
build_bgm_filter_chain,
|
||||
build_sidechain_mix_filter,
|
||||
@@ -17,408 +14,286 @@ from video_processing.bgm_mixer_pure import (
|
||||
validate_bgm_config,
|
||||
)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# should_loop_bgm 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── BGMPureConfig ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestShouldLoopBGM:
|
||||
"""BGM 循环判断测试."""
|
||||
class TestBGMPureConfig:
|
||||
def test_default_values(self):
|
||||
cfg = BGMPureConfig()
|
||||
assert cfg.volume == 0.3
|
||||
assert cfg.fade_in == 0.0
|
||||
assert cfg.fade_out == 0.0
|
||||
assert cfg.loop_enabled is True
|
||||
assert cfg.sidechain_enabled is False
|
||||
assert cfg.sidechain_ratio == 0.3
|
||||
assert cfg.sidechain_attack == 0.02
|
||||
assert cfg.sidechain_release == 0.5
|
||||
assert cfg.sidechain_threshold == -25.0
|
||||
|
||||
def test_need_loop_when_much_shorter(self):
|
||||
"""BGM 远短于目标时长,需要循环."""
|
||||
assert should_loop_bgm(10, 100, True) is True
|
||||
def test_custom_values(self):
|
||||
cfg = BGMPureConfig(
|
||||
volume=0.5,
|
||||
fade_in=1.0,
|
||||
fade_out=2.0,
|
||||
loop_enabled=False,
|
||||
sidechain_enabled=True,
|
||||
sidechain_ratio=0.5,
|
||||
)
|
||||
assert cfg.volume == 0.5
|
||||
assert cfg.loop_enabled is False
|
||||
assert cfg.sidechain_enabled is True
|
||||
assert cfg.sidechain_ratio == 0.5
|
||||
|
||||
def test_no_loop_when_long_enough(self):
|
||||
"""BGM 够长,不需要循环."""
|
||||
assert should_loop_bgm(100, 100, True) is False
|
||||
|
||||
def test_no_loop_when_just_slightly_shorter(self):
|
||||
"""BGM 只差一点点(>90%),不循环."""
|
||||
assert should_loop_bgm(95, 100, True) is False
|
||||
# ── should_loop_bgm ─────────────────────────────────────────────────────────
|
||||
|
||||
def test_threshold_90_percent(self):
|
||||
"""刚好 90% 阈值,不循环(<90% 才循环)."""
|
||||
assert should_loop_bgm(90, 100, True) is False
|
||||
|
||||
def test_just_below_threshold(self):
|
||||
"""略低于 90%,需要循环."""
|
||||
assert should_loop_bgm(89, 100, True) is True
|
||||
class TestShouldLoopBgm:
|
||||
def test_loop_enabled_much_shorter(self):
|
||||
# BGM 10秒,目标60秒 → 需要循环
|
||||
assert should_loop_bgm(10, 60) is True
|
||||
|
||||
def test_loop_disabled(self):
|
||||
"""禁用循环,即使 BGM 很短也不循环."""
|
||||
assert should_loop_bgm(10, 100, False) is False
|
||||
assert should_loop_bgm(10, 60, loop_enabled=False) is False
|
||||
|
||||
def test_bgm_longer_than_target(self):
|
||||
# BGM 100秒,目标60秒 → 不需要循环
|
||||
assert should_loop_bgm(100, 60) is False
|
||||
|
||||
def test_bgm_slightly_shorter_no_loop(self):
|
||||
# BGM 58秒,目标60秒 → 58 > 60*0.9=54,不需要循环
|
||||
assert should_loop_bgm(58, 60) is False
|
||||
|
||||
def test_bgm_significantly_shorter_loops(self):
|
||||
# BGM 50秒,目标60秒 → 50 < 54,需要循环
|
||||
assert should_loop_bgm(50, 60) is True
|
||||
|
||||
def test_zero_bgm_duration(self):
|
||||
"""BGM 时长为 0,不循环."""
|
||||
assert should_loop_bgm(0, 100, True) is False
|
||||
assert should_loop_bgm(0, 60) is False
|
||||
|
||||
def test_negative_bgm_duration(self):
|
||||
"""BGM 时长为负,不循环."""
|
||||
assert should_loop_bgm(-5, 100, True) is False
|
||||
assert should_loop_bgm(-1, 60) is False
|
||||
|
||||
def test_zero_target_duration(self):
|
||||
"""目标时长为 0,不循环."""
|
||||
assert should_loop_bgm(10, 0, True) is False
|
||||
assert should_loop_bgm(10, 0) is False
|
||||
|
||||
def test_negative_target_duration(self):
|
||||
"""目标时长为负,不循环."""
|
||||
assert should_loop_bgm(10, -10, True) is False
|
||||
assert should_loop_bgm(10, -1) is False
|
||||
|
||||
def test_exact_90_percent_no_loop(self):
|
||||
# 边界:bgm == target * 0.9 → 不小于,不循环
|
||||
assert should_loop_bgm(54, 60) is False
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# calculate_loop_count 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── calculate_loop_count ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateLoopCount:
|
||||
"""循环次数计算测试."""
|
||||
def test_exact_fit_returns_1(self):
|
||||
assert calculate_loop_count(60, 60) == 1
|
||||
|
||||
def test_exact_multiple(self):
|
||||
"""刚好整数倍."""
|
||||
# 100/10 = 10, +2 = 12
|
||||
assert calculate_loop_count(10, 100) == 12
|
||||
def test_bgm_longer_returns_1(self):
|
||||
assert calculate_loop_count(100, 60) == 1
|
||||
|
||||
def test_not_exact_multiple(self):
|
||||
"""不是整数倍."""
|
||||
# 100/30 = 3, +2 = 5
|
||||
assert calculate_loop_count(30, 100) == 5
|
||||
def test_needs_3_loops_plus_2_margin(self):
|
||||
# 60/20 = 3 + 2 = 5
|
||||
assert calculate_loop_count(20, 60) == 5
|
||||
|
||||
def test_bgm_longer_than_target(self):
|
||||
"""BGM 比目标长,至少 1 次."""
|
||||
assert calculate_loop_count(200, 100) == 1
|
||||
def test_needs_2_loops_plus_2_margin(self):
|
||||
# 60/30 = 2 + 2 = 4
|
||||
assert calculate_loop_count(30, 60) == 4
|
||||
|
||||
def test_zero_bgm_duration(self):
|
||||
"""BGM 时长为 0,返回 1."""
|
||||
assert calculate_loop_count(0, 100) == 1
|
||||
assert calculate_loop_count(0, 60) == 1
|
||||
|
||||
def test_negative_bgm_duration(self):
|
||||
"""BGM 时长为负,返回 1."""
|
||||
assert calculate_loop_count(-5, 100) == 1
|
||||
assert calculate_loop_count(-1, 60) == 1
|
||||
|
||||
def test_zero_target_duration(self):
|
||||
"""目标时长为 0,返回 1."""
|
||||
assert calculate_loop_count(10, 0) == 1
|
||||
|
||||
def test_negative_target_duration(self):
|
||||
"""目标时长为负,返回 1."""
|
||||
assert calculate_loop_count(10, -10) == 1
|
||||
assert calculate_loop_count(10, -1) == 1
|
||||
|
||||
def test_very_short_bgm(self):
|
||||
"""非常短的 BGM,循环次数多."""
|
||||
# 100/1 = 100, +2 = 102
|
||||
assert calculate_loop_count(1, 100) == 102
|
||||
def test_minimum_is_1(self):
|
||||
assert calculate_loop_count(10, 5) == 1
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# build_bgm_filter_chain 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── build_bgm_filter_chain ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildBGMFilterChain:
|
||||
"""BGM 预处理滤镜链构建测试."""
|
||||
class TestBuildBgmFilterChain:
|
||||
def test_basic_structure(self):
|
||||
result = build_bgm_filter_chain(100, 60)
|
||||
parts = result.split(",")
|
||||
# 至少有 atrim + asetpts
|
||||
assert any("atrim=" in p for p in parts)
|
||||
assert "asetpts=N/SR/TB" in parts
|
||||
|
||||
def test_basic_volume_only(self):
|
||||
"""只有音量调节."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=0.5,
|
||||
)
|
||||
def test_volume_filter_applied(self):
|
||||
result = build_bgm_filter_chain(100, 60, volume=0.5)
|
||||
assert "volume=0.500" in result
|
||||
assert "aloop" not in result
|
||||
assert "afade=t=in" not in result
|
||||
assert "afade=t=out" not in result
|
||||
assert "atrim=0:100.000" in result
|
||||
assert "asetpts=N/SR/TB" in result
|
||||
|
||||
def test_with_loop(self):
|
||||
"""需要循环的情况."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=10,
|
||||
target_duration=100,
|
||||
volume=0.3,
|
||||
loop_enabled=True,
|
||||
)
|
||||
assert "aloop=loop=" in result
|
||||
assert "volume=0.300" in result
|
||||
def test_volume_one_omitted(self):
|
||||
result = build_bgm_filter_chain(100, 60, volume=1.0)
|
||||
assert "volume=" not in result
|
||||
|
||||
def test_no_loop_when_disabled(self):
|
||||
"""禁用循环,即使 BGM 短也不循环."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=10,
|
||||
target_duration=100,
|
||||
volume=0.3,
|
||||
loop_enabled=False,
|
||||
)
|
||||
assert "aloop" not in result
|
||||
def test_volume_clamped(self):
|
||||
# volume=2.0钳制到1.0,1.0等于默认值所以被跳过
|
||||
result = build_bgm_filter_chain(100, 60, volume=2.0)
|
||||
assert "volume=" not in result # 钳制到1.0后与默认相同,跳过
|
||||
# 用0.5验证音量过滤器本身存在
|
||||
result2 = build_bgm_filter_chain(100, 60, volume=0.5)
|
||||
assert "volume=0.500" in result2
|
||||
|
||||
def test_fade_in_only(self):
|
||||
"""只有淡入."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=1.0,
|
||||
fade_in=2.5,
|
||||
)
|
||||
assert "afade=t=in:st=0:d=2.500" in result
|
||||
assert "afade=t=out" not in result
|
||||
assert "volume=" not in result # volume=1.0 不加
|
||||
def test_volume_zero(self):
|
||||
result = build_bgm_filter_chain(100, 60, volume=0.0)
|
||||
assert "volume=0.000" in result
|
||||
|
||||
def test_fade_out_only(self):
|
||||
"""只有淡出."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=1.0,
|
||||
fade_out=3.0,
|
||||
)
|
||||
assert "afade=t=out:st=97.000:d=3.000" in result
|
||||
assert "afade=t=in" not in result
|
||||
|
||||
def test_fade_in_and_out(self):
|
||||
"""淡入+淡出."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=1.0,
|
||||
fade_in=1.5,
|
||||
fade_out=2.0,
|
||||
)
|
||||
def test_fade_in_applied(self):
|
||||
result = build_bgm_filter_chain(100, 60, fade_in=1.5)
|
||||
assert "afade=t=in:st=0:d=1.500" in result
|
||||
assert "afade=t=out:st=98.000:d=2.000" in result
|
||||
|
||||
def test_volume_1_0_skipped(self):
|
||||
"""音量为 1.0 时不添加 volume 滤镜."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=1.0,
|
||||
)
|
||||
assert "volume=" not in result
|
||||
def test_fade_in_zero_skipped(self):
|
||||
result = build_bgm_filter_chain(100, 60, fade_in=0)
|
||||
assert "afade=t=in" not in result
|
||||
|
||||
def test_volume_0(self):
|
||||
"""音量为 0."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=0.0,
|
||||
)
|
||||
assert "volume=0.000" in result
|
||||
def test_fade_out_applied(self):
|
||||
result = build_bgm_filter_chain(100, 60, fade_out=2.0)
|
||||
assert "afade=t=out:st=58.000:d=2.000" in result
|
||||
|
||||
def test_volume_clamped_high(self):
|
||||
"""音量超过 1.0 被钳制."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=1.5,
|
||||
)
|
||||
assert "volume=1.000" not in result # 1.0不加
|
||||
# 钳制到1.0后和1.0一样,不加volume滤镜
|
||||
# 但因为abs(1.0 - 1.0) < 0.001,所以不添加
|
||||
assert "volume=" not in result
|
||||
|
||||
def test_volume_clamped_low(self):
|
||||
"""音量为负被钳制到 0."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=100,
|
||||
volume=-0.5,
|
||||
)
|
||||
assert "volume=0.000" in result
|
||||
|
||||
def test_fade_out_longer_than_duration(self):
|
||||
"""淡出时长超过总时长,不加淡出."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=10,
|
||||
volume=1.0,
|
||||
fade_out=20.0,
|
||||
)
|
||||
def test_fade_out_longer_than_target_skipped(self):
|
||||
result = build_bgm_filter_chain(100, 10, fade_out=20)
|
||||
# fade_out >= safe_target,不做淡出
|
||||
assert "afade=t=out" not in result
|
||||
|
||||
def test_fade_out_equal_to_duration(self):
|
||||
"""淡出时长等于总时长,不加淡出."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=200,
|
||||
target_duration=10,
|
||||
volume=1.0,
|
||||
fade_out=10.0,
|
||||
)
|
||||
assert "afade=t=out" not in result
|
||||
def test_loop_applied_when_needed(self):
|
||||
result = build_bgm_filter_chain(10, 60)
|
||||
assert "aloop=loop=" in result
|
||||
|
||||
def test_zero_target_duration_fallback(self):
|
||||
"""目标时长为 0,兜底 5 秒."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=3,
|
||||
target_duration=0,
|
||||
volume=0.5,
|
||||
)
|
||||
def test_no_loop_when_bgm_long(self):
|
||||
result = build_bgm_filter_chain(100, 60)
|
||||
assert "aloop=" not in result
|
||||
|
||||
def test_loop_disabled(self):
|
||||
result = build_bgm_filter_chain(10, 60, loop_enabled=False)
|
||||
assert "aloop=" not in result
|
||||
|
||||
def test_trim_to_target_duration(self):
|
||||
result = build_bgm_filter_chain(100, 60)
|
||||
assert "atrim=0:60.000" in result
|
||||
|
||||
def test_zero_target_uses_fallback(self):
|
||||
result = build_bgm_filter_chain(100, 0)
|
||||
# 兜底5秒
|
||||
assert "atrim=0:5.000" in result
|
||||
|
||||
def test_negative_target_duration_fallback(self):
|
||||
"""目标时长为负,兜底 5 秒."""
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=3,
|
||||
target_duration=-5,
|
||||
volume=0.5,
|
||||
)
|
||||
def test_negative_target_uses_fallback(self):
|
||||
result = build_bgm_filter_chain(100, -5)
|
||||
assert "atrim=0:5.000" in result
|
||||
|
||||
def test_full_chain_with_all_effects(self):
|
||||
"""完整滤镜链:循环+音量+淡入淡出+截断+重置."""
|
||||
def test_all_features_combined(self):
|
||||
result = build_bgm_filter_chain(
|
||||
bgm_duration=10,
|
||||
target_duration=100,
|
||||
volume=0.4,
|
||||
bgm_duration=15,
|
||||
target_duration=60,
|
||||
volume=0.3,
|
||||
fade_in=1.0,
|
||||
fade_out=2.0,
|
||||
loop_enabled=True,
|
||||
)
|
||||
parts = result.split(",")
|
||||
# 顺序:aloop -> volume -> afade in -> afade out -> atrim -> asetpts
|
||||
assert len(parts) >= 6
|
||||
assert "aloop" in parts[0]
|
||||
assert "volume" in parts[1]
|
||||
assert "afade=t=in" in parts[2]
|
||||
assert "afade=t=out" in parts[3]
|
||||
assert "atrim" in parts[4]
|
||||
assert "asetpts" in parts[5]
|
||||
assert "aloop=loop=" in result
|
||||
assert "volume=0.300" in result
|
||||
assert "afade=t=in" in result
|
||||
assert "afade=t=out" in result
|
||||
assert "atrim=0:60.000" in result
|
||||
assert "asetpts=N/SR/TB" in result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# calculate_sidechain_ratio 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── calculate_sidechain_ratio ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateSidechainRatio:
|
||||
"""Sidechain 压缩比计算测试."""
|
||||
def test_zero_ratio_minimum(self):
|
||||
assert calculate_sidechain_ratio(0) == 2.0
|
||||
|
||||
def test_default_ratio_0_3(self):
|
||||
"""默认 0.3."""
|
||||
# 1 / (1 - 0.3) = 1.428... 但下限是 2.0
|
||||
assert calculate_sidechain_ratio(0.3) == pytest.approx(2.0, rel=0.01)
|
||||
|
||||
def test_ratio_0_5(self):
|
||||
"""比例 0.5."""
|
||||
# 1 / (1 - 0.5) = 2.0
|
||||
assert calculate_sidechain_ratio(0.5) == pytest.approx(2.0, rel=0.01)
|
||||
|
||||
def test_ratio_0_8(self):
|
||||
"""比例 0.8."""
|
||||
# 1 / (1 - 0.8) = 5.0
|
||||
assert calculate_sidechain_ratio(0.8) == pytest.approx(5.0, rel=0.01)
|
||||
|
||||
def test_ratio_0_9(self):
|
||||
"""比例 0.9."""
|
||||
# 1 / (1 - 0.9) = 10.0
|
||||
assert calculate_sidechain_ratio(0.9) == pytest.approx(10.0, rel=0.01)
|
||||
|
||||
def test_ratio_0(self):
|
||||
"""比例 0,返回下限 2.0."""
|
||||
assert calculate_sidechain_ratio(0.0) == 2.0
|
||||
|
||||
def test_ratio_negative(self):
|
||||
"""比例为负,返回下限 2.0."""
|
||||
def test_negative_clamped(self):
|
||||
assert calculate_sidechain_ratio(-0.5) == 2.0
|
||||
|
||||
def test_ratio_1_0(self):
|
||||
"""比例 1.0,返回上限 10.0."""
|
||||
def test_one_ratio_maximum(self):
|
||||
assert calculate_sidechain_ratio(1.0) == 10.0
|
||||
|
||||
def test_ratio_greater_than_1(self):
|
||||
"""比例超过 1.0,返回上限 10.0."""
|
||||
assert calculate_sidechain_ratio(2.0) == 10.0
|
||||
def test_above_one_clamped(self):
|
||||
assert calculate_sidechain_ratio(1.5) == 10.0
|
||||
|
||||
def test_mid_value(self):
|
||||
# ratio = 1/(1-0.5) = 2.0
|
||||
result = calculate_sidechain_ratio(0.5)
|
||||
assert abs(result - 2.0) < 0.01
|
||||
|
||||
def test_high_value(self):
|
||||
# 1/(1-0.9) = 10 → 钳制到10
|
||||
assert calculate_sidechain_ratio(0.9) == 10.0
|
||||
|
||||
def test_03_default(self):
|
||||
# 1/(1-0.3) = 1.428... → 钳制到2.0
|
||||
result = calculate_sidechain_ratio(0.3)
|
||||
assert result >= 2.0
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# build_simple_mix_filter 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── build_simple_mix_filter ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildSimpleMixFilter:
|
||||
"""普通混音滤镜构建测试."""
|
||||
|
||||
def test_contains_amix(self):
|
||||
"""包含 amix."""
|
||||
def test_contains_inputs_and_output(self):
|
||||
result = build_simple_mix_filter()
|
||||
assert "[0:a][1:a]" in result
|
||||
assert "amix=inputs=2" in result
|
||||
|
||||
def test_contains_volume_compensation(self):
|
||||
"""包含 volume=2 补偿."""
|
||||
result = build_simple_mix_filter()
|
||||
assert "duration=first" in result
|
||||
assert "[final]" in result
|
||||
assert "volume=2" in result
|
||||
|
||||
def test_output_label(self):
|
||||
"""输出标签为 [final]."""
|
||||
result = build_simple_mix_filter()
|
||||
assert "[final]" in result
|
||||
|
||||
def test_duration_first(self):
|
||||
"""duration=first,以主音频时长为准."""
|
||||
result = build_simple_mix_filter()
|
||||
assert "duration=first" in result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# build_sidechain_mix_filter 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── build_sidechain_mix_filter ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildSidechainMixFilter:
|
||||
"""Sidechain 混音滤镜构建测试."""
|
||||
|
||||
def test_contains_sidechaincompress(self):
|
||||
"""包含 sidechaincompress."""
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "sidechaincompress=" in result
|
||||
assert "[1:a][0:a]sidechaincompress" in result
|
||||
|
||||
def test_threshold_param(self):
|
||||
"""threshold 参数正确."""
|
||||
def test_threshold_in_db(self):
|
||||
result = build_sidechain_mix_filter(threshold=-30.0)
|
||||
assert "threshold=-30.0dB" in result
|
||||
|
||||
def test_attack_param(self):
|
||||
"""attack 参数正确."""
|
||||
result = build_sidechain_mix_filter(attack=0.05)
|
||||
assert "attack=0.050" in result
|
||||
|
||||
def test_release_param(self):
|
||||
"""release 参数正确."""
|
||||
result = build_sidechain_mix_filter(release=0.8)
|
||||
assert "release=0.800" in result
|
||||
|
||||
def test_knee_param(self):
|
||||
"""knee=6 参数."""
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "knee=6" in result
|
||||
def test_attack_and_release(self):
|
||||
result = build_sidechain_mix_filter(attack=0.01, release=0.3)
|
||||
assert "attack=0.010" in result
|
||||
assert "release=0.300" in result
|
||||
|
||||
def test_contains_amix(self):
|
||||
"""包含 amix 混音."""
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "amix=inputs=2" in result
|
||||
assert "duration=first" in result
|
||||
|
||||
def test_volume_compensation(self):
|
||||
"""volume=1.5 轻微补偿."""
|
||||
def test_contains_volume_compensation(self):
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "volume=1.5" in result
|
||||
|
||||
def test_bgmc_comp_label(self):
|
||||
"""包含 [bgm_comp] 中间标签."""
|
||||
def test_output_label(self):
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "[final]" in result
|
||||
|
||||
def test_bgm_comp_label(self):
|
||||
result = build_sidechain_mix_filter()
|
||||
assert "[bgm_comp]" in result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# normalize_bgm_config 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── normalize_bgm_config ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestNormalizeBGMConfig:
|
||||
"""配置规范化测试."""
|
||||
|
||||
def test_empty_dict_defaults(self):
|
||||
"""空字典返回默认值."""
|
||||
class TestNormalizeBgmConfig:
|
||||
def test_default_values(self):
|
||||
result = normalize_bgm_config({})
|
||||
assert result["volume"] == 0.3
|
||||
assert result["fade_in"] == 0.0
|
||||
@@ -426,233 +301,184 @@ class TestNormalizeBGMConfig:
|
||||
assert result["loop_enabled"] is True
|
||||
assert result["sidechain_enabled"] is False
|
||||
assert result["sidechain_ratio"] == 0.3
|
||||
assert result["sidechain_attack"] == 0.02
|
||||
assert result["sidechain_release"] == 0.5
|
||||
assert result["sidechain_threshold"] == -25.0
|
||||
|
||||
def test_volume_clamped(self):
|
||||
"""音量钳制."""
|
||||
result = normalize_bgm_config({"volume": 1.5})
|
||||
result = normalize_bgm_config({"volume": 2.0})
|
||||
assert result["volume"] == 1.0
|
||||
result2 = normalize_bgm_config({"volume": -0.5})
|
||||
assert result2["volume"] == 0.0
|
||||
result = normalize_bgm_config({"volume": -1.0})
|
||||
assert result["volume"] == 0.0
|
||||
|
||||
def test_fade_in_negative(self):
|
||||
"""淡入为负钳制到 0."""
|
||||
result = normalize_bgm_config({"fade_in": -1})
|
||||
def test_fade_in_clamped_to_zero(self):
|
||||
result = normalize_bgm_config({"fade_in": -5})
|
||||
assert result["fade_in"] == 0.0
|
||||
|
||||
def test_fade_out_negative(self):
|
||||
"""淡出为负钳制到 0."""
|
||||
result = normalize_bgm_config({"fade_out": -1})
|
||||
def test_fade_out_clamped_to_zero(self):
|
||||
result = normalize_bgm_config({"fade_out": -5})
|
||||
assert result["fade_out"] == 0.0
|
||||
|
||||
def test_sidechain_ratio_clamped(self):
|
||||
"""sidechain_ratio 钳制."""
|
||||
result = normalize_bgm_config({"sidechain_ratio": 1.5})
|
||||
assert result["sidechain_ratio"] == 1.0
|
||||
result2 = normalize_bgm_config({"sidechain_ratio": -0.1})
|
||||
assert result2["sidechain_ratio"] == 0.0
|
||||
def test_loop_enabled_bool_conversion(self):
|
||||
assert normalize_bgm_config({"loop_enabled": True})["loop_enabled"] is True
|
||||
assert normalize_bgm_config({"loop_enabled": False})["loop_enabled"] is False
|
||||
assert normalize_bgm_config({"loop_enabled": 1})["loop_enabled"] is True
|
||||
assert normalize_bgm_config({"loop_enabled": 0})["loop_enabled"] is False
|
||||
|
||||
def test_sidechain_attack_min(self):
|
||||
"""attack 最小值 0.001."""
|
||||
def test_sidechain_ratio_clamped(self):
|
||||
result = normalize_bgm_config({"sidechain_ratio": 2.0})
|
||||
assert result["sidechain_ratio"] == 1.0
|
||||
result = normalize_bgm_config({"sidechain_ratio": -1.0})
|
||||
assert result["sidechain_ratio"] == 0.0
|
||||
|
||||
def test_sidechain_attack_minimum(self):
|
||||
result = normalize_bgm_config({"sidechain_attack": 0})
|
||||
assert result["sidechain_attack"] == 0.001
|
||||
|
||||
def test_sidechain_release_min(self):
|
||||
"""release 最小值 0.01."""
|
||||
def test_sidechain_release_minimum(self):
|
||||
result = normalize_bgm_config({"sidechain_release": 0})
|
||||
assert result["sidechain_release"] == 0.01
|
||||
|
||||
def test_sidechain_threshold_pass_through(self):
|
||||
result = normalize_bgm_config({"sidechain_threshold": -40.0})
|
||||
assert result["sidechain_threshold"] == -40.0
|
||||
|
||||
def test_string_values_converted(self):
|
||||
"""字符串数值被转换."""
|
||||
result = normalize_bgm_config(
|
||||
{
|
||||
"volume": "0.5",
|
||||
"fade_in": "2.0",
|
||||
"fade_in": "1.0",
|
||||
"sidechain_ratio": "0.7",
|
||||
}
|
||||
)
|
||||
assert result["volume"] == 0.5
|
||||
assert result["fade_in"] == 2.0
|
||||
|
||||
def test_loop_enabled_truthy(self):
|
||||
"""loop_enabled 真值转换."""
|
||||
result = normalize_bgm_config({"loop_enabled": 1})
|
||||
assert result["loop_enabled"] is True
|
||||
result2 = normalize_bgm_config({"loop_enabled": 0})
|
||||
assert result2["loop_enabled"] is False
|
||||
|
||||
def test_preserves_unknown_keys(self):
|
||||
"""未知 key 不保留."""
|
||||
result = normalize_bgm_config({"unknown_key": "value", "volume": 0.5})
|
||||
assert "unknown_key" not in result
|
||||
assert result["volume"] == 0.5
|
||||
assert result["fade_in"] == 1.0
|
||||
assert result["sidechain_ratio"] == 0.7
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# validate_bgm_config 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── validate_bgm_config ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateBGMConfig:
|
||||
"""配置验证测试."""
|
||||
|
||||
class TestValidateBgmConfig:
|
||||
def test_valid_config(self):
|
||||
"""合法配置."""
|
||||
ok, errors = validate_bgm_config(
|
||||
{
|
||||
"volume": 0.5,
|
||||
"fade_in": 1.0,
|
||||
"fade_out": 2.0,
|
||||
"sidechain_ratio": 0.3,
|
||||
}
|
||||
)
|
||||
assert ok is True
|
||||
assert len(errors) == 0
|
||||
valid, errors = validate_bgm_config({"volume": 0.3})
|
||||
assert valid is True
|
||||
assert errors == []
|
||||
|
||||
def test_volume_not_number(self):
|
||||
"""volume 不是数字."""
|
||||
ok, errors = validate_bgm_config({"volume": "high"})
|
||||
assert ok is False
|
||||
def test_invalid_volume_type(self):
|
||||
valid, errors = validate_bgm_config({"volume": "abc"})
|
||||
assert valid is False
|
||||
assert any("volume" in e for e in errors)
|
||||
|
||||
def test_volume_out_of_range(self):
|
||||
"""volume 超出范围."""
|
||||
ok, errors = validate_bgm_config({"volume": 1.5})
|
||||
assert ok is False
|
||||
valid, errors = validate_bgm_config({"volume": -0.1})
|
||||
assert valid is False
|
||||
assert any("volume" in e for e in errors)
|
||||
valid, errors = validate_bgm_config({"volume": 1.1})
|
||||
assert valid is False
|
||||
assert any("volume" in e for e in errors)
|
||||
|
||||
def test_fade_in_negative(self):
|
||||
"""fade_in 为负."""
|
||||
ok, errors = validate_bgm_config({"fade_in": -1})
|
||||
assert ok is False
|
||||
def test_volume_at_boundaries(self):
|
||||
assert validate_bgm_config({"volume": 0})[0] is True
|
||||
assert validate_bgm_config({"volume": 1})[0] is True
|
||||
|
||||
def test_invalid_fade_in_type(self):
|
||||
valid, errors = validate_bgm_config({"fade_in": "abc"})
|
||||
assert valid is False
|
||||
assert any("fade_in" in e for e in errors)
|
||||
|
||||
def test_fade_out_negative(self):
|
||||
"""fade_out 为负."""
|
||||
ok, errors = validate_bgm_config({"fade_out": -1})
|
||||
assert ok is False
|
||||
def test_negative_fade_in(self):
|
||||
valid, errors = validate_bgm_config({"fade_in": -1})
|
||||
assert valid is False
|
||||
assert any("fade_in" in e for e in errors)
|
||||
|
||||
def test_invalid_fade_out_type(self):
|
||||
valid, errors = validate_bgm_config({"fade_out": "abc"})
|
||||
assert valid is False
|
||||
assert any("fade_out" in e for e in errors)
|
||||
|
||||
def test_negative_fade_out(self):
|
||||
valid, errors = validate_bgm_config({"fade_out": -1})
|
||||
assert valid is False
|
||||
assert any("fade_out" in e for e in errors)
|
||||
|
||||
def test_invalid_sidechain_ratio_type(self):
|
||||
valid, errors = validate_bgm_config({"sidechain_ratio": "abc"})
|
||||
assert valid is False
|
||||
assert any("sidechain_ratio" in e for e in errors)
|
||||
|
||||
def test_sidechain_ratio_out_of_range(self):
|
||||
"""sidechain_ratio 超出范围."""
|
||||
ok, errors = validate_bgm_config({"sidechain_ratio": 2.0})
|
||||
assert ok is False
|
||||
valid, errors = validate_bgm_config({"sidechain_ratio": -0.1})
|
||||
assert valid is False
|
||||
assert any("sidechain_ratio" in e for e in errors)
|
||||
valid, errors = validate_bgm_config({"sidechain_ratio": 1.1})
|
||||
assert valid is False
|
||||
assert any("sidechain_ratio" in e for e in errors)
|
||||
|
||||
def test_multiple_errors(self):
|
||||
"""多个错误同时报告."""
|
||||
ok, errors = validate_bgm_config(
|
||||
valid, errors = validate_bgm_config(
|
||||
{
|
||||
"volume": 2.0,
|
||||
"volume": "bad",
|
||||
"fade_in": -1,
|
||||
"sidechain_ratio": -0.5,
|
||||
"sidechain_ratio": 2.0,
|
||||
}
|
||||
)
|
||||
assert ok is False
|
||||
assert valid is False
|
||||
assert len(errors) >= 3
|
||||
|
||||
def test_empty_config_valid(self):
|
||||
"""空配置(全用默认值)视为合法."""
|
||||
ok, errors = validate_bgm_config({})
|
||||
assert ok is True
|
||||
assert len(errors) == 0
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# calculate_fade_out_start 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── calculate_fade_out_start ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateFadeOutStart:
|
||||
"""淡出开始时间计算测试."""
|
||||
|
||||
def test_normal_case(self):
|
||||
"""正常情况."""
|
||||
assert calculate_fade_out_start(100, 3) == pytest.approx(97.0)
|
||||
assert calculate_fade_out_start(60, 2) == 58.0
|
||||
|
||||
def test_zero_fade_out(self):
|
||||
"""淡出时长为 0,返回 None."""
|
||||
assert calculate_fade_out_start(100, 0) is None
|
||||
assert calculate_fade_out_start(60, 0) is None
|
||||
|
||||
def test_negative_fade_out(self):
|
||||
"""淡出时长为负,返回 None."""
|
||||
assert calculate_fade_out_start(100, -1) is None
|
||||
|
||||
def test_zero_duration(self):
|
||||
"""总时长为 0,返回 None."""
|
||||
assert calculate_fade_out_start(0, 3) is None
|
||||
|
||||
def test_fade_out_longer_than_duration(self):
|
||||
"""淡出超过总时长,返回 None."""
|
||||
assert calculate_fade_out_start(10, 20) is None
|
||||
|
||||
def test_fade_out_equal_to_duration(self):
|
||||
"""淡出等于总时长,返回 None."""
|
||||
assert calculate_fade_out_start(10, 10) is None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# estimate_bgm_processing_duration 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEstimateBGMProcessingDuration:
|
||||
"""BGM 处理时长估算测试."""
|
||||
|
||||
def test_normal_case_with_loop(self):
|
||||
"""正常循环情况,输出目标时长."""
|
||||
assert estimate_bgm_processing_duration(10, 100, True) == 100
|
||||
|
||||
def test_bgm_longer_no_loop(self):
|
||||
"""BGM 够长,不循环,截断到目标时长."""
|
||||
assert estimate_bgm_processing_duration(200, 100, False) == 100
|
||||
|
||||
def test_bgm_shorter_no_loop(self):
|
||||
"""BGM 短但不循环,仍然截断到目标时长(实际会更短,但 atrim 会截断)."""
|
||||
assert estimate_bgm_processing_duration(10, 100, False) == 100
|
||||
assert calculate_fade_out_start(60, -1) is None
|
||||
|
||||
def test_zero_target(self):
|
||||
"""目标时长为 0,兜底 5 秒."""
|
||||
assert estimate_bgm_processing_duration(10, 0, True) == 5.0
|
||||
assert calculate_fade_out_start(0, 2) is None
|
||||
|
||||
def test_negative_target(self):
|
||||
"""目标时长为负,兜底 5 秒."""
|
||||
assert estimate_bgm_processing_duration(10, -5, True) == 5.0
|
||||
assert calculate_fade_out_start(-5, 2) is None
|
||||
|
||||
def test_fade_longer_than_target(self):
|
||||
assert calculate_fade_out_start(10, 20) is None
|
||||
|
||||
def test_fade_equal_to_target(self):
|
||||
assert calculate_fade_out_start(10, 10) is None
|
||||
|
||||
def test_float_values(self):
|
||||
assert calculate_fade_out_start(60.5, 2.5) == 58.0
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# BGMPureConfig 测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── estimate_bgm_processing_duration ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBGMPureConfig:
|
||||
"""BGMPureConfig 数据类测试."""
|
||||
class TestEstimateBgmProcessingDuration:
|
||||
def test_bgm_longer_no_loop(self):
|
||||
assert estimate_bgm_processing_duration(100, 60, loop_enabled=False) == 60.0
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = BGMPureConfig()
|
||||
assert config.volume == 0.3
|
||||
assert config.fade_in == 0.0
|
||||
assert config.fade_out == 0.0
|
||||
assert config.loop_enabled is True
|
||||
assert config.sidechain_enabled is False
|
||||
assert config.sidechain_ratio == 0.3
|
||||
assert config.sidechain_attack == 0.02
|
||||
assert config.sidechain_release == 0.5
|
||||
assert config.sidechain_threshold == -25.0
|
||||
def test_bgm_longer_with_loop(self):
|
||||
# 够长但允许循环,仍然截断到target
|
||||
assert estimate_bgm_processing_duration(100, 60, loop_enabled=True) == 60.0
|
||||
|
||||
def test_custom_values(self):
|
||||
"""自定义值."""
|
||||
config = BGMPureConfig(
|
||||
volume=0.7,
|
||||
fade_in=1.0,
|
||||
fade_out=2.0,
|
||||
loop_enabled=False,
|
||||
sidechain_enabled=True,
|
||||
sidechain_ratio=0.5,
|
||||
sidechain_attack=0.05,
|
||||
sidechain_release=0.8,
|
||||
sidechain_threshold=-30.0,
|
||||
)
|
||||
assert config.volume == 0.7
|
||||
assert config.loop_enabled is False
|
||||
assert config.sidechain_enabled is True
|
||||
assert config.sidechain_threshold == -30.0
|
||||
def test_bgm_shorter_with_loop(self):
|
||||
assert estimate_bgm_processing_duration(10, 60, loop_enabled=True) == 60.0
|
||||
|
||||
def test_bgm_shorter_no_loop(self):
|
||||
# 需要循环但不允许 → 截断到target
|
||||
assert estimate_bgm_processing_duration(10, 60, loop_enabled=False) == 60.0
|
||||
|
||||
def test_zero_target_fallback(self):
|
||||
assert estimate_bgm_processing_duration(100, 0) == 5.0
|
||||
|
||||
def test_negative_target_fallback(self):
|
||||
assert estimate_bgm_processing_duration(100, -5) == 5.0
|
||||
|
||||
def test_equal_duration(self):
|
||||
assert estimate_bgm_processing_duration(60, 60) == 60.0
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
"""视频拼接引擎纯逻辑单元测试."""
|
||||
"""concat_engine_pure 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from video_processing.concat_engine_pure import (
|
||||
from apps.worker.video_processing.concat_engine_pure import (
|
||||
build_concat_filter,
|
||||
build_fps_filter,
|
||||
build_scale_pad_filter,
|
||||
build_setpts_filter,
|
||||
build_single_segment_filter_chain,
|
||||
calculate_scaled_size,
|
||||
can_use_stream_copy,
|
||||
@@ -20,200 +20,203 @@ from video_processing.concat_engine_pure import (
|
||||
validate_video_path,
|
||||
)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 帧率解析测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── parse_fps ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestParseFps:
|
||||
"""parse_fps 测试."""
|
||||
|
||||
def test_integer_fps(self):
|
||||
"""整数帧率."""
|
||||
assert parse_fps(30) == 30.0
|
||||
|
||||
def test_float_fps(self):
|
||||
"""浮点帧率."""
|
||||
assert parse_fps(29.97) == pytest.approx(29.97)
|
||||
|
||||
def test_string_integer(self):
|
||||
"""字符串整数."""
|
||||
assert parse_fps("30") == 30.0
|
||||
|
||||
def test_string_fraction(self):
|
||||
"""分数字符串(30/1)."""
|
||||
assert parse_fps("30/1") == 30.0
|
||||
|
||||
def test_fraction_24000_1001(self):
|
||||
"""23.976 帧率."""
|
||||
result = parse_fps("24000/1001")
|
||||
assert result == pytest.approx(23.976, rel=0.01)
|
||||
|
||||
def test_none_input(self):
|
||||
"""None 输入返回默认值."""
|
||||
def test_none_returns_default(self):
|
||||
assert parse_fps(None) == 30.0
|
||||
|
||||
def test_empty_string(self):
|
||||
"""空字符串返回默认值."""
|
||||
assert parse_fps("") == 30.0
|
||||
def test_integer_value(self):
|
||||
assert parse_fps(30) == 30.0
|
||||
assert parse_fps(24) == 24.0
|
||||
|
||||
def test_invalid_string(self):
|
||||
"""无效字符串."""
|
||||
assert parse_fps("abc") == 30.0
|
||||
def test_float_value(self):
|
||||
assert parse_fps(29.97) == 29.97
|
||||
|
||||
def test_string_integer(self):
|
||||
assert parse_fps("30") == 30.0
|
||||
assert parse_fps(" 60 ") == 60.0 # 带空格
|
||||
|
||||
def test_string_fraction(self):
|
||||
assert parse_fps("30/1") == 30.0
|
||||
assert abs(parse_fps("24000/1001") - 23.976) < 0.01
|
||||
|
||||
def test_zero_denominator(self):
|
||||
"""分母为 0."""
|
||||
assert parse_fps("30/0") == 30.0
|
||||
|
||||
def test_empty_string(self):
|
||||
assert parse_fps("") == 30.0
|
||||
assert parse_fps(" ") == 30.0
|
||||
|
||||
def test_invalid_string(self):
|
||||
assert parse_fps("abc") == 30.0
|
||||
assert parse_fps("30fps") == 30.0
|
||||
|
||||
def test_negative_fps(self):
|
||||
"""负帧率."""
|
||||
assert parse_fps(-30) == -30.0
|
||||
|
||||
def test_zero_fps(self):
|
||||
assert parse_fps(0) == 0.0
|
||||
|
||||
|
||||
# ── format_fps_filter ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFormatFpsFilter:
|
||||
"""format_fps_filter 测试."""
|
||||
|
||||
def test_integer_fps(self):
|
||||
"""整数帧率."""
|
||||
assert format_fps_filter(30.0) == "fps=30"
|
||||
|
||||
def test_float_fps(self):
|
||||
"""浮点帧率."""
|
||||
def test_near_integer_fps(self):
|
||||
# 接近整数时用整数形式(注意:int(fps)是截断不是四舍五入)
|
||||
assert format_fps_filter(30.0001) == "fps=30"
|
||||
assert format_fps_filter(30.0005) == "fps=30" # int(30.0005)=30
|
||||
|
||||
def test_non_integer_fps(self):
|
||||
result = format_fps_filter(23.976)
|
||||
assert result.startswith("fps=")
|
||||
assert "23.976" in result
|
||||
|
||||
def test_float_precision(self):
|
||||
result = format_fps_filter(29.97)
|
||||
assert result.startswith("fps=")
|
||||
assert "29.97" in result
|
||||
# 三位小数
|
||||
parts = result.split("=")[1]
|
||||
assert len(parts.split(".")[1]) == 3
|
||||
|
||||
def test_near_integer(self):
|
||||
"""接近整数."""
|
||||
assert format_fps_filter(30.0001) == "fps=30"
|
||||
def test_one_fps(self):
|
||||
assert format_fps_filter(1.0) == "fps=1"
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 输出参数计算测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── resolve_output_params ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveOutputParams:
|
||||
"""resolve_output_params 测试."""
|
||||
|
||||
def test_all_specified(self):
|
||||
"""全部显式指定."""
|
||||
def test_config_specified(self):
|
||||
w, h, fps = resolve_output_params(1920, 1080, 60.0)
|
||||
assert w == 1920
|
||||
assert h == 1080
|
||||
assert fps == 60.0
|
||||
|
||||
def test_no_specified_use_defaults(self):
|
||||
"""全部未指定,用默认值."""
|
||||
w, h, fps = resolve_output_params(0, 0, 0)
|
||||
assert w == 1080
|
||||
assert h == 1920
|
||||
assert fps == 30.0
|
||||
|
||||
def test_use_first_video_info(self):
|
||||
"""用第一段视频信息."""
|
||||
def test_fallback_to_first_video_info(self):
|
||||
info = {"width": 1280, "height": 720, "r_frame_rate": "24/1"}
|
||||
w, h, fps = resolve_output_params(0, 0, 0, info)
|
||||
assert w == 1280
|
||||
assert h == 720
|
||||
assert fps == 24.0
|
||||
|
||||
def test_partial_specified(self):
|
||||
"""部分指定,未指定的用探测值."""
|
||||
def test_fallback_to_defaults(self):
|
||||
w, h, fps = resolve_output_params(0, 0, 0)
|
||||
assert w == 1080 # default_width
|
||||
assert h == 1920 # default_height
|
||||
assert fps == 30.0
|
||||
|
||||
def test_partial_config(self):
|
||||
# 宽度配置了,高度和帧率用探测的
|
||||
info = {"width": 1280, "height": 720, "r_frame_rate": "24/1"}
|
||||
w, h, fps = resolve_output_params(1920, 0, 0, info)
|
||||
assert w == 1920 # 指定的
|
||||
assert h == 720 # 探测的
|
||||
assert w == 1920
|
||||
assert h == 720
|
||||
assert fps == 24.0
|
||||
|
||||
def test_zero_size_clamped(self):
|
||||
"""零尺寸被钳制."""
|
||||
w, h, fps = resolve_output_params(0, 0, 0, {})
|
||||
assert w >= 1
|
||||
assert h >= 1
|
||||
assert fps >= 1.0
|
||||
|
||||
def test_custom_defaults(self):
|
||||
"""自定义默认值."""
|
||||
w, h, fps = resolve_output_params(0, 0, 0, None, 640, 480, 25.0)
|
||||
w, h, fps = resolve_output_params(
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
default_width=640,
|
||||
default_height=480,
|
||||
default_fps=25.0,
|
||||
)
|
||||
assert w == 640
|
||||
assert h == 480
|
||||
assert fps == 25.0
|
||||
|
||||
def test_minimum_size(self):
|
||||
w, h, fps = resolve_output_params(0, 0, 0, {"width": 0, "height": 0, "r_frame_rate": "0/1"})
|
||||
assert w >= 1
|
||||
assert h >= 1
|
||||
assert fps >= 1.0
|
||||
|
||||
def test_fps_fraction_in_info(self):
|
||||
info = {"width": 1920, "height": 1080, "r_frame_rate": "24000/1001"}
|
||||
_, _, fps = resolve_output_params(0, 0, 0, info)
|
||||
assert abs(fps - 23.976) < 0.01
|
||||
|
||||
|
||||
# ── calculate_scaled_size ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCalculateScaledSize:
|
||||
"""calculate_scaled_size 测试."""
|
||||
|
||||
def test_same_ratio(self):
|
||||
"""比例相同."""
|
||||
sw, sh, ox, oy = calculate_scaled_size(1920, 1080, 1920, 1080)
|
||||
assert sw == 1920
|
||||
assert sh == 1080
|
||||
assert ox == 0
|
||||
assert oy == 0
|
||||
|
||||
def test_wider_source(self):
|
||||
"""源更宽,上下填黑边."""
|
||||
def test_wider_source_pad_top_bottom(self):
|
||||
# 源是16:9,目标是9:16竖屏 → 上下填黑边
|
||||
sw, sh, ox, oy = calculate_scaled_size(1920, 1080, 1080, 1920)
|
||||
assert sw == 1080 # 以宽度为准
|
||||
assert sh < 1920 # 高度按比例
|
||||
assert sh == 607 # 1080 * 1080 / 1920 = 607.5 → 607
|
||||
assert ox == 0
|
||||
assert oy > 0 # 垂直居中
|
||||
|
||||
def test_taller_source(self):
|
||||
"""源更高,左右填黑边."""
|
||||
def test_taller_source_pad_left_right(self):
|
||||
# 源是9:16竖屏,目标是16:9横屏 → 左右填黑边
|
||||
sw, sh, ox, oy = calculate_scaled_size(1080, 1920, 1920, 1080)
|
||||
assert sh == 1080 # 以高度为准
|
||||
assert sw < 1920 # 宽度按比例
|
||||
assert sw == 607 # 1080 * 1080 / 1920 = 607.5 → 607
|
||||
assert ox > 0 # 水平居中
|
||||
assert oy == 0
|
||||
|
||||
def test_zero_source(self):
|
||||
"""零尺寸源."""
|
||||
sw, sh, ox, oy = calculate_scaled_size(0, 0, 100, 100)
|
||||
assert sw == 100
|
||||
assert sh == 100
|
||||
|
||||
def test_scale_down(self):
|
||||
"""缩小."""
|
||||
sw, sh, ox, oy = calculate_scaled_size(1920, 1080, 640, 360)
|
||||
assert sw == 640
|
||||
assert sh == 360
|
||||
def test_zero_source_size(self):
|
||||
sw, sh, ox, oy = calculate_scaled_size(0, 0, 1920, 1080)
|
||||
assert sw == 1920
|
||||
assert sh == 1080
|
||||
assert ox == 0
|
||||
assert oy == 0
|
||||
|
||||
def test_scale_up(self):
|
||||
"""放大."""
|
||||
def test_negative_source_size(self):
|
||||
sw, sh, ox, oy = calculate_scaled_size(-1, -1, 1920, 1080)
|
||||
assert sw == 1920
|
||||
assert sh == 1080
|
||||
assert ox == 0
|
||||
assert oy == 0
|
||||
|
||||
def test_target_same_ratio_different_size(self):
|
||||
# 比例相同,尺寸不同 → 直接缩放到目标大小
|
||||
sw, sh, ox, oy = calculate_scaled_size(640, 360, 1920, 1080)
|
||||
assert sw == 1920
|
||||
assert sh == 1080
|
||||
assert ox == 0
|
||||
assert oy == 0
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# stream copy 判断测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── can_use_stream_copy ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCanUseStreamCopy:
|
||||
"""can_use_stream_copy 测试."""
|
||||
def test_force_reencode_false(self):
|
||||
assert can_use_stream_copy([], 1920, 1080, 30.0, force_reencode=True) is False
|
||||
|
||||
def test_identical_segments(self):
|
||||
"""所有段参数相同,可以 stream copy."""
|
||||
def test_empty_segments(self):
|
||||
assert can_use_stream_copy([], 1920, 1080, 30.0) is False
|
||||
|
||||
def test_single_segment_matching_params(self):
|
||||
segs = [{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"}]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is True
|
||||
|
||||
def test_multiple_segments_same_params(self):
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is True
|
||||
|
||||
def test_force_reencode(self):
|
||||
"""强制重编码."""
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0, force_reencode=True) is False
|
||||
|
||||
def test_different_codec(self):
|
||||
"""编码不同."""
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
{"codec_name": "hevc", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
@@ -221,7 +224,6 @@ class TestCanUseStreamCopy:
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is False
|
||||
|
||||
def test_different_resolution(self):
|
||||
"""分辨率不同."""
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
{"codec_name": "h264", "width": 1280, "height": 720, "r_frame_rate": "30/1"},
|
||||
@@ -229,306 +231,349 @@ class TestCanUseStreamCopy:
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is False
|
||||
|
||||
def test_different_fps(self):
|
||||
"""帧率不同."""
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "60/1"},
|
||||
]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is False
|
||||
|
||||
def test_target_differs(self):
|
||||
"""目标参数与源不同."""
|
||||
segs = [
|
||||
{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"},
|
||||
]
|
||||
assert can_use_stream_copy(segs, 1280, 720, 30.0) is False
|
||||
|
||||
def test_empty_segments(self):
|
||||
"""空列表."""
|
||||
assert can_use_stream_copy([], 1920, 1080, 30.0) is False
|
||||
|
||||
def test_single_segment(self):
|
||||
"""单段."""
|
||||
def test_target_differs_from_source(self):
|
||||
segs = [{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "30/1"}]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 30.0) is True
|
||||
# 目标分辨率不同
|
||||
assert can_use_stream_copy(segs, 1280, 720, 30.0) is False
|
||||
# 目标帧率不同
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 60.0) is False
|
||||
|
||||
def test_fps_fraction_match(self):
|
||||
segs = [{"codec_name": "h264", "width": 1920, "height": 1080, "r_frame_rate": "24000/1001"}]
|
||||
assert can_use_stream_copy(segs, 1920, 1080, 23.976) is True
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 文件列表生成测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── generate_concat_file_list ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerateConcatFileList:
|
||||
"""generate_concat_file_list 测试."""
|
||||
|
||||
def test_single_file(self):
|
||||
"""单个文件."""
|
||||
result = generate_concat_file_list(["/a.mp4"])
|
||||
assert "file '/a.mp4'" in result
|
||||
assert result.endswith("\n")
|
||||
result = generate_concat_file_list(["/tmp/video.mp4"])
|
||||
assert result == "file '/tmp/video.mp4'\n"
|
||||
|
||||
def test_multiple_files(self):
|
||||
"""多个文件."""
|
||||
result = generate_concat_file_list(["/a.mp4", "/b.mp4", "/c.mp4"])
|
||||
lines = result.strip().split("\n")
|
||||
assert len(lines) == 3
|
||||
assert lines[0] == "file '/a.mp4'"
|
||||
assert lines[1] == "file '/b.mp4'"
|
||||
assert lines[2] == "file '/c.mp4'"
|
||||
assert result.endswith("\n")
|
||||
|
||||
def test_escapes_single_quotes(self):
|
||||
result = generate_concat_file_list(["/path/with'quote.mp4"])
|
||||
# 单引号转义: '\''
|
||||
assert "'\\''" in result
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表."""
|
||||
result = generate_concat_file_list([])
|
||||
assert result == "\n"
|
||||
|
||||
def test_path_with_single_quote(self):
|
||||
"""路径包含单引号(转义)."""
|
||||
result = generate_concat_file_list(["/path/to/file's.mp4"])
|
||||
# 单引号应该被转义
|
||||
assert "'\\''" in result or file
|
||||
assert "file '" in result
|
||||
|
||||
def test_path_with_spaces(self):
|
||||
"""路径包含空格."""
|
||||
result = generate_concat_file_list(["/path/to/my video.mp4"])
|
||||
assert "my video" in result
|
||||
result = generate_concat_file_list(["/path/to/video file.mp4"])
|
||||
assert "file '/path/to/video file.mp4'" in result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 滤镜构建测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── build_scale_pad_filter ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildScalePadFilter:
|
||||
"""scale+pad 滤镜测试."""
|
||||
|
||||
def test_contains_scale(self):
|
||||
"""包含 scale."""
|
||||
result = build_scale_pad_filter(1920, 1080)
|
||||
assert "scale=" in result
|
||||
|
||||
def test_contains_pad(self):
|
||||
"""包含 pad."""
|
||||
result = build_scale_pad_filter(1920, 1080)
|
||||
assert "pad=" in result
|
||||
assert "1920:1080" in result
|
||||
|
||||
def test_force_original_aspect_ratio(self):
|
||||
"""保持宽高比."""
|
||||
def test_basic_filter(self):
|
||||
result = build_scale_pad_filter(1920, 1080)
|
||||
assert "scale=1920:1080" in result
|
||||
assert "force_original_aspect_ratio=decrease" in result
|
||||
assert "pad=1920:1080" in result
|
||||
assert "black" in result
|
||||
assert "(ow-iw)/2" in result
|
||||
assert "(oh-ih)/2" in result
|
||||
|
||||
def test_black_padding(self):
|
||||
"""黑边填充."""
|
||||
result = build_scale_pad_filter(1920, 1080)
|
||||
assert ":black" in result
|
||||
def test_different_resolution(self):
|
||||
result = build_scale_pad_filter(1080, 1920)
|
||||
assert "scale=1080:1920" in result
|
||||
assert "pad=1080:1920" in result
|
||||
|
||||
def test_ignores_source_size(self):
|
||||
# src_w/src_h 目前不影响输出,都是用表达式
|
||||
result1 = build_scale_pad_filter(1920, 1080)
|
||||
result2 = build_scale_pad_filter(1920, 1080, src_w=1280, src_h=720)
|
||||
assert result1 == result2
|
||||
|
||||
|
||||
# ── build_fps_filter ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildFpsFilter:
|
||||
"""fps 滤镜测试."""
|
||||
|
||||
def test_integer_fps(self):
|
||||
"""整数帧率."""
|
||||
assert build_fps_filter(30.0) == "fps=30"
|
||||
|
||||
def test_float_fps(self):
|
||||
"""浮点帧率."""
|
||||
result = build_fps_filter(29.97)
|
||||
assert result.startswith("fps=")
|
||||
|
||||
|
||||
class TestBuildConcatFilter:
|
||||
"""concat 滤镜测试."""
|
||||
# ── build_setpts_filter ─────────────────────────────────────────────────────
|
||||
|
||||
def test_two_inputs_with_audio(self):
|
||||
"""两路输入,有音频."""
|
||||
result = build_concat_filter(2, has_audio=True)
|
||||
assert "[0:v][0:a][1:v][1:a]concat=n=2:v=1:a=1" in result
|
||||
|
||||
class TestBuildSetptsFilter:
|
||||
def test_returns_correct_string(self):
|
||||
assert build_setpts_filter() == "setpts=PTS-STARTPTS"
|
||||
|
||||
|
||||
# ── build_concat_filter ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildConcatFilter:
|
||||
def test_zero_inputs(self):
|
||||
assert build_concat_filter(0) == ""
|
||||
|
||||
def test_single_input_with_audio(self):
|
||||
result = build_concat_filter(1)
|
||||
assert "[0:v][0:a]" in result
|
||||
assert "concat=n=1:v=1:a=1" in result
|
||||
assert "[concat_v][concat_a]" in result
|
||||
|
||||
def test_three_inputs_video_only(self):
|
||||
"""三路输入,无音频."""
|
||||
result = build_concat_filter(3, has_audio=False)
|
||||
assert "[0:v][1:v][2:v]concat=n=3:v=1:a=0" in result
|
||||
def test_single_input_no_audio(self):
|
||||
result = build_concat_filter(1, has_audio=False)
|
||||
assert "[0:v]" in result
|
||||
assert "concat=n=1:v=1:a=0" in result
|
||||
assert "[concat_v]" in result
|
||||
assert "[concat_a]" not in result
|
||||
|
||||
def test_single_input(self):
|
||||
"""单路输入."""
|
||||
result = build_concat_filter(1, has_audio=True)
|
||||
assert "[0:v][0:a]concat=n=1:v=1:a=1" in result
|
||||
def test_multiple_inputs_with_audio(self):
|
||||
result = build_concat_filter(3)
|
||||
assert "[0:v][0:a][1:v][1:a][2:v][2:a]" in result
|
||||
assert "concat=n=3:v=1:a=1" in result
|
||||
|
||||
def test_zero_inputs(self):
|
||||
"""零输入."""
|
||||
assert build_concat_filter(0) == ""
|
||||
def test_multiple_inputs_no_audio(self):
|
||||
result = build_concat_filter(3, has_audio=False)
|
||||
assert "[0:v][1:v][2:v]" in result
|
||||
assert "concat=n=3:v=1:a=0" in result
|
||||
|
||||
def test_negative_inputs(self):
|
||||
assert build_concat_filter(-1) == ""
|
||||
|
||||
|
||||
# ── build_single_segment_filter_chain ───────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildSingleSegmentFilterChain:
|
||||
"""单段滤镜链测试."""
|
||||
|
||||
def test_with_audio(self):
|
||||
"""有音频."""
|
||||
result = build_single_segment_filter_chain(1920, 1080, 30.0, 0)
|
||||
assert "scale=" in result
|
||||
assert "fps=" in result
|
||||
assert "setpts=PTS-STARTPTS" in result
|
||||
assert "asetpts=PTS-STARTPTS" in result
|
||||
# 视频链
|
||||
assert "[0:v]" in result
|
||||
assert "[v0]" in result
|
||||
assert "scale=1920:1080" in result
|
||||
assert "fps=30" in result
|
||||
assert "setpts=PTS-STARTPTS" in result
|
||||
# 音频链
|
||||
assert "[0:a]" in result
|
||||
assert "[a0]" in result
|
||||
assert "asetpts=PTS-STARTPTS" in result
|
||||
# 用分号分隔
|
||||
assert ";" in result
|
||||
|
||||
def test_video_only(self):
|
||||
"""无音频."""
|
||||
result = build_single_segment_filter_chain(1920, 1080, 30.0, 1, has_audio=False)
|
||||
assert "scale=" in result
|
||||
assert "setpts=" in result
|
||||
assert "asetpts" not in result
|
||||
assert "[v1]" in result
|
||||
def test_without_audio(self):
|
||||
result = build_single_segment_filter_chain(1920, 1080, 30.0, 2, has_audio=False)
|
||||
assert "[2:v]" in result
|
||||
assert "[v2]" in result
|
||||
assert "[2:a]" not in result
|
||||
assert ";" not in result # 没有音频就没有分号
|
||||
|
||||
def test_segment_index_in_labels(self):
|
||||
"""段索引在标签中."""
|
||||
result = build_single_segment_filter_chain(1920, 1080, 30.0, 5)
|
||||
assert "[5:v]" in result
|
||||
assert "[v5]" in result
|
||||
def test_segment_index_propagated(self):
|
||||
for idx in [0, 5, 10]:
|
||||
result = build_single_segment_filter_chain(1920, 1080, 30.0, idx)
|
||||
assert f"[{idx}:v]" in result
|
||||
assert f"[v{idx}]" in result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 配置验证测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── validate_concat_config ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateConcatConfig:
|
||||
"""配置验证测试."""
|
||||
|
||||
def test_valid_config(self):
|
||||
"""合法配置."""
|
||||
config = {
|
||||
"segments": [{"video_path": "/a.mp4"}, {"video_path": "/b.mp4"}],
|
||||
"segments": [
|
||||
{"video_path": "/a.mp4"},
|
||||
{"video_path": "/b.mp4"},
|
||||
],
|
||||
"output_width": 1920,
|
||||
"output_height": 1080,
|
||||
"output_fps": 30,
|
||||
}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is True
|
||||
assert len(errors) == 0
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is True
|
||||
assert errors == []
|
||||
|
||||
def test_no_segments(self):
|
||||
valid, errors = validate_concat_config({})
|
||||
assert valid is False
|
||||
assert any("至少需要一个" in e for e in errors)
|
||||
|
||||
def test_empty_segments(self):
|
||||
"""空段列表."""
|
||||
ok, errors = validate_concat_config({"segments": []})
|
||||
assert ok is False
|
||||
assert any("至少需要" in e or "视频段" in e for e in errors)
|
||||
valid, errors = validate_concat_config({"segments": []})
|
||||
assert valid is False
|
||||
assert len(errors) >= 1
|
||||
|
||||
def test_missing_video_path(self):
|
||||
"""缺少 video_path."""
|
||||
config = {"segments": [{"video_path": "/a.mp4"}, {}]}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is False
|
||||
config = {"segments": [{"video_path": ""}]}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is False
|
||||
assert any("video_path" in e for e in errors)
|
||||
|
||||
def test_negative_width(self):
|
||||
"""负宽度."""
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_width": -100}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is False
|
||||
def test_multiple_missing_paths(self):
|
||||
config = {
|
||||
"segments": [
|
||||
{"video_path": "/a.mp4"},
|
||||
{"video_path": ""},
|
||||
{"video_path": ""},
|
||||
]
|
||||
}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is False
|
||||
path_errors = [e for e in errors if "video_path" in e]
|
||||
assert len(path_errors) == 2
|
||||
|
||||
def test_negative_output_width(self):
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_width": -1}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is False
|
||||
assert any("output_width" in e for e in errors)
|
||||
|
||||
def test_negative_height(self):
|
||||
"""负高度."""
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_height": -100}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is False
|
||||
def test_negative_output_height(self):
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_height": -1}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is False
|
||||
assert any("output_height" in e for e in errors)
|
||||
|
||||
def test_negative_fps(self):
|
||||
"""负帧率."""
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_fps": -30}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is False
|
||||
def test_negative_output_fps(self):
|
||||
config = {"segments": [{"video_path": "/a.mp4"}], "output_fps": -1}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is False
|
||||
assert any("output_fps" in e for e in errors)
|
||||
|
||||
def test_zero_output_params_ok(self):
|
||||
"""零输出参数合法(表示自动探测)."""
|
||||
config = {"segments": [{"video_path": "/a.mp4"}]}
|
||||
ok, errors = validate_concat_config(config)
|
||||
assert ok is True
|
||||
def test_zero_output_params_valid(self):
|
||||
# 0值表示未指定,是合法的
|
||||
config = {
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"output_width": 0,
|
||||
"output_height": 0,
|
||||
"output_fps": 0,
|
||||
}
|
||||
valid, errors = validate_concat_config(config)
|
||||
assert valid is True
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 路径验证测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── validate_video_path ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateVideoPath:
|
||||
"""视频路径验证测试."""
|
||||
|
||||
def test_empty_path(self):
|
||||
"""空路径."""
|
||||
ok, msg = validate_video_path("", "/work")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
valid, err = validate_video_path("", "/work")
|
||||
assert valid is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_path_traversal(self):
|
||||
"""路径遍历."""
|
||||
ok, msg = validate_video_path("../etc/passwd", "/work")
|
||||
assert ok is False
|
||||
assert "回溯" in msg or ".." in msg
|
||||
def test_relative_path_valid(self):
|
||||
valid, err = validate_video_path("video.mp4", "/work")
|
||||
assert valid is True
|
||||
assert err == ""
|
||||
|
||||
def test_valid_relative_path(self):
|
||||
"""相对路径(不检查边界)."""
|
||||
ok, msg = validate_video_path("video.mp4", "/work")
|
||||
assert ok is True
|
||||
def test_relative_path_with_subdir(self):
|
||||
valid, err = validate_video_path("sub/video.mp4", "/work")
|
||||
assert valid is True
|
||||
|
||||
def test_valid_absolute_path(self):
|
||||
"""绝对路径在工作目录内."""
|
||||
ok, msg = validate_video_path("/work/sub/video.mp4", "/work")
|
||||
assert ok is True
|
||||
def test_path_traversal_rejected(self):
|
||||
valid, err = validate_video_path("../secret.mp4", "/work")
|
||||
assert valid is False
|
||||
assert ".." in err
|
||||
|
||||
def test_path_outside_work_dir(self):
|
||||
"""路径在工作目录外."""
|
||||
ok, msg = validate_video_path("/etc/passwd", "/work")
|
||||
assert ok is False
|
||||
assert "工作目录" in msg
|
||||
def test_nested_path_traversal_rejected(self):
|
||||
valid, err = validate_video_path("sub/../../secret.mp4", "/work")
|
||||
assert valid is False
|
||||
|
||||
def test_absolute_path_inside_workdir(self):
|
||||
valid, err = validate_video_path("/work/sub/video.mp4", "/work")
|
||||
assert valid is True
|
||||
|
||||
def test_absolute_path_outside_workdir(self):
|
||||
valid, err = validate_video_path("/etc/passwd", "/work")
|
||||
assert valid is False
|
||||
assert "工作目录内" in err
|
||||
|
||||
def test_path_object_input(self):
|
||||
valid, err = validate_video_path(Path("video.mp4"), Path("/work"))
|
||||
assert valid is True
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# 工具函数测试
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# ── estimate_total_duration ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEstimateTotalDuration:
|
||||
"""总时长估算测试."""
|
||||
def test_single_segment(self):
|
||||
assert estimate_total_duration([{"duration": 10.5}]) == 10.5
|
||||
|
||||
def test_multiple_segments(self):
|
||||
"""多段视频."""
|
||||
segs = [{"duration": 10}, {"duration": 20.5}, {"duration": 5}]
|
||||
assert estimate_total_duration(segs) == pytest.approx(35.5)
|
||||
segs = [
|
||||
{"duration": 10},
|
||||
{"duration": 20.5},
|
||||
{"duration": 5.5},
|
||||
]
|
||||
assert estimate_total_duration(segs) == 36.0
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表."""
|
||||
assert estimate_total_duration([]) == 0.0
|
||||
|
||||
def test_invalid_duration_skipped(self):
|
||||
"""无效时长跳过."""
|
||||
segs = [{"duration": 10}, {"duration": "abc"}, {"duration": 20}]
|
||||
assert estimate_total_duration(segs) == pytest.approx(30.0)
|
||||
def test_missing_duration_field(self):
|
||||
segs = [{"path": "a.mp4"}, {"duration": 10}]
|
||||
assert estimate_total_duration(segs) == 10.0
|
||||
|
||||
def test_missing_duration(self):
|
||||
"""缺 duration 字段."""
|
||||
segs = [{}, {"duration": 10}]
|
||||
assert estimate_total_duration(segs) == pytest.approx(10.0)
|
||||
def test_invalid_duration_skipped(self):
|
||||
segs = [
|
||||
{"duration": 10},
|
||||
{"duration": "abc"},
|
||||
{"duration": 20},
|
||||
]
|
||||
assert estimate_total_duration(segs) == 30.0
|
||||
|
||||
def test_string_duration(self):
|
||||
segs = [{"duration": "15.5"}]
|
||||
assert estimate_total_duration(segs) == 15.5
|
||||
|
||||
def test_negative_duration(self):
|
||||
segs = [{"duration": -5}]
|
||||
assert estimate_total_duration(segs) == -5.0
|
||||
|
||||
|
||||
# ── count_valid_segments ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCountValidSegments:
|
||||
"""有效段统计测试."""
|
||||
|
||||
def test_all_valid(self):
|
||||
"""全部有效."""
|
||||
segs = [{"video_path": "/a.mp4"}, {"video_path": "/b.mp4"}]
|
||||
segs = [
|
||||
{"video_path": "/a.mp4"},
|
||||
{"video_path": "/b.mp4"},
|
||||
]
|
||||
assert count_valid_segments(segs) == 2
|
||||
|
||||
def test_some_invalid(self):
|
||||
"""部分无效."""
|
||||
segs = [{"video_path": "/a.mp4"}, {}, {"video_path": ""}]
|
||||
assert count_valid_segments(segs) == 1
|
||||
segs = [
|
||||
{"video_path": "/a.mp4"},
|
||||
{"video_path": ""},
|
||||
{"video_path": "/c.mp4"},
|
||||
]
|
||||
assert count_valid_segments(segs) == 2
|
||||
|
||||
def test_none_valid(self):
|
||||
segs = [
|
||||
{"video_path": ""},
|
||||
{"other_field": "x"},
|
||||
]
|
||||
assert count_valid_segments(segs) == 0
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表."""
|
||||
assert count_valid_segments([]) == 0
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+79
-288
@@ -1,9 +1,6 @@
|
||||
"""通用分页器单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
"""pagination 单元测试."""
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from packages.application.common.pagination import (
|
||||
PaginatedResponse,
|
||||
@@ -12,377 +9,171 @@ from packages.application.common.pagination import (
|
||||
paginate,
|
||||
)
|
||||
|
||||
# ── PaginationParams ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationParams:
|
||||
"""PaginationParams 测试"""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确"""
|
||||
params = PaginationParams()
|
||||
assert params.page == 1
|
||||
assert params.page_size == 20
|
||||
|
||||
def test_offset_first_page(self):
|
||||
"""第一页 offset 为 0"""
|
||||
def test_custom_values(self):
|
||||
params = PaginationParams(page=3, page_size=50)
|
||||
assert params.page == 3
|
||||
assert params.page_size == 50
|
||||
|
||||
def test_offset_calculation(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
assert params.offset == 0
|
||||
|
||||
def test_offset_second_page(self):
|
||||
"""第二页 offset 计算正确"""
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
assert params.offset == 20
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
assert params.offset == 40
|
||||
|
||||
def test_offset_custom_page_size(self):
|
||||
"""自定义 page_size 的 offset"""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
assert params.offset == 20
|
||||
params = PaginationParams(page=10, page_size=50)
|
||||
assert params.offset == 450
|
||||
|
||||
def test_limit_equals_page_size(self):
|
||||
"""limit 等于 page_size"""
|
||||
params = PaginationParams(page_size=50)
|
||||
assert params.limit == 50
|
||||
params = PaginationParams(page_size=30)
|
||||
assert params.limit == 30
|
||||
|
||||
def test_page_must_be_at_least_1(self):
|
||||
"""page 不能小于 1"""
|
||||
with pytest.raises(ValidationError):
|
||||
with pytest.raises(ValueError):
|
||||
PaginationParams(page=0)
|
||||
|
||||
def test_page_negative_raises(self):
|
||||
"""page 不能为负数"""
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page=-1)
|
||||
|
||||
def test_page_size_must_be_at_least_1(self):
|
||||
"""page_size 不能小于 1"""
|
||||
with pytest.raises(ValidationError):
|
||||
with pytest.raises(ValueError):
|
||||
PaginationParams(page_size=0)
|
||||
|
||||
def test_page_size_max_100(self):
|
||||
"""page_size 最大 100"""
|
||||
with pytest.raises(ValidationError):
|
||||
with pytest.raises(ValueError):
|
||||
PaginationParams(page_size=101)
|
||||
|
||||
def test_page_size_100_is_valid(self):
|
||||
"""page_size=100 是合法的"""
|
||||
params = PaginationParams(page_size=100)
|
||||
assert params.page_size == 100
|
||||
|
||||
# ── PaginationMeta ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationMeta:
|
||||
"""PaginationMeta 测试"""
|
||||
|
||||
def test_from_params_first_page(self):
|
||||
"""第一页元数据"""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=25)
|
||||
|
||||
assert meta.page == 1
|
||||
assert meta.page_size == 10
|
||||
assert meta.total == 25
|
||||
assert meta.total_pages == 3
|
||||
assert meta.total_pages == 3 # ceil(25/10)
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_from_params_last_page(self):
|
||||
"""最后一页元数据"""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=25)
|
||||
|
||||
assert meta.page == 3
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_from_params_middle_page(self):
|
||||
"""中间页元数据"""
|
||||
params = PaginationParams(page=2, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
|
||||
assert meta.page == 2
|
||||
assert meta.total_pages == 5
|
||||
meta = PaginationMeta.from_params(params, total=25)
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_from_params_single_page(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=5)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_from_params_zero_total(self):
|
||||
"""总数为 0 时"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
|
||||
assert meta.total == 0
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_from_params_exact_multiple(self):
|
||||
"""总数刚好是 page_size 的整数倍"""
|
||||
def test_from_params_exact_page_size(self):
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=30)
|
||||
|
||||
assert meta.total_pages == 3
|
||||
|
||||
def test_from_params_single_page(self):
|
||||
"""单页即可放下所有数据"""
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
|
||||
meta = PaginationMeta.from_params(params, total=10)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_from_params_one_extra(self):
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=11)
|
||||
assert meta.total_pages == 2
|
||||
|
||||
|
||||
# ── PaginatedResponse ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginatedResponse:
|
||||
"""PaginatedResponse 测试"""
|
||||
|
||||
def test_create_success(self):
|
||||
"""创建分页响应"""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
data = [1, 2, 3]
|
||||
|
||||
response = PaginatedResponse.create(data, params, total=25)
|
||||
|
||||
assert response.data == [1, 2, 3]
|
||||
def test_create_response(self):
|
||||
params = PaginationParams(page=1, page_size=5)
|
||||
data = [1, 2, 3, 4, 5]
|
||||
response = PaginatedResponse.create(data, params, total=15)
|
||||
assert response.data == data
|
||||
assert response.pagination.page == 1
|
||||
assert response.pagination.total == 25
|
||||
assert response.pagination.total == 15
|
||||
assert response.pagination.total_pages == 3
|
||||
|
||||
def test_create_empty_data(self):
|
||||
"""空数据分页响应"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
response = PaginatedResponse.create([], params, total=0)
|
||||
|
||||
assert response.data == []
|
||||
assert response.pagination.total == 0
|
||||
assert response.pagination.total_pages == 0
|
||||
# ── paginate function ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginateFunction:
|
||||
"""paginate 函数测试(内存分页)"""
|
||||
|
||||
class TestPaginate:
|
||||
def test_first_page(self):
|
||||
"""第一页分页"""
|
||||
items = list(range(30))
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert result.data == list(range(10))
|
||||
assert result.pagination.total == 30
|
||||
assert result.pagination.total_pages == 3
|
||||
assert result.pagination.has_next is True
|
||||
assert result.pagination.has_prev is False
|
||||
|
||||
def test_second_page(self):
|
||||
"""第二页分页"""
|
||||
items = list(range(30))
|
||||
params = PaginationParams(page=2, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert result.data == list(range(10, 20))
|
||||
assert result.pagination.page == 2
|
||||
|
||||
def test_last_page(self):
|
||||
"""最后一页分页"""
|
||||
items = list(range(25))
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert result.data == list(range(20, 25))
|
||||
assert len(result.data) == 5
|
||||
assert result.pagination.has_next is False
|
||||
assert result.pagination.has_prev is True
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表分页"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
result = paginate([], params)
|
||||
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
|
||||
def test_page_beyond_total(self):
|
||||
"""页码超出总数"""
|
||||
def test_single_page(self):
|
||||
items = list(range(5))
|
||||
params = PaginationParams(page=10, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 5
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_custom_page_size(self):
|
||||
"""自定义每页数量"""
|
||||
items = list(range(100))
|
||||
params = PaginationParams(page=1, page_size=50)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert len(result.data) == 50
|
||||
assert result.pagination.total_pages == 2
|
||||
|
||||
def test_single_item(self):
|
||||
"""单条数据"""
|
||||
items = ["only_one"]
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert result.data == ["only_one"]
|
||||
assert result.pagination.total == 1
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_generic_type_preserved(self):
|
||||
"""泛型类型数据正确"""
|
||||
items = [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert len(result.data) == 2
|
||||
assert result.data[0]["id"] == 1
|
||||
|
||||
|
||||
# ── PaginationParams 补充边界 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationParamsEdgeCases:
|
||||
"""PaginationParams 补充边界场景."""
|
||||
|
||||
def test_page_size_1_minimum(self):
|
||||
"""page_size=1 是允许的最小值."""
|
||||
params = PaginationParams(page_size=1)
|
||||
assert params.page_size == 1
|
||||
assert params.limit == 1
|
||||
|
||||
def test_page_size_100_maximum(self):
|
||||
"""page_size=100 是允许的最大值."""
|
||||
params = PaginationParams(page_size=100)
|
||||
assert params.page_size == 100
|
||||
|
||||
def test_offset_page_1_size_100(self):
|
||||
"""第1页每页100条 offset=0."""
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
assert params.offset == 0
|
||||
|
||||
def test_offset_page_100_size_100(self):
|
||||
"""第100页每页100条 offset=9900."""
|
||||
params = PaginationParams(page=100, page_size=100)
|
||||
assert params.offset == 9900
|
||||
|
||||
def test_large_page_number_accepted(self):
|
||||
"""极大页码(超过实际页数)允许."""
|
||||
params = PaginationParams(page=999999, page_size=20)
|
||||
assert params.page == 999999
|
||||
assert params.offset == (999999 - 1) * 20
|
||||
|
||||
|
||||
# ── PaginationMeta 补充边界 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationMetaEdgeCases:
|
||||
"""PaginationMeta 补充边界场景."""
|
||||
|
||||
def test_total_0_page_1(self):
|
||||
"""total=0, page=1 时 total_pages=0, 无上下页."""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_total_0_page_beyond(self):
|
||||
"""total=0, page>1 时 has_prev=True(因为page>1)."""
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_exact_last_page(self):
|
||||
"""刚好是最后一页时 has_next=False."""
|
||||
params = PaginationParams(page=5, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 5
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_one_more_than_exact(self):
|
||||
"""比整数页多1条时总页数+1."""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=51)
|
||||
assert meta.total_pages == 6
|
||||
|
||||
def test_page_exactly_total_pages(self):
|
||||
"""page == total_pages 时 has_next=False."""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=30)
|
||||
assert meta.has_next is False
|
||||
|
||||
def test_total_1_page_1_size_1(self):
|
||||
"""1条数据1页."""
|
||||
params = PaginationParams(page=1, page_size=1)
|
||||
meta = PaginationMeta.from_params(params, total=1)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
|
||||
# ── paginate 补充边界 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginateEdgeCases:
|
||||
"""paginate 补充边界场景."""
|
||||
|
||||
def test_single_item_list(self):
|
||||
"""单元素列表."""
|
||||
result = paginate([42], PaginationParams(page=1, page_size=10))
|
||||
assert result.data == [42]
|
||||
assert result.pagination.total == 1
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_page_exactly_last(self):
|
||||
"""刚好在最后一页."""
|
||||
items = list(range(25))
|
||||
result = paginate(items, PaginationParams(page=3, page_size=10))
|
||||
assert result.data == list(range(20, 25))
|
||||
assert result.pagination.has_next is False
|
||||
|
||||
def test_page_past_end_returns_empty(self):
|
||||
"""页码超过总数返回空."""
|
||||
items = list(range(5))
|
||||
result = paginate(items, PaginationParams(page=10, page_size=10))
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 5
|
||||
|
||||
def test_empty_list_page_1(self):
|
||||
"""空列表第1页."""
|
||||
result = paginate([], PaginationParams(page=1, page_size=10))
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
|
||||
def test_page_size_1_iterates_all(self):
|
||||
"""page_size=1 时每页1条."""
|
||||
items = ["a", "b", "c"]
|
||||
r1 = paginate(items, PaginationParams(page=1, page_size=1))
|
||||
r2 = paginate(items, PaginationParams(page=2, page_size=1))
|
||||
r3 = paginate(items, PaginationParams(page=3, page_size=1))
|
||||
assert r1.data == ["a"]
|
||||
assert r2.data == ["b"]
|
||||
assert r3.data == ["c"]
|
||||
|
||||
def test_does_not_mutate_input(self):
|
||||
"""不修改输入列表."""
|
||||
items = [1, 2, 3, 4, 5]
|
||||
original = items[:]
|
||||
paginate(items, PaginationParams(page=1, page_size=2))
|
||||
assert items == original
|
||||
|
||||
def test_page_size_greater_than_total(self):
|
||||
"""每页条数大于总数."""
|
||||
items = list(range(5))
|
||||
result = paginate(items, PaginationParams(page=1, page_size=100))
|
||||
assert result.data == items
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_empty_list(self):
|
||||
items = []
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
result = paginate(items, params)
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
|
||||
def test_page_beyond_end(self):
|
||||
items = list(range(5))
|
||||
params = PaginationParams(page=10, page_size=10)
|
||||
result = paginate(items, params)
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 5
|
||||
|
||||
def test_page_size_larger_than_items(self):
|
||||
items = list(range(5))
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
result = paginate(items, params)
|
||||
assert result.data == items
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_middle_page(self):
|
||||
items = list(range(100))
|
||||
params = PaginationParams(page=5, page_size=10)
|
||||
result = paginate(items, params)
|
||||
assert result.data == list(range(40, 50))
|
||||
assert result.pagination.has_next is True
|
||||
assert result.pagination.has_prev is True
|
||||
|
||||
+413
-699
File diff suppressed because it is too large
Load Diff
Executable
+453
@@ -0,0 +1,453 @@
|
||||
"""shared.ai_service 单元测试.
|
||||
|
||||
主要测试纯逻辑部分:_parse_recommend_response / _fallback_recommend_clips / _call_ai_cover_service.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from shared.ai_service import (
|
||||
_call_ai_cover_service,
|
||||
_fallback_recommend_clips,
|
||||
_parse_recommend_response,
|
||||
)
|
||||
|
||||
# ── _parse_recommend_response 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestParseRecommendResponseBasic:
|
||||
"""基础解析测试."""
|
||||
|
||||
def test_parse_valid_json(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{
|
||||
"clip_type": "intro",
|
||||
"order": 0,
|
||||
"text_content": "开场",
|
||||
"duration": 3.0,
|
||||
"transition_effect": "fade",
|
||||
"asset_id": "asset1",
|
||||
"start_time": 0.0,
|
||||
"config": {},
|
||||
},
|
||||
{
|
||||
"clip_type": "outro",
|
||||
"order": 1,
|
||||
"text_content": "结尾",
|
||||
"duration": 2.0,
|
||||
"transition_effect": "fade",
|
||||
"asset_id": "",
|
||||
"start_time": 0.0,
|
||||
"config": {},
|
||||
},
|
||||
],
|
||||
"title": "测试视频",
|
||||
"confidence": 0.85,
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["asset1"], 30.0)
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 2
|
||||
assert result["confidence"] == 0.85
|
||||
assert result["total_duration"] == 5.0
|
||||
assert result["config"]["title"]["text"] == "测试视频"
|
||||
assert result["config"]["title"]["ai_auto"] is True
|
||||
|
||||
def test_parse_none_returns_none(self):
|
||||
result = _parse_recommend_response(None, ["a1"], 30.0) # type: ignore[arg-type]
|
||||
assert result is None
|
||||
|
||||
def test_parse_empty_string_returns_none(self):
|
||||
result = _parse_recommend_response("", ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
def test_parse_whitespace_only_returns_none(self):
|
||||
result = _parse_recommend_response(" ", ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
def test_parse_invalid_json_returns_none(self):
|
||||
result = _parse_recommend_response("not json", ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
def test_parse_non_dict_json_returns_none(self):
|
||||
result = _parse_recommend_response("[1, 2, 3]", ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestParseRecommendResponseClips:
|
||||
"""clips 解析测试."""
|
||||
|
||||
def test_parse_no_clips_returns_none(self):
|
||||
content = json.dumps({"title": "test", "clips": []})
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
def test_parse_clips_not_list_returns_none(self):
|
||||
content = json.dumps({"clips": "not a list"})
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is None
|
||||
|
||||
def test_parse_clips_sorted_by_order(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "outro", "order": 2, "duration": 2, "asset_id": "a1"},
|
||||
{"clip_type": "intro", "order": 0, "duration": 3, "asset_id": "a1"},
|
||||
{"clip_type": "showcase", "order": 1, "duration": 5, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 3
|
||||
assert result["clips"][0]["clip_type"] == "intro"
|
||||
assert result["clips"][1]["clip_type"] == "showcase"
|
||||
assert result["clips"][2]["clip_type"] == "outro"
|
||||
|
||||
def test_parse_clips_renumbered_continuously(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 10, "duration": 2, "asset_id": "a1"},
|
||||
{"clip_type": "outro", "order": 20, "duration": 2, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["order"] == 0
|
||||
assert result["clips"][1]["order"] == 1
|
||||
|
||||
def test_parse_skips_invalid_clip_dicts(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"},
|
||||
"not a dict",
|
||||
{"clip_type": "outro", "order": 2, "duration": 2, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 2
|
||||
|
||||
|
||||
class TestParseRecommendResponseFields:
|
||||
"""各字段解析与边界测试."""
|
||||
|
||||
def test_parse_duration_clamped_min(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 0.5, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["duration"] == 1.0
|
||||
|
||||
def test_parse_duration_clamped_max(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 100, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["duration"] == 30.0
|
||||
|
||||
def test_parse_start_time_clamped_min(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1", "start_time": -5.0},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["start_time"] == 0.0
|
||||
|
||||
def test_parse_asset_id_not_in_list_empty(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "unknown_asset"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1", "a2"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["asset_id"] == ""
|
||||
|
||||
def test_parse_asset_id_in_list_kept(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a2"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1", "a2"], 30.0)
|
||||
assert result is not None
|
||||
assert result["clips"][0]["asset_id"] == "a2"
|
||||
|
||||
def test_parse_default_values(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"order": 0},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
clip = result["clips"][0]
|
||||
assert clip["clip_type"] == "showcase"
|
||||
assert clip["text_content"] == ""
|
||||
assert clip["duration"] == 3.0
|
||||
assert clip["transition_effect"] == "cut"
|
||||
assert clip["asset_id"] == ""
|
||||
assert clip["start_time"] == 0.0
|
||||
assert clip["config"] == {}
|
||||
|
||||
|
||||
class TestParseRecommendResponseMarkdown:
|
||||
"""Markdown 代码块包裹的 JSON 测试."""
|
||||
|
||||
def test_parse_markdown_json(self):
|
||||
content = (
|
||||
"```json\n"
|
||||
+ json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
"title": "md test",
|
||||
}
|
||||
)
|
||||
+ "\n```"
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 1
|
||||
assert result["config"]["title"]["text"] == "md test"
|
||||
|
||||
def test_parse_backticks_no_language(self):
|
||||
content = (
|
||||
"```\n"
|
||||
+ json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
}
|
||||
)
|
||||
+ "\n```"
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert len(result["clips"]) == 1
|
||||
|
||||
|
||||
class TestParseRecommendResponseConfidence:
|
||||
"""confidence 解析测试."""
|
||||
|
||||
def test_parse_confidence_normal(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
"confidence": 0.85,
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["confidence"] == 0.85
|
||||
|
||||
def test_parse_confidence_clamped_min(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
"confidence": -0.5,
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["confidence"] == 0.0
|
||||
|
||||
def test_parse_confidence_clamped_max(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
"confidence": 1.5,
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["confidence"] == 1.0
|
||||
|
||||
def test_parse_confidence_default(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["confidence"] == 0.7
|
||||
|
||||
|
||||
class TestParseRecommendResponseConfig:
|
||||
"""config 生成测试."""
|
||||
|
||||
def test_parse_no_title_no_ai_auto(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
# 没有 title 时,config 的 title.text 保持默认(DEFAULT_EDIT_PLAN_CONFIG 中的值)
|
||||
assert "title" in result["config"]
|
||||
|
||||
def test_parse_config_is_deep_copy(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [{"clip_type": "intro", "order": 0, "duration": 2, "asset_id": "a1"}],
|
||||
"title": "test",
|
||||
}
|
||||
)
|
||||
result1 = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
result2 = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
# 修改其中一个不影响另一个
|
||||
result1["config"]["title"]["text"] = "modified"
|
||||
assert result2["config"]["title"]["text"] != "modified"
|
||||
|
||||
|
||||
class TestParseRecommendResponseTotalDuration:
|
||||
"""total_duration 计算测试."""
|
||||
|
||||
def test_parse_total_duration_sum(self):
|
||||
content = json.dumps(
|
||||
{
|
||||
"clips": [
|
||||
{"clip_type": "intro", "order": 0, "duration": 3.5, "asset_id": "a1"},
|
||||
{"clip_type": "showcase", "order": 1, "duration": 5.2, "asset_id": "a1"},
|
||||
{"clip_type": "outro", "order": 2, "duration": 2.0, "asset_id": "a1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
result = _parse_recommend_response(content, ["a1"], 30.0)
|
||||
assert result is not None
|
||||
assert result["total_duration"] == pytest.approx(10.7, abs=0.01)
|
||||
|
||||
|
||||
# ── _fallback_recommend_clips 测试 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFallbackRecommendClips:
|
||||
"""本地降级推荐方案测试."""
|
||||
|
||||
def test_fallback_returns_dict_with_clips(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2"], "one_take", 30.0)
|
||||
assert "clips" in result
|
||||
assert "config" in result
|
||||
assert "total_duration" in result
|
||||
assert "confidence" in result
|
||||
|
||||
def test_fallback_has_intro_and_outro(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2"], "one_take", 30.0)
|
||||
clips = result["clips"]
|
||||
assert clips[0]["clip_type"] == "intro"
|
||||
assert clips[-1]["clip_type"] == "outro"
|
||||
|
||||
def test_fallback_showcase_count_matches_assets(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2", "a3"], "one_take", 30.0)
|
||||
showcase_clips = [c for c in result["clips"] if c["clip_type"] == "showcase"]
|
||||
assert len(showcase_clips) == 3
|
||||
|
||||
def test_fallback_no_assets_still_works(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", [], "one_take", 30.0)
|
||||
assert len(result["clips"]) >= 2 # 至少有intro和outro
|
||||
|
||||
def test_fallback_intro_uses_first_asset(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2"], "one_take", 30.0)
|
||||
assert result["clips"][0]["asset_id"] == "a1"
|
||||
|
||||
def test_fallback_outro_has_empty_asset(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1"], "one_take", 30.0)
|
||||
assert result["clips"][-1]["asset_id"] == ""
|
||||
|
||||
def test_fallback_confidence_in_range(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1"], "one_take", 30.0)
|
||||
assert 0.75 <= result["confidence"] <= 0.95
|
||||
|
||||
def test_fallback_title_contains_asset_count(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2", "a3"], "one_take", 30.0)
|
||||
assert "3" in result["config"]["title"]["text"]
|
||||
assert result["config"]["title"]["ai_auto"] is True
|
||||
|
||||
def test_fallback_total_duration_matches(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2"], "one_take", 30.0)
|
||||
total = sum(c["duration"] for c in result["clips"])
|
||||
assert result["total_duration"] == round(total, 1)
|
||||
|
||||
def test_fallback_orders_are_sequential(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _fallback_recommend_clips("plan1", "tmpl1", ["a1", "a2", "a3"], "one_take", 30.0)
|
||||
orders = [c["order"] for c in result["clips"]]
|
||||
assert orders == list(range(len(result["clips"])))
|
||||
|
||||
|
||||
# ── _call_ai_cover_service 测试 ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAiCoverService:
|
||||
"""AI封面生成服务测试."""
|
||||
|
||||
def test_cover_type_upload(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _call_ai_cover_service("plan1", ["a1"], "upload")
|
||||
assert result["type"] == "upload"
|
||||
assert result["image_url"] == ""
|
||||
|
||||
def test_cover_type_manual_with_frame_time(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _call_ai_cover_service("plan1", ["a1"], "manual", frame_time=5.5)
|
||||
assert result["type"] == "manual"
|
||||
assert result["frame_time"] == 5.5
|
||||
assert "5.5" in result["image_url"]
|
||||
|
||||
def test_cover_type_ai_frame(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
with patch("shared.ai_service.random.uniform", side_effect=[5.0, 0.9]):
|
||||
result = _call_ai_cover_service("plan1", ["a1"], "ai_frame")
|
||||
assert result["type"] == "ai_frame"
|
||||
assert result["frame_time"] == 5.0
|
||||
assert result["confidence"] == 0.9
|
||||
assert "plan1" in result["image_url"]
|
||||
|
||||
def test_cover_type_ai_regenerate(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _call_ai_cover_service("plan1", ["a1"], "ai_regenerate")
|
||||
assert result["type"] == "ai_frame"
|
||||
|
||||
def test_cover_frame_time_in_range(self):
|
||||
with patch("shared.ai_service.time.sleep"):
|
||||
result = _call_ai_cover_service("plan1", ["a1"], "ai_frame")
|
||||
assert 1.0 <= result["frame_time"] <= 10.0
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,412 +1,100 @@
|
||||
"""文本分段工具单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
"""text_splitter 单元测试."""
|
||||
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
|
||||
|
||||
class TestSplitText:
|
||||
"""split_text 函数测试"""
|
||||
|
||||
def test_empty_string_returns_empty_list(self):
|
||||
"""空字符串返回空列表"""
|
||||
def test_empty_text_returns_empty(self):
|
||||
assert split_text("") == []
|
||||
|
||||
def test_whitespace_only_returns_empty_list(self):
|
||||
"""纯空白字符返回空列表"""
|
||||
assert split_text(" \n \t ") == []
|
||||
|
||||
def test_short_text_returns_single_segment(self):
|
||||
"""短文本直接返回单段"""
|
||||
text = "这是一段短文本。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert result == [text]
|
||||
|
||||
def test_text_length_equals_max_chars(self):
|
||||
"""文本长度恰好等于 max_chars 时返回单段"""
|
||||
text = "a" * 100
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 100
|
||||
|
||||
def test_splits_on_sentence_boundary(self):
|
||||
"""在句子边界处分段"""
|
||||
# 构造长文本,确保超过 max_chars
|
||||
sentences = ["今天天气真好。我们一起去公园散步吧。", "公园里有很多花。还有很多小朋友在玩耍。"] * 10
|
||||
text = "".join(sentences)
|
||||
|
||||
result = split_text(text, max_chars=200)
|
||||
|
||||
assert len(result) >= 2
|
||||
# 每段都不超过 max_chars
|
||||
for seg in result:
|
||||
assert len(seg) <= 200
|
||||
|
||||
def test_all_segments_within_max_chars(self):
|
||||
"""所有分段都不超过 max_chars"""
|
||||
text = "这是第一句话。这是第二句话。这是第三句话。这是第四句话。这是第五句话。" * 10
|
||||
|
||||
result = split_text(text, max_chars=100)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_long_single_sentence_hard_cut(self):
|
||||
"""超长单句会被硬切"""
|
||||
text = "a" * 1000 # 没有标点
|
||||
|
||||
result = split_text(text, max_chars=200)
|
||||
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 200
|
||||
|
||||
def test_newline_is_sentence_end(self):
|
||||
"""换行符作为句子结束符"""
|
||||
text = "第一行内容\n第二行内容\n第三行内容" * 10
|
||||
|
||||
result = split_text(text, max_chars=50)
|
||||
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 50
|
||||
|
||||
def test_chinese_punctuation(self):
|
||||
"""中文标点(。!?;)作为句子结束符"""
|
||||
text = "你好!今天吃什么?我吃米饭;你呢?我也吃米饭。" * 10
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
def test_english_punctuation(self):
|
||||
"""英文标点(.!?;)作为句子结束符"""
|
||||
text = "Hello! How are you? I'm fine; thank you. Good bye." * 10
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
def test_merged_short_segments(self):
|
||||
"""过短的段落会被合并"""
|
||||
# 构造很多短句
|
||||
text = "你好。再见。谢谢。抱歉。好的。不行。可以。去吧。" * 5 # 每句3-4字
|
||||
|
||||
result = split_text(text, max_chars=100)
|
||||
|
||||
# 合并后段数应该比单纯按句切的少
|
||||
assert len(result) < len(text) // 3 # 粗略估计
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_preserves_content(self):
|
||||
"""分段后内容总和与原文基本一致(忽略strip的空白)"""
|
||||
text = "这是测试文本。包含多个句子。用来验证分段正确性。" * 5
|
||||
|
||||
result = split_text(text, max_chars=50)
|
||||
|
||||
# 合并所有分段,去掉空白后应该与原文去掉空白后基本一致
|
||||
combined = "".join(result).replace(" ", "")
|
||||
original = text.strip().replace(" ", "")
|
||||
assert combined == original
|
||||
|
||||
def test_custom_max_chars(self):
|
||||
"""支持自定义 max_chars"""
|
||||
text = "测试" * 100 # 200字
|
||||
|
||||
result_50 = split_text(text, max_chars=50)
|
||||
result_100 = split_text(text, max_chars=100)
|
||||
|
||||
# max_chars 越小,段数应该越多
|
||||
assert len(result_50) >= len(result_100)
|
||||
|
||||
def test_single_char_text(self):
|
||||
"""单字符文本"""
|
||||
assert split_text("好", max_chars=10) == ["好"]
|
||||
|
||||
def test_text_with_only_punctuation(self):
|
||||
"""纯标点文本"""
|
||||
text = "。。。。。。。。。。" # 10个句号
|
||||
result = split_text(text, max_chars=5)
|
||||
|
||||
assert len(result) >= 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 5
|
||||
|
||||
def test_mixed_content(self):
|
||||
"""中英文混合内容"""
|
||||
text = "今天的天气是 sunny and warm。我们去了 park 玩。真的很开心!" * 5
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
|
||||
# ── 短文本与空文本补充 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextEmptyAndShort:
|
||||
"""空文本与短文本补充场景."""
|
||||
|
||||
def test_whitespace_only_returns_empty(self):
|
||||
"""纯空白文本返回空列表."""
|
||||
def test_whitespace_only(self):
|
||||
assert split_text(" \n\t ") == []
|
||||
|
||||
def test_single_char(self):
|
||||
"""单字符文本."""
|
||||
assert split_text("好", max_chars=10) == ["好"]
|
||||
|
||||
def test_exactly_max_chars_no_split(self):
|
||||
"""刚好等于 max_chars 不分割."""
|
||||
text = "a" * 100
|
||||
result = split_text(text, max_chars=100)
|
||||
def test_short_text_single_segment(self):
|
||||
text = "你好世界。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert result[0] == text
|
||||
|
||||
def test_one_over_max_chars_splits(self):
|
||||
"""超过 max_chars 1 个字符就会分割."""
|
||||
text = "a" * 101
|
||||
def test_exact_max_chars(self):
|
||||
text = "a" * 500
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 500
|
||||
|
||||
def test_splits_on_sentence_boundary(self):
|
||||
# 两个长句子,各300字左右,超过50字阈值
|
||||
sent1 = "你" * 300 + "。"
|
||||
sent2 = "我" * 300 + "。"
|
||||
text = sent1 + sent2
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 2
|
||||
assert result[0] == sent1
|
||||
assert result[1] == sent2
|
||||
|
||||
def test_long_sentence_hard_cut(self):
|
||||
# 一个超长句子,没有句末标点,会被硬切
|
||||
text = "长" * 800
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) >= 2
|
||||
assert all(len(seg) <= 500 for seg in result)
|
||||
# 合起来应该等于原文本
|
||||
assert "".join(result) == text
|
||||
|
||||
def test_short_segments_merged(self):
|
||||
# 多个短句应该被合并
|
||||
sentences = [f"第{i}句。" for i in range(10)]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=200)
|
||||
# 每句5字左右,10句才50字,应该合并成1段
|
||||
assert len(result) < 10
|
||||
assert len(result[0]) <= 200
|
||||
|
||||
def test_preserves_content(self):
|
||||
text = "今天天气真好。我们去公园玩吧!你觉得怎么样?好的,走吧。"
|
||||
result = split_text(text, max_chars=20)
|
||||
# 合并后内容应一致
|
||||
assert "".join(result) == text
|
||||
|
||||
def test_multiple_punctuation_types(self):
|
||||
# 构造足够长的文本触发分段
|
||||
text = "第一" * 30 + "。" + "第二" * 30 + "!" + "第三" * 30 + "?" + "第四" * 30 + ";"
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text
|
||||
|
||||
def test_none_raises(self):
|
||||
"""None 输入抛 AttributeError(strip 失败)."""
|
||||
with pytest.raises(AttributeError):
|
||||
split_text(None)
|
||||
def test_custom_max_chars(self):
|
||||
text = "a" * 100 + "。" + "b" * 100 + "。"
|
||||
result = split_text(text, max_chars=150)
|
||||
assert len(result) == 2
|
||||
assert "a" in result[0]
|
||||
assert "b" in result[1]
|
||||
|
||||
|
||||
# ── 句子边界分段补充 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextSentenceBoundaries:
|
||||
"""句子边界分段补充场景."""
|
||||
|
||||
def test_split_on_fullwidth_period(self):
|
||||
"""全角句号分段."""
|
||||
text = "第一句很长的内容。" * 20
|
||||
result = split_text(text, max_chars=60)
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 60
|
||||
|
||||
def test_split_on_fullwidth_question(self):
|
||||
"""全角问号分段."""
|
||||
text = "你知道这是为什么吗?" + "是的。" * 20
|
||||
result = split_text(text, max_chars=60)
|
||||
assert len(result) > 1
|
||||
|
||||
def test_split_on_fullwidth_exclamation(self):
|
||||
"""全角感叹号分段."""
|
||||
text = "真是太棒了!" + "内容。" * 20
|
||||
result = split_text(text, max_chars=60)
|
||||
assert len(result) > 1
|
||||
|
||||
def test_split_on_newline(self):
|
||||
"""换行符分段."""
|
||||
lines = ["这是第一行很长的一段文字内容" * 3 for _ in range(5)]
|
||||
text = "\n".join(lines)
|
||||
result = split_text(text, max_chars=80)
|
||||
assert len(result) > 1
|
||||
|
||||
def test_split_on_semicolon(self):
|
||||
"""全角分号分段."""
|
||||
text = "第一项内容;" + "其他内容。" * 20
|
||||
result = split_text(text, max_chars=60)
|
||||
assert len(result) > 1
|
||||
|
||||
def test_english_period_splits(self):
|
||||
"""英文句号分段."""
|
||||
text = "Hello world. " * 30
|
||||
result = split_text(text, max_chars=80)
|
||||
assert len(result) > 1
|
||||
|
||||
def test_short_sentences_stay_merged(self):
|
||||
"""短句(都 < 50字的句子不会单独成段,会累积到一起."""
|
||||
text = "你好。我好。大家好。"
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
# ── 长句强制切段补充 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextLongSentenceForce:
|
||||
"""超长单句强制切段补充."""
|
||||
|
||||
def test_no_punctuation_forced_split(self):
|
||||
"""完全没有标点的超长文本硬切."""
|
||||
text = "字" * 300
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) == 3
|
||||
for seg in result:
|
||||
assert len(seg) == 100
|
||||
|
||||
def test_force_split_preserves_content(self):
|
||||
"""硬切不丢字符."""
|
||||
text = "a" * 250
|
||||
result = split_text(text, max_chars=100)
|
||||
assert sum(len(s) for s in result) == 250
|
||||
|
||||
def test_mixed_long_and_short(self):
|
||||
"""长句短句混合."""
|
||||
long_part = "非常长的句子没有标点符号" * 15
|
||||
text = long_part + "。结尾。"
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
|
||||
# ── 短段合并补充 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextShortSegmentMerge:
|
||||
"""短段合并补充场景."""
|
||||
|
||||
def test_multiple_short_sentences_merged(self):
|
||||
"""多个短句合并成一段."""
|
||||
sentences = ["你好。", "我好。", "大家好。", "天气好。", "心情好。"]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_short_tail_merged(self):
|
||||
"""尾部短段被合并到前一段."""
|
||||
# 前面一段接近 max_chars,尾部很短
|
||||
long_part = "一二三四五六七八九十" * 9 + "。" # ~90字
|
||||
tail = "完。" # 2字
|
||||
text = long_part + tail
|
||||
result = split_text(text, max_chars=100)
|
||||
# 尾部短的应该被合并
|
||||
assert len(result) <= 2
|
||||
|
||||
|
||||
# ── 边界情况补充 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextEdgeCases:
|
||||
"""边界情况补充."""
|
||||
|
||||
def test_only_punctuation(self):
|
||||
"""纯标点符号."""
|
||||
text = "。。。。。"
|
||||
result = split_text(text, max_chars=10)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_mixed_chinese_english(self):
|
||||
"""中英文混合."""
|
||||
text = "你好Hello。World!" * 20
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_strip_whitespace(self):
|
||||
"""首尾空白被去除."""
|
||||
text = " 你好世界。 "
|
||||
result = split_text(text, max_chars=100)
|
||||
assert result == ["你好世界。"]
|
||||
|
||||
def test_total_length_preserved(self):
|
||||
"""分段后总长度等于原文 strip 后长度."""
|
||||
text = "这是一段用于测试的文本内容。" * 20
|
||||
result = split_text(text, max_chars=100)
|
||||
def test_newline_as_sentence_end(self):
|
||||
text = "第一段\n第二段\n第三段"
|
||||
result = split_text(text, max_chars=50)
|
||||
assert len(result) >= 1
|
||||
assert "".join(result) == text.strip()
|
||||
|
||||
def test_custom_small_max_chars(self):
|
||||
"""很小的 max_chars."""
|
||||
text = "一二三四五六七八九十。" * 5
|
||||
def test_minimum_segment_length(self):
|
||||
# 句子太短(<50字)不会立即分段
|
||||
text = "短句一。短句二。短句三。"
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_trailing_content_added(self):
|
||||
# 最后一段不完整的句子也要加上
|
||||
text = "完整的句子。剩余内容"
|
||||
result = split_text(text, max_chars=50)
|
||||
assert "".join(result) == text
|
||||
|
||||
def test_no_empty_segments(self):
|
||||
text = "。。。。。" # 全是标点
|
||||
result = split_text(text, max_chars=2)
|
||||
assert all(len(seg) > 0 for seg in result)
|
||||
|
||||
def test_chinese_and_english_mixed(self):
|
||||
text = "Hello世界。这是测试Test文本。Mixed混合。"
|
||||
result = split_text(text, max_chars=20)
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 20
|
||||
|
||||
|
||||
# ── 更多边界场景补充 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitTextMoreEdgeCases:
|
||||
"""更多边界场景补充"""
|
||||
|
||||
def test_max_chars_one(self):
|
||||
"""max_chars=1 每个字符一段"""
|
||||
text = "一二三四五"
|
||||
result = split_text(text, max_chars=1)
|
||||
assert len(result) == 5
|
||||
for seg in result:
|
||||
assert len(seg) == 1
|
||||
|
||||
def test_consecutive_newlines(self):
|
||||
"""连续多个换行符"""
|
||||
text = "第一段\n\n\n第二段\n\n第三段"
|
||||
result = split_text(text, max_chars=100)
|
||||
# 合并后应该是一段(内容不长且合并逻辑会被合并)
|
||||
assert len(result) >= 1
|
||||
assert "第一段" in result[0]
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_only_newlines_only(self):
|
||||
"""只有换行符(纯空白被strip掉返回空"""
|
||||
assert split_text("\n\n\n\n") == []
|
||||
|
||||
def test_leading_trailing_whitespace(self):
|
||||
"""首尾空白被去除"""
|
||||
text = " 你好世界。 "
|
||||
result = split_text(text, max_chars=100)
|
||||
assert result == ["你好世界。"]
|
||||
|
||||
def test_very_long_single_sentence_many_segments(self):
|
||||
"""超长单句被切成很多段"""
|
||||
text = "字" * 1000
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) == 10
|
||||
for seg in result:
|
||||
assert len(seg) == 100
|
||||
|
||||
def test_mixed_punctuation_types(self):
|
||||
"""全角半角标点混合"""
|
||||
text = "你好!再见。谢谢?抱歉;好的"
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_last_segment_short_merged_to_previous(self):
|
||||
"""尾部极短段被合并到前一段"""
|
||||
# 构造第一段接近max_chars,结尾有个短句尾巴
|
||||
long_part = "一二三四五六七八九十" * 9 + "。" # ~90字
|
||||
tail = "完" # 1字
|
||||
text = long_part + tail
|
||||
result = split_text(text, max_chars=100)
|
||||
# 尾巴应该被合并
|
||||
combined = "".join(result)
|
||||
assert combined == text.strip()
|
||||
assert len(result) <= 2
|
||||
|
||||
def test_all_short_sentences_merged_into_one(self):
|
||||
"""大量短句全部合并成一段"""
|
||||
sentences = ["你好。", "我好。", "他好。", "大家好。", "才是真的好。"]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_punctuation_only_long(self):
|
||||
"""很长的纯标点文本"""
|
||||
text = "。" * 200
|
||||
result = split_text(text, max_chars=50)
|
||||
assert len(result) >= 4
|
||||
for seg in result:
|
||||
assert len(seg) <= 50
|
||||
|
||||
def test_tab_not_sentence_end(self):
|
||||
"""制表符不是句子结束符"""
|
||||
text = "这是一段\t包含制表符的文本内容" + "字" * 100
|
||||
result = split_text(text, max_chars=50)
|
||||
# 制表符不在句子结束符集合中,不会触发分段
|
||||
# 制表符会保留在分段内容中
|
||||
has_tab = any("\t" in seg for seg in result)
|
||||
assert has_tab
|
||||
assert len(result) >= 2
|
||||
assert "".join(result) == text
|
||||
|
||||
Reference in New Issue
Block a user