from __future__ import annotations import hashlib from pathlib import Path import numpy as np import pytest from openmusic_analysis.analyzers.clap import ( ClapGlobalAudioAnalyzer, ClapTemporalAudioAnalyzer, ) from openmusic_analysis.analyzers.lyrics import BGEM3LyricsAnalyzer, LyricsPreprocessor from openmusic_analysis.application import MusicAnalysisService from openmusic_analysis.audio.decoder import AudioMetadata, DecodedAudio from openmusic_analysis.errors import AudioDecodeError from openmusic_analysis.registry import ModelRegistry from openmusic_analysis.settings import GlobalAudioConfig, LyricsConfig, TemporalAudioConfig class FakeAudioEncoder: dimension = 4 loaded = True def __init__(self) -> None: self.calls: list[list[np.ndarray]] = [] async def ready(self) -> None: return None async def encode(self, windows: list[np.ndarray]) -> np.ndarray: self.calls.append(windows) rows = [] for window in windows: rows.append( [ float(np.mean(window)), float(np.std(window)), float(window[0]), float(window[-1]) + 1.0, ] ) return np.asarray(rows, dtype=np.float32) class FailingAudioEncoder(FakeAudioEncoder): async def encode(self, windows: list[np.ndarray]) -> np.ndarray: raise RuntimeError("internal model detail must not leak") class FakeTextEncoder: dimension = 6 loaded = True async def ready(self) -> None: return None def count_tokens(self, text: str) -> int: return len(text.split()) + 2 def split_tokens(self, text: str, max_tokens: int) -> list[str]: words = text.split() size = max(1, max_tokens - 2) return [" ".join(words[index : index + size]) for index in range(0, len(words), size)] async def encode(self, texts: list[str], batch_size: int) -> np.ndarray: rows = [] for text in texts: digest = hashlib.sha256(text.encode("utf-8")).digest() rows.append([float(value + 1) for value in digest[: self.dimension]]) return np.asarray(rows, dtype=np.float32) class FakeDecoder: canonical_sample_rate = 10 def __init__(self, waveform: np.ndarray | None = None) -> None: self.waveform = ( np.asarray(waveform, dtype=np.float32) if waveform is not None else np.linspace(-1.0, 1.0, 95, dtype=np.float32) ) self.calls = 0 def decode(self, source_path: str | Path) -> DecodedAudio: self.calls += 1 if Path(source_path).read_bytes().startswith(b"bad"): raise AudioDecodeError() return DecodedAudio( waveform=self.waveform, sample_rate=self.canonical_sample_rate, metadata=AudioMetadata( duration_ms=int(round(self.waveform.size * 1000 / self.canonical_sample_rate)), source_sample_rate=self.canonical_sample_rate, source_channels=1, source_format="fake", ), ) def make_service( *, waveform: np.ndarray | None = None, audio_encoder: FakeAudioEncoder | None = None, ) -> tuple[MusicAnalysisService, FakeDecoder, FakeAudioEncoder]: audio_encoder = audio_encoder or FakeAudioEncoder() decoder = FakeDecoder(waveform) global_config = GlobalAudioConfig( sample_rate=10, window_seconds=2.0, target_windows=4, minimum_audio_seconds=1.0, ) temporal_config = TemporalAudioConfig( sample_rate=10, window_seconds=2.0, hop_seconds=2.0, max_segments=4, minimum_audio_seconds=1.0, ) analyzers = [ ClapGlobalAudioAnalyzer(audio_encoder, global_config), ClapTemporalAudioAnalyzer(audio_encoder, temporal_config), BGEM3LyricsAnalyzer( FakeTextEncoder(), LyricsPreprocessor(), LyricsConfig(max_chunk_tokens=12, batch_size=4), ), ] service = MusicAnalysisService( registry=ModelRegistry(analyzers), decoder=decoder, device="cpu" ) return service, decoder, audio_encoder @pytest.fixture def service_bundle(): return make_service()