Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 6f0a5d63c7 |
@@ -6,6 +6,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
# ── 单轨时间计算 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
from packages.domain.speed_config import (
|
||||
DEFAULT_SPEED,
|
||||
SpeedConfig,
|
||||
|
||||
+537
@@ -0,0 +1,537 @@
|
||||
"""plan_generator_utils 单元测试 - wave166
|
||||
|
||||
覆盖:
|
||||
- distribute_assets 素材分配(4种模式 + 边界)
|
||||
- map_clip_types_for_mode clip类型映射(4种模式)
|
||||
- generate_default_clips 默认片段生成(4种模式 + 边界)
|
||||
- create_clips_from_configs 从模板配置创建
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.edit_plan_clip import EditPlanClip
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
from packages.domain.plan_generator_utils import (
|
||||
DEFAULT_CLIP_DURATION,
|
||||
create_clips_from_configs,
|
||||
distribute_assets,
|
||||
generate_default_clips,
|
||||
map_clip_types_for_mode,
|
||||
)
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
|
||||
|
||||
# ============================================================
|
||||
# 辅助函数
|
||||
# ============================================================
|
||||
|
||||
|
||||
def _make_main_clip(plan_id: str, order: int = 0) -> EditPlanClip:
|
||||
return EditPlanClip.create(
|
||||
plan_id=plan_id,
|
||||
clip_type=ClipType.MAIN.value,
|
||||
order=order,
|
||||
duration=5.0,
|
||||
)
|
||||
|
||||
|
||||
def _make_clips(plan_id: str, count: int, clip_type: str = "main") -> list[EditPlanClip]:
|
||||
return [
|
||||
EditPlanClip.create(
|
||||
plan_id=plan_id,
|
||||
clip_type=clip_type,
|
||||
order=i,
|
||||
duration=5.0,
|
||||
)
|
||||
for i in range(count)
|
||||
]
|
||||
|
||||
|
||||
# ============================================================
|
||||
# distribute_assets - ONE_TAKE
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestDistributeOneTake:
|
||||
def test_equal_count(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
assets = ["a1", "a2", "a3"]
|
||||
distribute_assets(clips, assets, EditingMode.ONE_TAKE.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
assert clips[2].asset_id == "a3"
|
||||
|
||||
def test_more_clips_than_assets(self):
|
||||
clips = _make_clips("p1", 5)
|
||||
assets = ["a1", "a2"]
|
||||
distribute_assets(clips, assets, EditingMode.ONE_TAKE.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
assert clips[2].asset_id == "" # 没分配到
|
||||
|
||||
def test_more_assets_than_clips(self):
|
||||
clips = _make_clips("p1", 2)
|
||||
assets = ["a1", "a2", "a3"]
|
||||
distribute_assets(clips, assets, EditingMode.ONE_TAKE.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
|
||||
def test_empty_assets(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
distribute_assets(clips, [], EditingMode.ONE_TAKE.value)
|
||||
for c in clips:
|
||||
assert c.asset_id == ""
|
||||
|
||||
def test_empty_clips(self):
|
||||
# 不报错即可
|
||||
distribute_assets([], ["a1", "a2"], EditingMode.ONE_TAKE.value)
|
||||
|
||||
def test_only_main_clips_get_assigned(self):
|
||||
# intro/outro 不应该被分配
|
||||
clips = []
|
||||
clips.append(EditPlanClip.create("p1", clip_type="intro", order=0, duration=3.0))
|
||||
clips.append(_make_main_clip("p1", order=1))
|
||||
clips.append(EditPlanClip.create("p1", clip_type="outro", order=2, duration=3.0))
|
||||
assets = ["a1"]
|
||||
distribute_assets(clips, assets, EditingMode.ONE_TAKE.value)
|
||||
assert clips[0].asset_id == "" # intro 无
|
||||
assert clips[1].asset_id == "a1" # main 有
|
||||
assert clips[2].asset_id == "" # outro 无
|
||||
|
||||
|
||||
# ============================================================
|
||||
# distribute_assets - PIP
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestDistributePip:
|
||||
def test_first_asset_to_main(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
# 第一个main是背景,其余改为overlay
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assets = ["a1", "a2", "a3"]
|
||||
distribute_assets(clips, assets, EditingMode.PIP.value)
|
||||
assert clips[0].asset_id == "a1" # main → 背景
|
||||
assert clips[1].asset_id == "a2" # overlay
|
||||
assert clips[2].asset_id == "a3" # overlay
|
||||
|
||||
def test_single_asset(self):
|
||||
clips = _make_clips("p1", 1)
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assets = ["a1"]
|
||||
distribute_assets(clips, assets, EditingMode.PIP.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
|
||||
def test_only_main_clip_with_no_overlays(self):
|
||||
clips = _make_clips("p1", 1)
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assets = ["a1", "a2", "a3"] # 多余素材
|
||||
distribute_assets(clips, assets, EditingMode.PIP.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# distribute_assets - VOICE_OVER
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestDistributeVoiceOver:
|
||||
def test_assets_to_main_clips(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
assets = ["a1", "a2", "a3"]
|
||||
distribute_assets(clips, assets, EditingMode.VOICE_OVER.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
assert clips[2].asset_id == "a3"
|
||||
|
||||
def test_more_clips_than_assets(self):
|
||||
clips = _make_clips("p1", 5)
|
||||
assets = ["a1", "a2"]
|
||||
distribute_assets(clips, assets, EditingMode.VOICE_OVER.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
assert clips[2].asset_id == ""
|
||||
|
||||
|
||||
# ============================================================
|
||||
# distribute_assets - VOICE_PIP
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestDistributeVoicePip:
|
||||
def test_three_assets_three_roles(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assets = ["a1", "a2", "a3"]
|
||||
distribute_assets(clips, assets, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].clip_type == "corner_voice"
|
||||
assert clips[1].asset_id == "a2"
|
||||
assert clips[2].clip_type == "b_roll"
|
||||
assert clips[2].asset_id == "a3"
|
||||
|
||||
def test_single_asset(self):
|
||||
clips = _make_clips("p1", 1)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assets = ["a1"]
|
||||
distribute_assets(clips, assets, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[0].asset_id == "a1"
|
||||
|
||||
def test_two_assets(self):
|
||||
clips = _make_clips("p1", 2)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assets = ["a1", "a2"]
|
||||
distribute_assets(clips, assets, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# distribute_assets - 边界情况
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestDistributeEdgeCases:
|
||||
def test_unknown_mode_falls_back_to_one_take(self):
|
||||
clips = _make_clips("p1", 2)
|
||||
assets = ["a1", "a2"]
|
||||
distribute_assets(clips, assets, "unknown_mode")
|
||||
assert clips[0].asset_id == "a1"
|
||||
assert clips[1].asset_id == "a2"
|
||||
|
||||
def test_none_clips_no_crash(self):
|
||||
# 空列表
|
||||
distribute_assets([], ["a1"], EditingMode.ONE_TAKE.value)
|
||||
|
||||
def test_none_assets_no_crash(self):
|
||||
clips = _make_clips("p1", 2)
|
||||
distribute_assets(clips, [], EditingMode.ONE_TAKE.value)
|
||||
for c in clips:
|
||||
assert c.asset_id == ""
|
||||
|
||||
|
||||
# ============================================================
|
||||
# map_clip_types_for_mode
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestMapClipTypesForMode:
|
||||
def test_one_take_unchanged(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
original_types = [c.clip_type for c in clips]
|
||||
map_clip_types_for_mode(clips, EditingMode.ONE_TAKE.value)
|
||||
assert [c.clip_type for c in clips] == original_types
|
||||
|
||||
def test_voice_over_unchanged(self):
|
||||
clips = _make_clips("p1", 3)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_OVER.value)
|
||||
for c in clips:
|
||||
assert c.clip_type == ClipType.MAIN.value
|
||||
|
||||
def test_pip_first_stays_main_rest_overlay(self):
|
||||
clips = _make_clips("p1", 4)
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assert clips[0].clip_type == ClipType.MAIN.value
|
||||
assert clips[1].clip_type == "overlay"
|
||||
assert clips[2].clip_type == "overlay"
|
||||
assert clips[3].clip_type == "overlay"
|
||||
|
||||
def test_pip_single_clip_stays_main(self):
|
||||
clips = _make_clips("p1", 1)
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assert clips[0].clip_type == ClipType.MAIN.value
|
||||
|
||||
def test_voice_pip_mapping(self):
|
||||
clips = _make_clips("p1", 5)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[1].clip_type == "corner_voice"
|
||||
assert clips[2].clip_type == "b_roll"
|
||||
assert clips[3].clip_type == "b_roll"
|
||||
assert clips[4].clip_type == "b_roll"
|
||||
|
||||
def test_voice_pip_two_clips(self):
|
||||
clips = _make_clips("p1", 2)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[1].clip_type == "corner_voice"
|
||||
|
||||
def test_voice_pip_single_clip(self):
|
||||
clips = _make_clips("p1", 1)
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "background"
|
||||
|
||||
def test_non_main_clips_unchanged(self):
|
||||
clips = [
|
||||
EditPlanClip.create("p1", clip_type="intro", order=0, duration=3.0),
|
||||
_make_main_clip("p1", order=1),
|
||||
EditPlanClip.create("p1", clip_type="outro", order=2, duration=3.0),
|
||||
]
|
||||
map_clip_types_for_mode(clips, EditingMode.VOICE_PIP.value)
|
||||
assert clips[0].clip_type == "intro"
|
||||
assert clips[1].clip_type == "background" # main 被改了
|
||||
assert clips[2].clip_type == "outro"
|
||||
|
||||
def test_empty_clips_no_error(self):
|
||||
map_clip_types_for_mode([], EditingMode.PIP.value)
|
||||
|
||||
def test_no_main_clips_no_error(self):
|
||||
clips = [
|
||||
EditPlanClip.create("p1", clip_type="intro", order=0, duration=3.0),
|
||||
]
|
||||
map_clip_types_for_mode(clips, EditingMode.PIP.value)
|
||||
assert clips[0].clip_type == "intro"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate_default_clips
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGenerateDefaultClipsOneTake:
|
||||
def test_basic(self):
|
||||
clips = generate_default_clips("p1", EditingMode.ONE_TAKE.value, 3)
|
||||
assert len(clips) == 3
|
||||
for c in clips:
|
||||
assert c.clip_type == ClipType.MAIN.value
|
||||
assert c.plan_id == "p1"
|
||||
|
||||
def test_order_sequential(self):
|
||||
clips = generate_default_clips("p1", EditingMode.ONE_TAKE.value, 5)
|
||||
for i, c in enumerate(clips):
|
||||
assert c.order == i
|
||||
|
||||
def test_zero_assets_at_least_one(self):
|
||||
clips = generate_default_clips("p1", EditingMode.ONE_TAKE.value, 0)
|
||||
assert len(clips) == 1
|
||||
|
||||
def test_negative_assets_at_least_one(self):
|
||||
clips = generate_default_clips("p1", EditingMode.ONE_TAKE.value, -5)
|
||||
assert len(clips) == 1
|
||||
|
||||
def test_default_duration(self):
|
||||
clips = generate_default_clips("p1", EditingMode.ONE_TAKE.value, 1)
|
||||
assert clips[0].duration == DEFAULT_CLIP_DURATION
|
||||
|
||||
|
||||
class TestGenerateDefaultClipsPip:
|
||||
def test_one_asset(self):
|
||||
clips = generate_default_clips("p1", EditingMode.PIP.value, 1)
|
||||
assert len(clips) == 1
|
||||
assert clips[0].clip_type == ClipType.MAIN.value
|
||||
|
||||
def test_three_assets(self):
|
||||
clips = generate_default_clips("p1", EditingMode.PIP.value, 3)
|
||||
assert len(clips) == 3
|
||||
assert clips[0].clip_type == ClipType.MAIN.value
|
||||
assert clips[1].clip_type == "overlay"
|
||||
assert clips[2].clip_type == "overlay"
|
||||
|
||||
def test_zero_assets(self):
|
||||
clips = generate_default_clips("p1", EditingMode.PIP.value, 0)
|
||||
assert len(clips) >= 1
|
||||
assert clips[0].clip_type == ClipType.MAIN.value
|
||||
|
||||
|
||||
class TestGenerateDefaultClipsVoiceOver:
|
||||
def test_basic(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_OVER.value, 3)
|
||||
assert len(clips) == 3
|
||||
for c in clips:
|
||||
assert c.clip_type == ClipType.MAIN.value
|
||||
|
||||
def test_has_b_roll_config(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_OVER.value, 2)
|
||||
# VOICE_OVER 标记 role=b_roll
|
||||
assert clips[0].config.get("role") == "b_roll"
|
||||
|
||||
|
||||
class TestGenerateDefaultClipsVoicePip:
|
||||
def test_one_asset(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_PIP.value, 1)
|
||||
assert len(clips) == 1
|
||||
assert clips[0].clip_type == "background"
|
||||
|
||||
def test_two_assets(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_PIP.value, 2)
|
||||
assert len(clips) == 2
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[1].clip_type == "corner_voice"
|
||||
|
||||
def test_five_assets(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_PIP.value, 5)
|
||||
assert len(clips) == 5
|
||||
assert clips[0].clip_type == "background"
|
||||
assert clips[1].clip_type == "corner_voice"
|
||||
assert clips[2].clip_type == "b_roll"
|
||||
assert clips[3].clip_type == "b_roll"
|
||||
assert clips[4].clip_type == "b_roll"
|
||||
|
||||
def test_zero_assets(self):
|
||||
clips = generate_default_clips("p1", EditingMode.VOICE_PIP.value, 0)
|
||||
assert len(clips) >= 1
|
||||
assert clips[0].clip_type == "background"
|
||||
|
||||
|
||||
class TestGenerateDefaultClipsUnknownMode:
|
||||
def test_falls_back_to_one_take(self):
|
||||
clips = generate_default_clips("p1", "unknown_mode", 3)
|
||||
assert len(clips) == 3
|
||||
for c in clips:
|
||||
assert c.clip_type == ClipType.MAIN.value
|
||||
|
||||
|
||||
# ============================================================
|
||||
# create_clips_from_configs
|
||||
# ============================================================
|
||||
|
||||
|
||||
def _make_template_config(
|
||||
cfg_id: str,
|
||||
order: int,
|
||||
clip_type: ClipType = ClipType.MAIN,
|
||||
min_dur: float = 0,
|
||||
max_dur: float = 0,
|
||||
) -> TemplateClipConfig:
|
||||
return TemplateClipConfig(
|
||||
id=cfg_id,
|
||||
template_id="t1",
|
||||
clip_type=clip_type,
|
||||
order=order,
|
||||
min_duration=min_dur,
|
||||
max_duration=max_dur,
|
||||
transition_effect="cut",
|
||||
config={},
|
||||
)
|
||||
|
||||
|
||||
class TestCreateClipsFromConfigs:
|
||||
def test_empty_configs(self):
|
||||
result = create_clips_from_configs("p1", [])
|
||||
assert result == []
|
||||
|
||||
def test_single_config(self):
|
||||
configs = [_make_template_config("c1", 0, min_dur=3.0, max_dur=7.0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert len(result) == 1
|
||||
assert result[0].plan_id == "p1"
|
||||
assert result[0].template_clip_config_id == "c1"
|
||||
# 平均时长 = (3+7)/2 = 5.0
|
||||
assert result[0].duration == pytest.approx(5.0)
|
||||
|
||||
def test_duration_min_only(self):
|
||||
configs = [_make_template_config("c1", 0, min_dur=4.0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].duration == 4.0
|
||||
|
||||
def test_duration_max_only(self):
|
||||
configs = [_make_template_config("c1", 0, max_dur=6.0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].duration == 6.0
|
||||
|
||||
def test_duration_default_when_no_bounds(self):
|
||||
configs = [_make_template_config("c1", 0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].duration == DEFAULT_CLIP_DURATION
|
||||
|
||||
def test_sorted_by_order(self):
|
||||
configs = [
|
||||
_make_template_config("c_third", 2),
|
||||
_make_template_config("c_first", 0),
|
||||
_make_template_config("c_second", 1),
|
||||
]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert len(result) == 3
|
||||
assert result[0].template_clip_config_id == "c_first"
|
||||
assert result[1].template_clip_config_id == "c_second"
|
||||
assert result[2].template_clip_config_id == "c_third"
|
||||
assert result[0].order == 0
|
||||
assert result[1].order == 1
|
||||
assert result[2].order == 2
|
||||
|
||||
def test_clip_type_preserved(self):
|
||||
configs = [
|
||||
TemplateClipConfig(
|
||||
id="c_intro",
|
||||
template_id="t1",
|
||||
clip_type=ClipType.INTRO,
|
||||
order=0,
|
||||
min_duration=3.0,
|
||||
max_duration=3.0,
|
||||
transition_effect="cut",
|
||||
config={},
|
||||
),
|
||||
TemplateClipConfig(
|
||||
id="c_main",
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=1,
|
||||
min_duration=5.0,
|
||||
max_duration=5.0,
|
||||
transition_effect="cut",
|
||||
config={},
|
||||
),
|
||||
]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].clip_type == ClipType.INTRO.value
|
||||
assert result[1].clip_type == ClipType.MAIN.value
|
||||
|
||||
def test_playback_speed_from_config(self):
|
||||
configs = [
|
||||
TemplateClipConfig(
|
||||
id="c1",
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=0,
|
||||
min_duration=5.0,
|
||||
max_duration=5.0,
|
||||
transition_effect="cut",
|
||||
config={"playback_speed": 1.5},
|
||||
)
|
||||
]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].playback_speed == pytest.approx(1.5)
|
||||
|
||||
def test_speed_ratio_fallback(self):
|
||||
# 兼容 speed_ratio 字段名
|
||||
configs = [
|
||||
TemplateClipConfig(
|
||||
id="c1",
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=0,
|
||||
min_duration=5.0,
|
||||
max_duration=5.0,
|
||||
transition_effect="cut",
|
||||
config={"speed_ratio": 0.8},
|
||||
)
|
||||
]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].playback_speed == pytest.approx(0.8)
|
||||
|
||||
def test_default_playback_speed(self):
|
||||
configs = [_make_template_config("c1", 0, min_dur=5.0, max_dur=5.0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].playback_speed == pytest.approx(1.0)
|
||||
|
||||
def test_transition_effect_preserved(self):
|
||||
configs = [
|
||||
TemplateClipConfig(
|
||||
id="c1",
|
||||
template_id="t1",
|
||||
clip_type=ClipType.MAIN,
|
||||
order=0,
|
||||
min_duration=5.0,
|
||||
max_duration=5.0,
|
||||
transition_effect="fade",
|
||||
config={},
|
||||
)
|
||||
]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert result[0].transition_effect == "fade"
|
||||
|
||||
def test_returns_edit_plan_clip_objects(self):
|
||||
configs = [_make_template_config("c1", 0, min_dur=3.0, max_dur=5.0)]
|
||||
result = create_clips_from_configs("p1", configs)
|
||||
assert isinstance(result[0], EditPlanClip)
|
||||
@@ -1,567 +0,0 @@
|
||||
"""url_security 单测.
|
||||
|
||||
domain 层 URL 安全校验纯逻辑模块,0 网络依赖。
|
||||
覆盖 SSRF 防护、主机名校验、IP 检查、魔数校验等。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.url_security import (
|
||||
ALLOWED_IMAGE_MIME_TYPES,
|
||||
ALLOWED_PORTS,
|
||||
ALLOWED_SCHEMES,
|
||||
ALLOWED_VIDEO_MIME_TYPES,
|
||||
MAGIC_NUMBERS,
|
||||
MAX_URL_LENGTH,
|
||||
UrlSecurityError,
|
||||
check_internal_hostname,
|
||||
check_ssrf_ip,
|
||||
is_ip_address,
|
||||
is_trusted_domain,
|
||||
is_url_basic_safe,
|
||||
validate_magic_number,
|
||||
validate_url_basic,
|
||||
)
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# 常量与异常类
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestConstants:
|
||||
"""常量测试."""
|
||||
|
||||
def test_allowed_schemes(self):
|
||||
"""允许的 scheme 包含 http 和 https."""
|
||||
assert "http" in ALLOWED_SCHEMES
|
||||
assert "https" in ALLOWED_SCHEMES
|
||||
|
||||
def test_allowed_ports(self):
|
||||
"""允许的端口:80, 443."""
|
||||
assert 80 in ALLOWED_PORTS
|
||||
assert 443 in ALLOWED_PORTS
|
||||
|
||||
def test_max_url_length(self):
|
||||
"""最大 URL 长度 2048."""
|
||||
assert MAX_URL_LENGTH == 2048
|
||||
|
||||
def test_magic_numbers_has_common_formats(self):
|
||||
"""魔数表包含常见格式."""
|
||||
assert "image/jpeg" in MAGIC_NUMBERS
|
||||
assert "image/png" in MAGIC_NUMBERS
|
||||
assert "image/gif" in MAGIC_NUMBERS
|
||||
assert "video/mp4" in MAGIC_NUMBERS
|
||||
assert "audio/mpeg" in MAGIC_NUMBERS
|
||||
|
||||
|
||||
class TestUrlSecurityError:
|
||||
"""异常类测试."""
|
||||
|
||||
def test_is_value_error(self):
|
||||
"""UrlSecurityError 继承 ValueError."""
|
||||
assert issubclass(UrlSecurityError, ValueError)
|
||||
|
||||
def test_raise_with_message(self):
|
||||
"""抛出时携带错误信息."""
|
||||
with pytest.raises(UrlSecurityError, match="test error"):
|
||||
raise UrlSecurityError("test error")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# check_internal_hostname
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestCheckInternalHostname:
|
||||
"""内部主机名检查测试."""
|
||||
|
||||
def test_normal_domain_passes(self):
|
||||
"""普通外部域名通过."""
|
||||
check_internal_hostname("example.com")
|
||||
check_internal_hostname("www.google.com")
|
||||
|
||||
def test_localhost_blocked(self):
|
||||
"""localhost 被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="内部主机名"):
|
||||
check_internal_hostname("localhost")
|
||||
|
||||
def test_localhost_case_insensitive(self):
|
||||
"""大小写不敏感."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("LOCALHOST")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("LocalHost")
|
||||
|
||||
def test_localhost_localdomain_blocked(self):
|
||||
"""localhost.localdomain 被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("localhost.localdomain")
|
||||
|
||||
def test_metadata_blocked(self):
|
||||
"""metadata 被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("metadata")
|
||||
|
||||
def test_metadata_google_internal_blocked(self):
|
||||
"""GCP 元数据服务被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("metadata.google.internal")
|
||||
|
||||
def test_cloud_metadata_ip_blocked(self):
|
||||
"""云元数据 IP 169.254.169.254 被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("169.254.169.254")
|
||||
|
||||
def test_local_suffix_blocked(self):
|
||||
""".local 后缀域名被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="内网域名"):
|
||||
check_internal_hostname("myhost.local")
|
||||
|
||||
def test_internal_suffix_blocked(self):
|
||||
""".internal 后缀被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("svc.cluster.internal")
|
||||
|
||||
def test_localdomain_suffix_blocked(self):
|
||||
""".localdomain 后缀被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_internal_hostname("host.localdomain")
|
||||
|
||||
def test_com_domain_not_blocked(self):
|
||||
""".com 域名不被拦截."""
|
||||
check_internal_hostname("example.com")
|
||||
check_internal_hostname("sub.example.com")
|
||||
|
||||
def test_subdomain_of_public_domain_ok(self):
|
||||
"""公网域名的子域名正常."""
|
||||
check_internal_hostname("api.example.com")
|
||||
check_internal_hostname("cdn.assets.example.org")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# is_trusted_domain
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestIsTrustedDomain:
|
||||
"""可信域名匹配测试."""
|
||||
|
||||
def test_empty_trusted_domains_allows_all(self):
|
||||
"""空集合允许所有域名."""
|
||||
assert is_trusted_domain("anything.com", set()) is True
|
||||
assert is_trusted_domain("anywhere.org", set()) is True
|
||||
|
||||
def test_exact_match(self):
|
||||
"""精确匹配."""
|
||||
trusted = {"example.com", "example.org"}
|
||||
assert is_trusted_domain("example.com", trusted) is True
|
||||
assert is_trusted_domain("example.org", trusted) is True
|
||||
|
||||
def test_subdomain_match(self):
|
||||
"""子域名匹配."""
|
||||
trusted = {"example.com"}
|
||||
assert is_trusted_domain("api.example.com", trusted) is True
|
||||
assert is_trusted_domain("cdn.assets.example.com", trusted) is True
|
||||
|
||||
def test_no_match(self):
|
||||
"""不匹配."""
|
||||
trusted = {"example.com"}
|
||||
assert is_trusted_domain("other.com", trusted) is False
|
||||
assert is_trusted_domain("example.net", trusted) is False
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""大小写不敏感."""
|
||||
trusted = {"Example.COM"}
|
||||
assert is_trusted_domain("example.com", trusted) is True
|
||||
assert is_trusted_domain("API.EXAMPLE.COM", trusted) is True
|
||||
|
||||
def test_partial_match_no(self):
|
||||
"""域名部分相同但不是子域名不匹配."""
|
||||
trusted = {"example.com"}
|
||||
# fakeexample.com 不是 example.com 的子域名
|
||||
assert is_trusted_domain("fakeexample.com", trusted) is False
|
||||
|
||||
def test_none_trusted_domains(self):
|
||||
"""trusted_domains 为 None 时由调用方处理,空 set 全允许."""
|
||||
# 传空集合时全允许
|
||||
assert is_trusted_domain("a.com", set()) is True
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# check_ssrf_ip
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestCheckSsrIp:
|
||||
"""IP SSRF 检查测试."""
|
||||
|
||||
def test_public_ip_passes(self):
|
||||
"""公网 IP 通过."""
|
||||
check_ssrf_ip("8.8.8.8")
|
||||
check_ssrf_ip("1.1.1.1")
|
||||
check_ssrf_ip("114.114.114.114")
|
||||
|
||||
def test_loopback_blocked(self):
|
||||
"""回环地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="回环"):
|
||||
check_ssrf_ip("127.0.0.1")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("127.0.0.53")
|
||||
|
||||
def test_private_ip_blocked(self):
|
||||
"""私有内网 IP 被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="内网"):
|
||||
check_ssrf_ip("192.168.1.1")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("10.0.0.1")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("172.16.0.1")
|
||||
|
||||
def test_link_local_blocked(self):
|
||||
"""链路本地地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="链路本地"):
|
||||
check_ssrf_ip("169.254.169.254")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("169.254.1.1")
|
||||
|
||||
def test_multicast_blocked(self):
|
||||
"""组播地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="组播"):
|
||||
check_ssrf_ip("224.0.0.1")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("239.255.255.250")
|
||||
|
||||
def test_unspecified_blocked(self):
|
||||
"""未指定地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError, match="未指定"):
|
||||
check_ssrf_ip("0.0.0.0")
|
||||
|
||||
def test_ipv6_loopback_blocked(self):
|
||||
"""IPv6 回环地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("::1")
|
||||
|
||||
def test_ipv6_private_blocked(self):
|
||||
"""IPv6 内网地址被拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("fc00::1")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
check_ssrf_ip("fe80::1")
|
||||
|
||||
def test_ipv6_public_passes(self):
|
||||
"""IPv6 公网地址通过."""
|
||||
check_ssrf_ip("2001:4860:4860::8888")
|
||||
|
||||
def test_invalid_ip_raises_value_error(self):
|
||||
"""非法 IP 抛出 ValueError(不是 UrlSecurityError)."""
|
||||
with pytest.raises(ValueError):
|
||||
check_ssrf_ip("not-an-ip")
|
||||
with pytest.raises(ValueError):
|
||||
check_ssrf_ip("999.999.999.999")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# is_ip_address
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestIsIpAddress:
|
||||
"""IP 地址判断测试."""
|
||||
|
||||
def test_ipv4_true(self):
|
||||
"""IPv4 地址返回 True."""
|
||||
assert is_ip_address("127.0.0.1") is True
|
||||
assert is_ip_address("8.8.8.8") is True
|
||||
assert is_ip_address("0.0.0.0") is True
|
||||
|
||||
def test_ipv6_true(self):
|
||||
"""IPv6 地址返回 True."""
|
||||
assert is_ip_address("::1") is True
|
||||
assert is_ip_address("2001:db8::1") is True
|
||||
|
||||
def test_hostname_false(self):
|
||||
"""主机名返回 False."""
|
||||
assert is_ip_address("example.com") is False
|
||||
assert is_ip_address("localhost") is False
|
||||
assert is_ip_address("sub.domain.org") is False
|
||||
|
||||
def test_empty_string_false(self):
|
||||
"""空字符串返回 False."""
|
||||
assert is_ip_address("") is False
|
||||
|
||||
def test_invalid_ip_false(self):
|
||||
"""非法 IP 返回 False."""
|
||||
assert is_ip_address("999.999.999.999") is False
|
||||
assert is_ip_address("1234") is False
|
||||
assert is_ip_address("abc.def") is False
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# validate_url_basic
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestValidateUrlBasic:
|
||||
"""URL 基础校验测试."""
|
||||
|
||||
def test_normal_https_url_passes(self):
|
||||
"""正常 HTTPS URL 通过."""
|
||||
result = validate_url_basic("https://example.com/path")
|
||||
assert result == "https://example.com/path"
|
||||
|
||||
def test_normal_http_url_passes(self):
|
||||
"""正常 HTTP URL 通过."""
|
||||
result = validate_url_basic("http://example.com/path")
|
||||
assert result == "http://example.com/path"
|
||||
|
||||
def test_empty_url_rejected(self):
|
||||
"""空 URL 被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="为空"):
|
||||
validate_url_basic("")
|
||||
|
||||
def test_none_url_not_passed_as_str(self):
|
||||
"""None 作为 URL(这里只测空字符串)."""
|
||||
# 空字符串被拒
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("")
|
||||
|
||||
def test_too_long_url_rejected(self):
|
||||
"""超长 URL 被拒."""
|
||||
long_url = "https://example.com/" + "a" * 3000
|
||||
with pytest.raises(UrlSecurityError, match="过长"):
|
||||
validate_url_basic(long_url)
|
||||
|
||||
def test_invalid_scheme_rejected(self):
|
||||
"""非法 scheme 被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="scheme"):
|
||||
validate_url_basic("ftp://example.com/file")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("file:///etc/passwd")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("javascript:alert(1)")
|
||||
|
||||
def test_missing_hostname_rejected(self):
|
||||
"""缺少主机名被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="主机名"):
|
||||
validate_url_basic("https:///path")
|
||||
|
||||
def test_localhost_rejected(self):
|
||||
"""localhost 被拒."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("https://localhost/api")
|
||||
|
||||
def test_internal_domain_rejected(self):
|
||||
"""内网域名被拒."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("http://server.local/api")
|
||||
|
||||
def test_non_standard_port_rejected(self):
|
||||
"""非标准端口被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="端口"):
|
||||
validate_url_basic("https://example.com:8080/")
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("http://example.com:3000/")
|
||||
|
||||
def test_port_80_ok(self):
|
||||
"""80 端口允许."""
|
||||
validate_url_basic("http://example.com:80/path")
|
||||
|
||||
def test_port_443_ok(self):
|
||||
"""443 端口允许."""
|
||||
validate_url_basic("https://example.com:443/path")
|
||||
|
||||
def test_no_port_ok(self):
|
||||
"""无端口默认允许."""
|
||||
validate_url_basic("https://example.com/path")
|
||||
|
||||
def test_direct_ip_rejected_by_default(self):
|
||||
"""默认禁止直接 IP 访问."""
|
||||
with pytest.raises(UrlSecurityError, match="直接 IP"):
|
||||
validate_url_basic("https://8.8.8.8/path")
|
||||
|
||||
def test_direct_ip_allowed_when_enabled(self):
|
||||
"""allow_direct_ip=True 时允许公网 IP."""
|
||||
validate_url_basic("https://8.8.8.8/path", allow_direct_ip=True)
|
||||
|
||||
def test_direct_ip_private_still_blocked(self):
|
||||
"""即使 allow_direct_ip,内网 IP 仍被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="内网"):
|
||||
validate_url_basic("https://192.168.1.1/", allow_direct_ip=True)
|
||||
|
||||
def test_direct_ip_loopback_still_blocked(self):
|
||||
"""回环 IP 即使开启 direct_ip 也被拒."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("https://127.0.0.1/", allow_direct_ip=True)
|
||||
|
||||
def test_trusted_domains_pass(self):
|
||||
"""可信域名列表内的域名通过."""
|
||||
trusted = {"example.com", "cdn.com"}
|
||||
validate_url_basic("https://api.example.com/path", trusted_domains=trusted)
|
||||
validate_url_basic("https://cdn.com/asset.jpg", trusted_domains=trusted)
|
||||
|
||||
def test_untrusted_domain_rejected(self):
|
||||
"""不在可信域名列表中的域名被拒."""
|
||||
trusted = {"example.com"}
|
||||
with pytest.raises(UrlSecurityError, match="白名单"):
|
||||
validate_url_basic("https://evil.com/malware", trusted_domains=trusted)
|
||||
|
||||
def test_trusted_domain_subdomain_pass(self):
|
||||
"""可信域名的子域名通过."""
|
||||
trusted = {"example.com"}
|
||||
validate_url_basic("https://sub.example.com/a", trusted_domains=trusted)
|
||||
validate_url_basic("https://a.b.example.com/b", trusted_domains=trusted)
|
||||
|
||||
def test_return_value_is_original_url(self):
|
||||
"""返回原始 URL 字符串."""
|
||||
url = "https://example.com/path?query=value#frag"
|
||||
assert validate_url_basic(url) == url
|
||||
|
||||
def test_metadata_ip_rejected(self):
|
||||
"""云元数据 IP 被内部主机名检查拦截."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_url_basic("http://169.254.169.254/latest/meta-data/")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# is_url_basic_safe
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestIsUrlBasicSafe:
|
||||
"""便捷函数 is_url_basic_safe 测试."""
|
||||
|
||||
def test_safe_url_returns_true(self):
|
||||
"""安全 URL 返回 True."""
|
||||
assert is_url_basic_safe("https://example.com/") is True
|
||||
assert is_url_basic_safe("http://example.org/path") is True
|
||||
|
||||
def test_unsafe_url_returns_false(self):
|
||||
"""不安全 URL 返回 False."""
|
||||
assert is_url_basic_safe("https://localhost/") is False
|
||||
assert is_url_basic_safe("ftp://example.com/") is False
|
||||
assert is_url_basic_safe("") is False
|
||||
|
||||
def test_trusted_domains_param(self):
|
||||
"""支持 trusted_domains 参数."""
|
||||
trusted = {"example.com"}
|
||||
assert is_url_basic_safe("https://other.com/", trusted_domains=trusted) is False
|
||||
assert is_url_basic_safe("https://example.com/", trusted_domains=trusted) is True
|
||||
|
||||
def test_allow_direct_ip_param(self):
|
||||
"""支持 allow_direct_ip 参数."""
|
||||
assert is_url_basic_safe("https://8.8.8.8/") is False
|
||||
assert is_url_basic_safe("https://8.8.8.8/", allow_direct_ip=True) is True
|
||||
|
||||
def test_no_exceptions_raised(self):
|
||||
"""不抛出异常,只返回 bool."""
|
||||
# 各种边界情况都不抛异常
|
||||
try:
|
||||
is_url_basic_safe("")
|
||||
is_url_basic_safe("not a url")
|
||||
is_url_basic_safe("http://" + "a" * 3000)
|
||||
except UrlSecurityError:
|
||||
pytest.fail("is_url_basic_safe should not raise UrlSecurityError")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# validate_magic_number
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestValidateMagicNumber:
|
||||
"""魔数校验测试."""
|
||||
|
||||
def test_jpeg_valid(self):
|
||||
"""JPEG 文件通过."""
|
||||
# JPEG 文件头: FF D8 FF
|
||||
jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00"
|
||||
validate_magic_number(jpeg_header, {"image/jpeg"})
|
||||
|
||||
def test_png_valid(self):
|
||||
"""PNG 文件通过."""
|
||||
png_header = b"\x89PNG\r\n\x1a\n\x00\x00\x00"
|
||||
validate_magic_number(png_header, {"image/png"})
|
||||
|
||||
def test_gif_valid(self):
|
||||
"""GIF 文件通过(GIF89a 和 GIF87a)."""
|
||||
validate_magic_number(b"GIF89a...", {"image/gif"})
|
||||
validate_magic_number(b"GIF87a...", {"image/gif"})
|
||||
|
||||
def test_wav_valid(self):
|
||||
"""WAV 文件通过(RIFF + WAVE)."""
|
||||
wav_header = b"RIFF\x00\x00\x00\x00WAVEfmt "
|
||||
validate_magic_number(wav_header, {"audio/wav"})
|
||||
|
||||
def test_mp3_id3_valid(self):
|
||||
"""带 ID3 标签的 MP3 通过."""
|
||||
mp3_header = b"ID3\x03\x00\x00\x00\x00\x0f\x76"
|
||||
validate_magic_number(mp3_header, {"audio/mpeg"})
|
||||
|
||||
def test_mp3_sync_valid(self):
|
||||
"""不带 ID3 的 MP3(帧同步字)通过."""
|
||||
mp3_header = b"\xff\xfb\x90\x00" + b"\x00" * 32
|
||||
validate_magic_number(mp3_header, {"audio/mpeg"})
|
||||
|
||||
def test_ogg_valid(self):
|
||||
"""OGG 文件通过."""
|
||||
validate_magic_number(b"OggS\x00\x00...", {"audio/ogg"})
|
||||
|
||||
def test_flac_valid(self):
|
||||
"""FLAC 文件通过."""
|
||||
validate_magic_number(b"fLaC\x00\x00...", {"audio/flac"})
|
||||
|
||||
def test_webp_valid(self):
|
||||
"""WebP 文件通过(RIFF + WEBP)."""
|
||||
webp_header = b"RIFF\x00\x00\x00\x00WEBPVP8 "
|
||||
validate_magic_number(webp_header, {"image/webp"})
|
||||
|
||||
def test_bmp_valid(self):
|
||||
"""BMP 文件通过."""
|
||||
validate_magic_number(b"BM\x00\x00\x00\x00...", {"image/bmp"})
|
||||
|
||||
def test_mp4_valid(self):
|
||||
"""MP4 文件通过(ftyp 在偏移 4)."""
|
||||
mp4_header = b"\x00\x00\x00\x20ftypisom\x00\x00\x02\x00"
|
||||
validate_magic_number(mp4_header, {"video/mp4"})
|
||||
|
||||
def test_invalid_format_rejected(self):
|
||||
"""不匹配的格式被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="魔数"):
|
||||
validate_magic_number(b"hello world", {"image/jpeg"})
|
||||
|
||||
def test_empty_bytes_rejected(self):
|
||||
"""空字节被拒."""
|
||||
with pytest.raises(UrlSecurityError, match="为空"):
|
||||
validate_magic_number(b"", {"image/jpeg"})
|
||||
|
||||
def test_too_short_bytes_rejected(self):
|
||||
"""字节太短不匹配魔数时被拒."""
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_magic_number(b"\xff\xd8", {"image/jpeg"}) # 只2字节,不够JPEG魔数
|
||||
|
||||
def test_multiple_allowed_types(self):
|
||||
"""允许多种格式时任一匹配即通过."""
|
||||
jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00"
|
||||
validate_magic_number(jpeg_header, {"image/jpeg", "image/png", "image/gif"})
|
||||
|
||||
def test_wrong_type_rejected(self):
|
||||
"""用 PNG 魔数校验 JPEG 类型失败."""
|
||||
jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00"
|
||||
with pytest.raises(UrlSecurityError):
|
||||
validate_magic_number(jpeg_header, {"image/png"})
|
||||
|
||||
def test_unknown_mime_skipped(self):
|
||||
"""未知 MIME 类型(无对应魔数)不阻断."""
|
||||
# application/octet-stream 没有魔数定义,直接通过
|
||||
validate_magic_number(b"random bytes here", {"application/octet-stream"})
|
||||
|
||||
def test_allowed_image_mime_types_has_common(self):
|
||||
"""图片 MIME 白名单包含常见类型."""
|
||||
assert "image/jpeg" in ALLOWED_IMAGE_MIME_TYPES
|
||||
assert "image/png" in ALLOWED_IMAGE_MIME_TYPES
|
||||
|
||||
def test_allowed_video_mime_types_has_common(self):
|
||||
"""视频 MIME 白名单包含常见类型."""
|
||||
assert "video/mp4" in ALLOWED_VIDEO_MIME_TYPES
|
||||
@@ -6,6 +6,7 @@ domain 层纯逻辑模块,0 FFmpeg 依赖,快速轻量。
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.transition_presets import (
|
||||
|
||||
Reference in New Issue
Block a user