Files
xiaoxia-saas/tests/unit/test_noise_reduction_config.py

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