259 lines
10 KiB
Python
Executable File
259 lines
10 KiB
Python
Executable File
"""noise_reduction_config 领域模型单测."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
|
|
from packages.domain.noise_reduction_config import (
|
|
DEFAULT_LEVEL,
|
|
DEFAULT_NOISE_FLOOR,
|
|
MAX_NOISE_FLOOR,
|
|
MIN_NOISE_FLOOR,
|
|
NoiseReductionConfig,
|
|
NoiseReductionLevel,
|
|
apply_noise_reduction_if_needed,
|
|
build_afftdn_filter,
|
|
build_arnndn_filter,
|
|
get_level_names,
|
|
)
|
|
|
|
# ── NoiseReductionLevel 枚举测试 ───────────────────────────────────────────
|
|
|
|
|
|
class TestNoiseReductionLevel:
|
|
def test_four_levels(self):
|
|
assert len(NoiseReductionLevel) == 4
|
|
|
|
def test_level_values(self):
|
|
assert NoiseReductionLevel.LOW.value == "low"
|
|
assert NoiseReductionLevel.MEDIUM.value == "medium"
|
|
assert NoiseReductionLevel.HIGH.value == "high"
|
|
assert NoiseReductionLevel.CUSTOM.value == "custom"
|
|
|
|
def test_from_string(self):
|
|
assert NoiseReductionLevel("low") == NoiseReductionLevel.LOW
|
|
assert NoiseReductionLevel("medium") == NoiseReductionLevel.MEDIUM
|
|
assert NoiseReductionLevel("high") == NoiseReductionLevel.HIGH
|
|
assert NoiseReductionLevel("custom") == NoiseReductionLevel.CUSTOM
|
|
|
|
|
|
# ── NoiseReductionConfig.from_dict 测试 ───────────────────────────────────
|
|
|
|
|
|
class TestNoiseReductionConfigFromDict:
|
|
def test_none_returns_disabled(self):
|
|
cfg = NoiseReductionConfig.from_dict(None)
|
|
assert cfg.enabled is False
|
|
|
|
def test_empty_dict_returns_disabled(self):
|
|
cfg = NoiseReductionConfig.from_dict({})
|
|
assert cfg.enabled is False
|
|
|
|
def test_disabled_returns_disabled(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": False})
|
|
assert cfg.enabled is False
|
|
|
|
def test_enabled_default_params(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True})
|
|
assert cfg.enabled is True
|
|
assert cfg.level == NoiseReductionLevel.MEDIUM
|
|
assert cfg.noise_floor == DEFAULT_NOISE_FLOOR
|
|
assert cfg.voice_enhance is False
|
|
|
|
def test_custom_level(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": -30.0})
|
|
assert cfg.level == NoiseReductionLevel.CUSTOM
|
|
assert cfg.noise_floor == -30.0
|
|
|
|
def test_invalid_level_defaults_medium(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "invalid"})
|
|
assert cfg.level == NoiseReductionLevel.MEDIUM
|
|
|
|
def test_noise_floor_clamped_low(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": -100.0})
|
|
assert cfg.noise_floor == MIN_NOISE_FLOOR
|
|
|
|
def test_noise_floor_clamped_high(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": 0.0})
|
|
assert cfg.noise_floor == MAX_NOISE_FLOOR
|
|
|
|
def test_invalid_noise_floor_type_uses_default(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": "not_a_number"})
|
|
assert cfg.noise_floor == DEFAULT_NOISE_FLOOR
|
|
|
|
def test_voice_enhance_true(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "voice_enhance": True})
|
|
assert cfg.voice_enhance is True
|
|
|
|
def test_case_insensitive_level(self):
|
|
cfg = NoiseReductionConfig.from_dict({"enabled": True, "level": "HIGH"})
|
|
assert cfg.level == NoiseReductionLevel.HIGH
|
|
|
|
|
|
# ── has_effect / get_effective_noise_floor 测试 ───────────────────────────
|
|
|
|
|
|
class TestConfigProperties:
|
|
def test_disabled_no_effect(self):
|
|
cfg = NoiseReductionConfig(enabled=False)
|
|
assert cfg.has_effect() is False
|
|
|
|
def test_enabled_has_effect(self):
|
|
cfg = NoiseReductionConfig(enabled=True)
|
|
assert cfg.has_effect() is True
|
|
|
|
def test_effective_noise_floor_low(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.LOW)
|
|
assert cfg.get_effective_noise_floor() == -35.0
|
|
|
|
def test_effective_noise_floor_medium(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM)
|
|
assert cfg.get_effective_noise_floor() == -25.0
|
|
|
|
def test_effective_noise_floor_high(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.HIGH)
|
|
assert cfg.get_effective_noise_floor() == -15.0
|
|
|
|
def test_effective_noise_floor_custom(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.CUSTOM, noise_floor=-40.0)
|
|
assert cfg.get_effective_noise_floor() == -40.0
|
|
|
|
def test_get_level_params_medium(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM)
|
|
params = cfg.get_level_params()
|
|
assert params["nf"] == -25.0
|
|
assert params["tn"] == -10.0
|
|
assert params["tr"] == 50.0
|
|
|
|
def test_get_level_params_custom(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.CUSTOM, noise_floor=-30.0)
|
|
params = cfg.get_level_params()
|
|
assert params["nf"] == -30.0
|
|
assert "tn" in params
|
|
assert "tr" in params
|
|
|
|
|
|
# ── validate 测试 ─────────────────────────────────────────────────────────
|
|
|
|
|
|
class TestValidate:
|
|
def test_disabled_valid(self):
|
|
cfg = NoiseReductionConfig(enabled=False)
|
|
ok, msg = cfg.validate()
|
|
assert ok is True
|
|
assert msg == ""
|
|
|
|
def test_enabled_valid(self):
|
|
cfg = NoiseReductionConfig(enabled=True, noise_floor=-25.0)
|
|
ok, msg = cfg.validate()
|
|
assert ok is True
|
|
|
|
def test_noise_floor_out_of_range(self):
|
|
cfg = NoiseReductionConfig(enabled=True, noise_floor=-100.0)
|
|
ok, msg = cfg.validate()
|
|
assert ok is False
|
|
assert "noise_floor" in msg
|
|
|
|
|
|
# ── build_afftdn_filter 测试 ───────────────────────────────────────────────
|
|
|
|
|
|
class TestBuildAfftdnFilter:
|
|
def test_disabled_returns_anull(self):
|
|
cfg = NoiseReductionConfig(enabled=False)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "anull" in result
|
|
assert "[in]" in result
|
|
assert "[out]" in result
|
|
|
|
def test_medium_level(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM)
|
|
result = build_afftdn_filter(cfg, "[a]", "[nr]")
|
|
assert "afftdn=" in result
|
|
assert "nf=-25.0" in result or "nf=-25" in result
|
|
assert "[a]" in result
|
|
assert "[nr]" in result
|
|
|
|
def test_high_level(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.HIGH)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "afftdn=" in result
|
|
assert "nf=-15.0" in result or "nf=-15" in result
|
|
|
|
def test_low_level(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.LOW)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "afftdn=" in result
|
|
assert "nf=-35.0" in result or "nf=-35" in result
|
|
|
|
def test_custom_level(self):
|
|
cfg = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.CUSTOM, noise_floor=-40.0)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "afftdn=" in result
|
|
assert "nf=-40.0" in result or "nf=-40" in result
|
|
|
|
def test_voice_enhance_adds_filters(self):
|
|
cfg = NoiseReductionConfig(enabled=True, voice_enhance=True)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "highpass" in result
|
|
assert "acompressor" in result
|
|
assert "loudnorm" in result
|
|
|
|
def test_no_voice_enhance_no_extra_filters(self):
|
|
cfg = NoiseReductionConfig(enabled=True, voice_enhance=False)
|
|
result = build_afftdn_filter(cfg, "[in]", "[out]")
|
|
assert "highpass" not in result
|
|
assert "acompressor" not in result
|
|
|
|
|
|
# ── build_arnndn_filter 测试 ───────────────────────────────────────────────
|
|
|
|
|
|
class TestBuildArnndnFilter:
|
|
def test_disabled_returns_anull(self):
|
|
cfg = NoiseReductionConfig(enabled=False)
|
|
result = build_arnndn_filter(cfg, "[in]", "[out]", "model.rnnn")
|
|
assert "anull" in result
|
|
|
|
def test_enabled_returns_arnndn(self):
|
|
cfg = NoiseReductionConfig(enabled=True)
|
|
result = build_arnndn_filter(cfg, "[a]", "[nr]", "/path/to/model.rnnn")
|
|
assert "arnndn=" in result
|
|
assert "m=/path/to/model.rnnn" in result
|
|
assert "[a]" in result
|
|
assert "[nr]" in result
|
|
|
|
|
|
# ── apply_noise_reduction_if_needed 测试 ─────────────────────────────────
|
|
|
|
|
|
class TestApplyNoiseReductionIfNeeded:
|
|
def test_none_config_returns_none(self):
|
|
assert apply_noise_reduction_if_needed(None, "[in]", "[out]") is None
|
|
|
|
def test_disabled_returns_none(self):
|
|
assert apply_noise_reduction_if_needed({"enabled": False}, "[in]", "[out]") is None
|
|
|
|
def test_enabled_returns_filter(self):
|
|
result = apply_noise_reduction_if_needed({"enabled": True, "level": "medium"}, "[in]", "[out]")
|
|
assert result is not None
|
|
assert "afftdn" in result
|
|
|
|
def test_invalid_config_handles_exception(self):
|
|
# 异常情况应该返回 None 而不是抛出
|
|
result = apply_noise_reduction_if_needed("invalid", "[in]", "[out]")
|
|
assert result is None
|
|
|
|
|
|
# ── 工具函数测试 ───────────────────────────────────────────────────────────
|
|
|
|
|
|
class TestUtils:
|
|
def test_get_level_names_returns_four(self):
|
|
names = get_level_names()
|
|
assert len(names) == 4
|
|
assert "low" in names
|
|
assert "medium" in names
|
|
assert "high" in names
|
|
assert "custom" in names
|