CoolFace
Apppublic

NhanNguyen1309/audio-separation-api

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
test_key_api.py147 linesDownload Raw Back to tests
1from datetime import datetime2from pathlib import Path3from types import SimpleNamespace4from unittest import mock5import unittest6 7from app.chords import run_chord_detection_job8from app.main import build_chord_job_status_response9from app.models import (10    ChordJobRecord,11    ChordSegmentRecord,12    KeyCandidateRecord,13    KeyDetectionResultRecord,14    KeyScoreComponentsRecord,15    MusicAnalysisRecord,16    SnapConfigRecord,17    TimeSignatureRecord,18    VisualGridRecord,19)20 21 22def key_result() -> KeyDetectionResultRecord:23    components = KeyScoreComponentsRecord(0.8, 1.0, 0.9, 1.0, 0.8, 0.5, 1.0)24    candidate = KeyCandidateRecord("D", "D", "major", 0.88, components)25    return KeyDetectionResultRecord("D", "D", "major", 0.42, "ensemble", [candidate])26 27 28def music_analysis() -> MusicAnalysisRecord:29    return MusicAnalysisRecord(30        bpm=120.0,31        bpm_confidence=0.8,32        raw_bpm_candidates=[],33        beats=[0.0, 1.0, 2.0],34        downbeats=[0.0],35        bar_start_offset_sec=0.0,36        time_signature=TimeSignatureRecord(4, 4),37        beats_per_bar=4,38        visual_grid=VisualGridRecord(4, []),39        chords=[],40        chord_source="chordino",41        beat_source="fallback",42        warnings=[],43        original_key=key_result(),44        chord_audio_source="harmonic",45        snap_config=SnapConfigRecord(0.2, 0.3, 0.25, True),46        key_analysis_chords=[ChordSegmentRecord(0.0, 2.0, "D")],47        key_audio_source="original_mix",48        key_detection_source="original_mix_unsnapped",49    )50 51 52class KeyApiTests(unittest.TestCase):53    def test_chord_job_response_includes_original_key(self):54        now = datetime.now()55        job = ChordJobRecord(56            id="job-id",57            status="completed",58            progress=100,59            stage="completed",60            created_at=now,61            updated_at=now,62            input_path=Path("input.wav"),63            working_dir=Path("."),64            prepared_path=Path("prepared.wav"),65            result=[ChordSegmentRecord(0.0, 2.0, "D")],66            analysis=music_analysis(),67        )68 69        payload = build_chord_job_status_response(job).model_dump()70 71        self.assertEqual(payload["originalKey"]["key"], "D")72        self.assertEqual(payload["originalKey"]["method"], "ensemble")73        self.assertEqual(payload["originalKey"]["candidates"][0]["components"]["cadenceScore"], 0.8)74        self.assertEqual(payload["chordAudioSource"], "harmonic")75        self.assertEqual(payload["snapConfig"]["downbeatSnapToleranceSec"], 0.3)76        self.assertEqual(payload["snapConfig"]["forwardDownbeatToleranceSec"], 0.35)77        self.assertEqual(payload["snapConfig"]["backwardDownbeatToleranceSec"], 0.25)78        self.assertEqual(payload["keyAnalysisChords"][0]["chord"], "D")79        self.assertEqual(payload["keyAudioSource"], "original_mix")80        self.assertEqual(payload["keyDetectionSource"], "original_mix_unsnapped")81 82    def test_chord_job_runs_ensemble_after_postprocessing(self):83        harmonic_raw = [ChordSegmentRecord(0.0, 1.0, "A"), ChordSegmentRecord(1.0, 2.0, "E")]84        original_raw = [ChordSegmentRecord(0.0, 1.0, "D"), ChordSegmentRecord(1.0, 2.0, "A")]85        final = [ChordSegmentRecord(0.0, 2.0, "D")]86        analysis = music_analysis()87        analysis.original_key = None88        job = SimpleNamespace(input_path=Path("input.wav"), prepared_path=Path("prepared.wav"), working_dir=Path("."))89 90        with (91            mock.patch("app.chords.chord_registry.get_job", return_value=job),92            mock.patch("app.chords.chord_registry.update_job"),93            mock.patch("app.chords.ensure_chord_duration", return_value=2.0),94            mock.patch("app.chords.prepare_audio_for_chordino"),95            mock.patch(96                "app.chords.prepare_chord_detection_audio",97                return_value=SimpleNamespace(path=Path("harmonic.wav"), source="harmonic", warnings=[]),98            ),99            mock.patch("app.chords.run_chordino", side_effect=["harmonic csv", "original csv"]) as run_chordino,100            mock.patch("app.chords.parse_chordino_csv", side_effect=[harmonic_raw, original_raw]),101            mock.patch("app.chords.analyze_music_grid", return_value=analysis) as analyze_grid,102            mock.patch("app.chords.ChordPostprocessConfig.from_settings", return_value=SimpleNamespace(103                beats_per_bar=4,104                snap_tolerance_sec=0.2,105                downbeat_snap_tolerance_sec=0.3,106                forward_downbeat_tolerance_sec=0.35,107                backward_downbeat_tolerance_sec=0.25,108                visual_cell_snap_tolerance_sec=0.25,109                enable_downbeat_snap=True,110            )),111            mock.patch("app.chords.postprocess_chords_with_stats", return_value=(final, SimpleNamespace(112                raw_count=2,113                beat_count=2,114                snapped_count=1,115                min_duration_count=1,116                viterbi_count=1,117                guardrail_count=1,118                final_count=1,119                short_chords_removed=1,120                duplicate_merges=0,121                bar_density_merges=0,122            ))) as postprocess,123            mock.patch("app.chords.detect_original_key_ensemble", return_value=key_result()) as detect_key,124            mock.patch("app.chords.settings.key_use_harmonic_chords", False),125            mock.patch("app.chords.settings.key_use_snapped_chords", False),126        ):127            run_chord_detection_job("job-id")128 129        self.assertEqual(run_chordino.call_args_list, [mock.call(Path("harmonic.wav")), mock.call(Path("prepared.wav"))])130        analyze_grid.assert_called_once_with(Path("prepared.wav"), 2.0, harmonic_raw)131        self.assertEqual(132            detect_key.call_args_list,133            [134                mock.call("prepared.wav", original_raw, analysis.beats, analysis.downbeats, 2.0),135                mock.call("prepared.wav", harmonic_raw, analysis.beats, analysis.downbeats, 2.0),136            ],137        )138        self.assertEqual(postprocess.call_args.kwargs["downbeats"], analysis.downbeats)139        self.assertEqual(analysis.original_key.key, "D")140        self.assertEqual(analysis.key_analysis_chords, original_raw)141        self.assertEqual(analysis.key_audio_source, "original_mix")142        self.assertEqual(analysis.key_detection_source, "original_mix_unsnapped")143 144 145if __name__ == "__main__":146    unittest.main()147