from pathlib import Path from src import duration, ui def _estimate(monkeypatch, tmp_path, *, files=2, package_models=None, safe=False): paths = [] for index in range(files): path = tmp_path / f"item-{index}.wav" path.write_bytes(b"x") paths.append(path) monkeypatch.setattr(duration, "probe_duration", lambda path: 15.0) return duration.calculate_duration_estimate( paths, [], package_models or ["UVR_MDXNET_KARA_2.onnx"], "None", "avg_wave", "OPUS", 0, 0, duration.DURATION_MODE_SEMI_AUTO, 180, 12, 1.25, safe, ) def test_two_item_onnx_batch_declares_cold_start_aware_duration(monkeypatch, tmp_path): estimate = _estimate(monkeypatch, tmp_path) assert estimate.seconds == 78 assert estimate.first_inference_warmup_seconds == duration.ONNX_FIRST_INFERENCE_WARMUP_SECONDS assert estimate.additional_batch_item_overhead_seconds == duration.ADDITIONAL_BATCH_ITEM_DISPATCH_SECONDS assert estimate.calibration_revision == duration.DURATION_CALIBRATION_REVISION assert estimate.onnx_model_count == 1 def test_safe_mode_applies_after_cold_start_and_batch_overhead(monkeypatch, tmp_path): base = _estimate(monkeypatch, tmp_path, safe=False) safe = _estimate(monkeypatch, tmp_path, safe=True) assert safe.base_unclamped_seconds == base.base_unclamped_seconds assert safe.unclamped_seconds == base.unclamped_seconds * duration.SAFE_MODE_MULTIPLIER assert safe.seconds == duration._clamp_seconds(base.unclamped_seconds * duration.SAFE_MODE_MULTIPLIER) def test_non_onnx_profile_uses_non_onnx_warmup(monkeypatch, tmp_path): estimate = _estimate( monkeypatch, tmp_path, files=1, package_models=["UVR-De-Reverb-aufr33-jarredou.pth"], ) assert estimate.first_inference_warmup_seconds == duration.NON_ONNX_FIRST_INFERENCE_WARMUP_SECONDS assert estimate.onnx_model_count == 0 assert estimate.non_onnx_model_count == 1 def test_ui_wrapper_passes_explicit_batch_failure_policy(monkeypatch): captured = {} def fake_prepare(*args, **kwargs): captured["kwargs"] = kwargs return "state", "markdown", "preflight.json", "prepare.log" monkeypatch.setattr(ui, "prepare_job", fake_prepare) values = list(range(34)) + [False] result = ui.prepare_separation_with_progress(*values, progress=None) assert result[0] == "state" assert captured["kwargs"]["batch_continue_on_item_error"] is False