NhanNguyen1309/audio-separation-api
0
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 