import math from types import SimpleNamespace import numpy as np import pytest import torch from production import split_leading_clause, split_text_for_tts from quality_runtime import ( BASE_GENERATION_POLICY, SAFE_DURATION_GENERATION_POLICY, WHISPER_MODEL_ID, WHISPER_REVISION, CandidateObservation, CascadeResult, ChunkCandidateArtifact, FinalOutputRejectedError, LazyWhisperASR, NoQualifiedCandidateError, TrajectoryGateResult, WhisperRuntime, active_voiced_duration_seconds, candidate_chunk_transition_score, candidate_limit_for_chunk_budget, cosine_similarity, generation_policy_for_candidate_offset, load_pinned_whisper_runtime, prepare_candidate_audio, qualify_trajectory_with_joined_output, require_verified_final_output, resolve_request_seed, run_adaptive_cascade, select_k_candidate_sequences, speaker_embedding_from_audio, speaker_evidence_from_audio, transcribe_whisper, verify_candidate, verify_trajectory, ) class _Factory: def __init__(self, value): self.value = value self.calls = [] def from_pretrained(self, *args, **kwargs): self.calls.append((args, kwargs)) return self.value def test_candidate_generation_policy_mapping_uses_safe_duration_after_offset_zero(): base = generation_policy_for_candidate_offset(0) assert base is BASE_GENERATION_POLICY assert (base.name, base.cjk_cps, base.ascii_cps, base.hard_stop_margin_steps) == ( "base", 5.2, 4.6, 1, ) for offset in (1, 2, 5, 9): safe = generation_policy_for_candidate_offset(offset) assert safe is SAFE_DURATION_GENERATION_POLICY assert ( safe.name, safe.cjk_cps, safe.ascii_cps, safe.hard_stop_margin_steps, ) == ("safe_duration", 4.6, 4.0, 1) @pytest.mark.parametrize("offset", [-1, 0.5, 1.0, True, None, "one"]) def test_candidate_generation_policy_mapping_rejects_invalid_offsets(offset): with pytest.raises(ValueError, match="non-negative integer"): generation_policy_for_candidate_offset(offset) def test_explicit_request_seed_is_forwarded_without_calling_random_factory(): calls = [] def factory(limit): calls.append(limit) return 456 assert resolve_request_seed(123, factory) == 123 assert calls == [] assert resolve_request_seed(None, factory) == 456 assert calls == [2**31] @pytest.mark.parametrize("seed", [-1, 2**31, 1.0, True, "123", None]) def test_request_seed_validation_fails_closed(seed): factory = lambda _limit: 2**31 if seed is None else 0 with pytest.raises(ValueError, match=r"\[0, 2147483648\)"): resolve_request_seed(seed, factory) def test_quality_gate_accepts_only_eval_canonicalized_pronoun_homophones(): result = verify_candidate( CandidateObservation( target_text="她提醒我", transcript_text="他提醒我", audio_duration_seconds=1.0, ) ) assert result.passed assert result.comparison.cer == 0.0 assert result.comparison.prefix_cer == 0.0 assert result.comparison.suffix_cer == 0.0 class _FakeWhisperModel: def __init__(self): self.device = None self.evaluated = False self.generate_calls = [] def to(self, device): self.device = torch.device(device) return self def eval(self): self.evaluated = True return self def generate(self, features, **kwargs): self.generate_calls.append((features, kwargs)) return torch.tensor([[1, 2, 3]]).repeat(features.shape[0], 1) class _FakeProcessor: def __init__(self): self.calls = [] def __call__(self, waveform, **kwargs): copied = ( [segment.copy() for segment in waveform] if isinstance(waveform, list) else waveform.copy() ) self.calls.append((copied, kwargs)) batch_size = len(waveform) if isinstance(waveform, list) else 1 return SimpleNamespace(input_features=torch.ones(batch_size, 80, 10)) def batch_decode(self, token_ids, **kwargs): assert all(row == [1, 2, 3] for row in token_ids.tolist()) assert kwargs == {"skip_special_tokens": True} return [" 合成內容完整。 "] * token_ids.shape[0] class _StatsEncoder: def __init__(self, *, invalid=False): self.lengths = [] self.wav_lens = [] self.invalid = invalid def encode_batch(self, tensor, wav_lens=None): self.lengths.append(tensor.shape[-1]) if wav_lens is None: wav_lens = torch.ones(tensor.shape[0], device=tensor.device) self.wav_lens.append(wav_lens.detach().cpu().numpy()) if self.invalid: return torch.full( (tensor.shape[0], 1, 2), math.nan, device=tensor.device, ) rows = [] for row, relative_length in zip(tensor, wav_lens, strict=True): sample_count = max(1, int(round(float(relative_length) * tensor.shape[-1]))) active = row[:sample_count] rows.append( torch.stack( ( active.mean(), active.std(unbiased=False), active.abs().amax(), ) ) ) return torch.stack(rows) def _tone(sample_rate=16_000, seconds=1.0, amplitude=0.2, frequency=220.0): timeline = np.arange(round(sample_rate * seconds), dtype=np.float32) / sample_rate return amplitude * np.sin(2.0 * np.pi * frequency * timeline).astype(np.float32) def _gate_result(passed, score=0.0): return TrajectoryGateResult( passed=passed, candidate_results=(), score=score if passed else math.inf, rejection_reasons=() if passed else ("rejected",), ) def _chunk_verification(target, transcript, *, embedding, rms_db): observation = CandidateObservation( target_text=target, transcript_text=transcript, audio_duration_seconds=2.0, speaker_similarity=0.8, begin_speaker_similarity=0.8, end_speaker_similarity=0.8, ) return observation, ChunkCandidateArtifact( speaker_embedding=( None if embedding is None else np.asarray(embedding, dtype=np.float32) ), rms_db=rms_db, ) def _whole_speaker_verification(*, similarity, boundary_drop, passed=True): target = "完整而且穩定的候選內容" observation = CandidateObservation( target_text=target, transcript_text=target if passed else "錯誤內容", audio_duration_seconds=2.5, speaker_similarity=similarity, begin_speaker_similarity=0.60, end_speaker_similarity=0.60 - boundary_drop, ) artifact = ChunkCandidateArtifact( speaker_embedding=np.array([1.0, 0.0], dtype=np.float32), rms_db=-20.0, ) return verify_trajectory( [observation], chunk_artifacts=[artifact], min_speaker_similarity=0.10, max_boundary_speaker_drop=0.10, ) def _joined_verification(target, transcript): return verify_trajectory( [ CandidateObservation( target_text=target, transcript_text=transcript, audio_duration_seconds=3.0, speaker_similarity=0.8, begin_speaker_similarity=0.8, end_speaker_similarity=0.8, ) ] ) def test_pinned_whisper_loader_uses_exact_revision_without_global_download(): processor = _FakeProcessor() model = _FakeWhisperModel() processor_factory = _Factory(processor) model_factory = _Factory(model) runtime = load_pinned_whisper_runtime( device="cpu", processor_factory=processor_factory, model_factory=model_factory, ) assert runtime.processor is processor assert runtime.model is model assert runtime.device == torch.device("cpu") assert runtime.dtype == torch.float32 assert model.evaluated assert processor_factory.calls == [ ((WHISPER_MODEL_ID,), {"revision": WHISPER_REVISION}) ] model_args, model_kwargs = model_factory.calls[0] assert model_args == (WHISPER_MODEL_ID,) assert model_kwargs["revision"] == WHISPER_REVISION assert model_kwargs["torch_dtype"] == torch.float32 assert model_kwargs["use_safetensors"] is True def test_lazy_whisper_runtime_loads_once_and_transcription_is_deterministic(): processor = _FakeProcessor() model = _FakeWhisperModel() runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) loads = [] lazy_asr = LazyWhisperASR(lambda: loads.append("load") or runtime) audio_8khz = _tone(sample_rate=8_000, seconds=0.5) first = transcribe_whisper(audio_8khz, 8_000, lazy_asr=lazy_asr) second = transcribe_whisper(audio_8khz, 8_000, lazy_asr=lazy_asr) assert first == second == "合成內容完整。" assert loads == ["load"] assert processor.calls[0][0].shape == (8_000,) assert processor.calls[0][1] == { "sampling_rate": 16_000, "return_tensors": "pt", } _, generation_kwargs = model.generate_calls[0] assert generation_kwargs == { "language": "zh", "task": "transcribe", "do_sample": False, "num_beams": 1, "max_new_tokens": 128, } def test_long_whisper_audio_is_batched_below_thirty_second_limit(): processor = _FakeProcessor() model = _FakeWhisperModel() runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) audio = np.concatenate( ( _tone(seconds=21.0), np.zeros(16_000, dtype=np.float32), _tone(seconds=21.0, frequency=240.0), np.zeros(16_000, dtype=np.float32), _tone(seconds=21.0, frequency=260.0), ) ) transcript = transcribe_whisper(audio, 16_000, runtime=runtime, max_new_tokens=440) segments = processor.calls[0][0] assert isinstance(segments, list) assert len(segments) == 3 assert all(segment.size <= 30 * 16_000 for segment in segments) assert transcript == "合成內容完整。 合成內容完整。 合成內容完整。" assert len(model.generate_calls) == 1 assert model.generate_calls[0][1]["max_new_tokens"] == 440 def test_transcription_rejects_nonfinite_audio_and_conflicting_injection(): processor = _FakeProcessor() model = _FakeWhisperModel() runtime = WhisperRuntime(processor, model, torch.device("cpu"), torch.float32) with pytest.raises(ValueError, match="non-finite"): transcribe_whisper(np.array([0.0, math.nan]), 16_000, runtime=runtime) with pytest.raises(ValueError, match="either runtime or lazy_asr"): transcribe_whisper( np.ones(20), 16_000, runtime=runtime, lazy_asr=LazyWhisperASR(lambda: runtime), ) def test_candidate_audio_data_or_asr_value_error_becomes_rejection_evidence(): calls = [] def transcriber(waveform, sample_rate): calls.append((waveform.copy(), sample_rate)) return " 內容完整 " prepared = prepare_candidate_audio(_tone(seconds=0.25), 16_000, transcriber=transcriber) assert prepared is not None assert prepared.transcript_text == "內容完整" assert prepared.duration_seconds == pytest.approx(0.25) assert len(calls) == 1 assert prepare_candidate_audio( [0.0, math.nan], 16_000, transcriber=transcriber, ) is None assert prepare_candidate_audio( _tone(seconds=0.25), 16_000, transcriber=lambda *_: (_ for _ in ()).throw(ValueError("bad candidate")), ) is None assert len(calls) == 1 def test_candidate_audio_infrastructure_error_propagates_fail_closed(): with pytest.raises(RuntimeError, match="ASR unavailable"): prepare_candidate_audio( _tone(seconds=0.25), 16_000, transcriber=lambda *_: (_ for _ in ()).throw(RuntimeError("ASR unavailable")), ) def test_ecapa_ndarray_helper_trims_resamples_normalizes_and_handles_channels(): encoder = _StatsEncoder() tone = _tone(sample_rate=8_000, seconds=1.0) padded = np.concatenate((np.zeros(4_000), tone, np.zeros(4_000))).astype(np.float32) stereo = np.stack((padded, padded), axis=0) embedding = speaker_embedding_from_audio(stereo, 8_000, encoder) assert embedding.shape == (3,) assert np.isfinite(embedding).all() assert np.linalg.norm(embedding) == pytest.approx(1.0, abs=1e-6) assert 16_000 <= encoder.lengths[0] < 32_000 assert len(encoder.lengths) == 1 def test_ecapa_ndarray_helper_fails_closed_on_silence_nan_and_bad_embedding(): with pytest.raises(ValueError, match="active speech"): speaker_embedding_from_audio(np.zeros(16_000), 16_000, _StatsEncoder()) with pytest.raises(ValueError, match="non-finite samples"): speaker_embedding_from_audio( np.array([0.0, math.nan], dtype=np.float32), 16_000, _StatsEncoder(), ) with pytest.raises(ValueError, match="invalid embedding"): speaker_embedding_from_audio(_tone(), 16_000, _StatsEncoder(invalid=True)) def test_cosine_and_speaker_evidence_are_finite_and_directional(): encoder = _StatsEncoder() audio = np.concatenate((_tone(seconds=1.0), _tone(seconds=1.0, amplitude=0.1))) anchor = speaker_embedding_from_audio(audio, 16_000, encoder) evidence = speaker_evidence_from_audio( audio, 16_000, encoder, anchor, edge_seconds=0.75, ) # Anchor extraction is one call; whole/begin/end evidence is one batched # call instead of three repeated ECAPA forwards for the same candidate. assert len(encoder.lengths) == 2 assert encoder.wav_lens[-1].shape == (3,) assert evidence.similarity == pytest.approx(1.0, abs=1e-5) assert -1.0 <= evidence.begin_similarity <= 1.0 assert -1.0 <= evidence.end_similarity <= 1.0 assert evidence.boundary_drop == max( 0.0, evidence.begin_similarity - evidence.end_similarity, ) assert evidence.active_duration_seconds > 1.5 assert cosine_similarity(anchor, anchor) == pytest.approx(1.0) with pytest.raises(ValueError): cosine_similarity([0.0, 0.0], [0.0, 0.0]) def test_long_speaker_evidence_uses_one_bounded_window_batch(): encoder = _StatsEncoder() audio = _tone(seconds=20.0) anchor = np.array([0.0, 0.0, 1.0], dtype=np.float32) evidence = speaker_evidence_from_audio(audio, 16_000, encoder, anchor) assert len(encoder.lengths) == 1 assert encoder.lengths[0] <= 3 * 16_000 assert encoder.wav_lens[0].shape == (6,) assert evidence.active_duration_seconds == pytest.approx(20.0, abs=0.05) assert evidence.speaker_embedding.shape == (3,) assert math.isfinite(evidence.active_rms_db) def test_active_union_pace_ignores_one_or_three_second_internal_silence(): def waveform(pause_seconds): return np.concatenate( ( _tone(seconds=0.75), np.zeros(round(pause_seconds * 16_000), dtype=np.float32), _tone(seconds=0.75, frequency=260.0), ) ) one_second = active_voiced_duration_seconds(waveform(1.0), 16_000) three_seconds = active_voiced_duration_seconds(waveform(3.0), 16_000) assert one_second == pytest.approx(three_seconds, abs=1.0e-9) assert 8 / one_second == pytest.approx(8 / three_seconds, abs=1.0e-9) assert one_second < 1.6 def test_active_union_duration_controls_speaker_eligibility_at_1_5_seconds(): below = active_voiced_duration_seconds(_tone(seconds=1.49), 16_000) boundary = active_voiced_duration_seconds(_tone(seconds=1.50), 16_000) assert below == pytest.approx(1.49, abs=1.0e-9) assert boundary == pytest.approx(1.50, abs=1.0e-9) common = { "target_text": "這是一段完整而且穩定的合成內容", "transcript_text": "這是一段完整而且穩定的合成內容", } short_result = verify_candidate( CandidateObservation(**common, audio_duration_seconds=below) ) boundary_result = verify_candidate( CandidateObservation(**common, audio_duration_seconds=boundary) ) assert short_result.passed assert not short_result.speaker_gate_applied assert not boundary_result.passed assert boundary_result.speaker_gate_applied assert "missing_speaker_evidence" in boundary_result.rejection_reasons def test_short_candidate_uses_exact_semantics_but_bypasses_speaker_hard_gate(): passed = verify_candidate( CandidateObservation( target_text="謝謝你", transcript_text="謝謝你", audio_duration_seconds=1.49, ) ) failed = verify_candidate( CandidateObservation( target_text="謝謝你", transcript_text="謝謝", audio_duration_seconds=1.0, ) ) assert passed.passed assert not passed.speaker_gate_applied assert passed.speaker_similarity is None assert passed.score == 0.0 assert not failed.passed assert failed.rejection_reasons == ("semantic_gate",) def test_candidate_pace_gate_is_optional_and_fails_closed_when_enabled(): base = CandidateObservation( target_text="內容完整", transcript_text="內容完整", audio_duration_seconds=1.0, ) assert verify_candidate(base).passed assert verify_candidate( CandidateObservation(**{**base.__dict__, "pace_cps": 4.2}), max_pace_cps=4.3, ).passed missing = verify_candidate(base, max_pace_cps=4.3) fast = verify_candidate( CandidateObservation(**{**base.__dict__, "pace_cps": 4.31}), max_pace_cps=4.3, ) assert missing.rejection_reasons == ("missing_pace_evidence",) assert fast.rejection_reasons == ("pace_too_fast",) def test_long_candidate_requires_speaker_evidence_and_clamps_negative_drop(): observation = CandidateObservation( target_text="這是一段完整而且穩定的合成內容", transcript_text="這是一段完整而且穩定的合成內容", audio_duration_seconds=2.0, speaker_similarity=0.70, begin_speaker_similarity=0.50, end_speaker_similarity=0.60, ) result = verify_candidate(observation) assert result.passed assert result.speaker_gate_applied assert result.boundary_speaker_drop == 0.0 assert result.score == pytest.approx(0.05 * 0.30) missing = verify_candidate( CandidateObservation( target_text=observation.target_text, transcript_text=observation.transcript_text, audio_duration_seconds=2.0, ) ) assert not missing.passed assert "missing_speaker_evidence" in missing.rejection_reasons def test_candidate_rejects_boundary_drop_low_similarity_truncation_and_tail(): base = dict( target_text="這是一段完整而且穩定的合成內容", transcript_text="這是一段完整而且穩定的合成內容", audio_duration_seconds=2.0, speaker_similarity=0.7, begin_speaker_similarity=0.5, end_speaker_similarity=0.5, ) boundary = verify_candidate( CandidateObservation(**{**base, "begin_speaker_similarity": 0.8}) ) similarity = verify_candidate( CandidateObservation(**{**base, "speaker_similarity": 0.09}) ) truncation = verify_candidate(CandidateObservation(**{**base, "truncated": True})) tail = verify_candidate( CandidateObservation(**{**base, "transcript_text": base["target_text"] + "多講"}) ) assert boundary.rejection_reasons == ("boundary_speaker_drop",) assert similarity.rejection_reasons == ("speaker_similarity",) assert truncation.rejection_reasons == ("truncated",) assert "semantic_gate" in tail.rejection_reasons # ASR comparison canonicalizes both sides to Simplified Chinese. assert tail.comparison.extra_tail == "多讲" def test_candidate_and_trajectory_fail_closed_for_bad_values_or_any_bad_chunk(): good_short = CandidateObservation("內容完整", "內容完整", 1.0) bad_short = CandidateObservation("句尾完整", "句尾", 1.0) invalid = verify_candidate( CandidateObservation("內容完整", "內容完整", math.nan), max_cer=math.nan, ) passed_trajectory = verify_trajectory([good_short, good_short]) failed_trajectory = verify_trajectory([good_short, bad_short]) assert not invalid.passed assert "invalid_audio_duration" in invalid.rejection_reasons assert "invalid_gate_config" in invalid.rejection_reasons assert passed_trajectory.passed assert passed_trajectory.score == 0.0 assert not failed_trajectory.passed assert math.isinf(failed_trajectory.score) assert failed_trajectory.rejection_reasons == ("chunk_1:semantic_gate",) assert verify_trajectory([]).rejection_reasons == ("empty_trajectory",) def test_final_whole_output_gate_rejects_post_join_regression_without_fallback(): passed = verify_trajectory( [CandidateObservation("整段內容完整", "整段內容完整", 1.0)] ) rejected = verify_trajectory( [CandidateObservation("整段內容完整", "整段內容", 1.0)] ) assert require_verified_final_output(passed) is passed with pytest.raises(FinalOutputRejectedError, match="final output rejected"): require_verified_final_output(rejected) with pytest.raises(FinalOutputRejectedError, match="invalid result"): require_verified_final_output(None) def test_joined_output_rejection_preserves_local_results_for_sequence_dp(): observations = [] artifacts = [] for target in ("第一段完整", "第二段完整"): observation, artifact = _chunk_verification( target, target, embedding=[1.0, 0.0], rms_db=-20.0, ) observations.append(observation) artifacts.append(artifact) local = verify_trajectory(observations, chunk_artifacts=artifacts) joined_failed = _joined_verification("第一段完整第二段完整", "第一段完整") qualified = qualify_trajectory_with_joined_output(local, joined_failed) assert local.passed assert not qualified.passed assert math.isinf(qualified.score) assert qualified.candidate_results is local.candidate_results assert qualified.chunk_artifacts is local.chunk_artifacts assert qualified.rejection_reasons == ("joined_output:chunk_0:semantic_gate",) assert all(result.passed for result in qualified.candidate_results) def test_joined_output_qualification_returns_local_result_only_when_joined_passes(): local = verify_trajectory( [CandidateObservation("第一段完整", "第一段完整", 1.0)] ) joined = verify_trajectory( [CandidateObservation("第一段完整", "第一段完整", 1.0)] ) assert qualify_trajectory_with_joined_output(local, joined) is local def test_k_best_sequence_paths_are_cost_ranked_stable_distinct_and_bounded(): ranked = select_k_candidate_sequences( [[0.0, 0.2], [0.0, 0.1]], [[[0.4, 0.0], [0.0, 0.5]]], max_paths=3, ) assert [selection.candidate_indices for selection in ranked] == [ (0, 1), (1, 0), (0, 0), ] assert [selection.total_score for selection in ranked] == pytest.approx( [0.1, 0.2, 0.4] ) assert len({selection.candidate_indices for selection in ranked}) == 3 tied = select_k_candidate_sequences( [[0.0, 0.0], [0.0, 0.0]], [[[0.0, 0.0], [0.0, 0.0]]], max_paths=3, ) assert [selection.candidate_indices for selection in tied] == [ (0, 0), (0, 1), (1, 0), ] def test_k_best_sequence_selection_fails_closed_for_malformed_or_single_chunk_graphs(): assert select_k_candidate_sequences([[0.0, 0.1]], (), max_paths=3) == () assert select_k_candidate_sequences([[0.0], [0.0]], (), max_paths=3) == () assert ( select_k_candidate_sequences( [[0.0, 0.1], [0.0, 0.1]], [[[0.0]]], max_paths=3, ) == () ) for invalid in (0, 4, 1.0, True): with pytest.raises(ValueError, match="between 1 and 3"): select_k_candidate_sequences( [[0.0], [0.0]], [[[0.0]]], max_paths=invalid, ) def _locally_safe_join_rejected_verification(chunks, seed): # Crossed embeddings make the two mixed paths cheaper than either # same-candidate path, giving deterministic k-best callback order. first_candidate = seed % 2 == 0 similarities = (0.9, 0.2) if first_candidate else (0.2, 0.9) embeddings = ( ([1.0, 0.0], [0.0, 1.0]) if first_candidate else ([0.0, 1.0], [1.0, 0.0]) ) observations = [] artifacts = [] for target, similarity, embedding in zip( chunks, similarities, embeddings, strict=True, ): observations.append( CandidateObservation( target_text=target, transcript_text=target, audio_duration_seconds=2.0, speaker_similarity=similarity, begin_speaker_similarity=similarity, end_speaker_similarity=similarity, ) ) artifacts.append( ChunkCandidateArtifact( speaker_embedding=np.asarray(embedding, dtype=np.float32), rms_db=-20.0, ) ) local = verify_trajectory(observations, chunk_artifacts=artifacts) joined_failed = _joined_verification("".join(chunks), chunks[0]) return qualify_trajectory_with_joined_output(local, joined_failed) def test_k_best_sequence_final_callback_uses_rank_order_and_first_passing_path(): chunks = ("第一段完整", "第二段完整") callback_paths = [] def generator(candidate_chunks, seed): return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) def verifier(trajectory, candidate_chunks, seed): return _locally_safe_join_rejected_verification(candidate_chunks, seed) def final_verifier(sequence_result, candidate_chunks): callback_paths.append(sequence_result.chunk_candidate_indices) transcript = ( "".join(candidate_chunks) if sequence_result.chunk_candidate_indices == (1, 0) else candidate_chunks[0] ) return _joined_verification("".join(candidate_chunks), transcript) result = run_adaptive_cascade( chunks, 20, generator, verifier, max_candidates=2, sequence_final_verifier=final_verifier, max_sequence_paths=3, ) assert callback_paths == [(0, 1), (1, 0)] assert result.selection_mode == "sequence_dp" assert result.chunk_candidate_indices == (1, 0) assert result.sequence_path_rank == 2 assert result.sequence_paths_checked == 2 def test_k_best_sequence_all_paths_rejected_is_no_qualified_candidate(): chunks = ("第一段完整", "第二段完整") callback_paths = [] def generator(candidate_chunks, seed): return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) def final_verifier(sequence_result, candidate_chunks): callback_paths.append(sequence_result.chunk_candidate_indices) return _joined_verification("".join(candidate_chunks), candidate_chunks[0]) with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): run_adaptive_cascade( chunks, 20, generator, lambda trajectory, candidate_chunks, seed: ( _locally_safe_join_rejected_verification(candidate_chunks, seed) ), max_candidates=2, sequence_final_verifier=final_verifier, max_sequence_paths=3, ) assert callback_paths == [(0, 1), (1, 0), (0, 0)] assert len(set(callback_paths)) == 3 def test_k_best_sequence_callback_exception_aborts_fail_closed(): chunks = ("第一段完整", "第二段完整") callback_paths = [] def callback(sequence_result, candidate_chunks): callback_paths.append(sequence_result.chunk_candidate_indices) raise ValueError("ASR backend disappeared") with pytest.raises(RuntimeError, match="refusing unverified audio"): run_adaptive_cascade( chunks, 20, lambda candidate_chunks, seed: tuple(candidate_chunks), lambda trajectory, candidate_chunks, seed: ( _locally_safe_join_rejected_verification(candidate_chunks, seed) ), max_candidates=2, sequence_final_verifier=callback, ) assert callback_paths == [(0, 1)] def test_k_best_sequence_invalid_callback_result_aborts_fail_closed(): chunks = ("第一段完整", "第二段完整") with pytest.raises(RuntimeError, match="invalid result"): run_adaptive_cascade( chunks, 20, lambda candidate_chunks, seed: tuple(candidate_chunks), lambda trajectory, candidate_chunks, seed: ( _locally_safe_join_rejected_verification(candidate_chunks, seed) ), max_candidates=2, sequence_final_verifier=lambda *args: None, ) def test_k_best_sequence_inconsistent_callback_evidence_never_passes(): chunks = ("第一段完整", "第二段完整") failed_candidate = verify_candidate( CandidateObservation("完整內容", "錯誤內容", 1.0) ) with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): run_adaptive_cascade( chunks, 20, lambda candidate_chunks, seed: tuple(candidate_chunks), lambda trajectory, candidate_chunks, seed: ( _locally_safe_join_rejected_verification(candidate_chunks, seed) ), max_candidates=2, sequence_final_verifier=lambda *args: TrajectoryGateResult( passed=True, candidate_results=(failed_candidate,), score=0.0, rejection_reasons=(), ), ) def test_sequence_final_callback_is_not_used_for_single_chunk_or_safe_whole(): callback_calls = [] with pytest.raises(NoQualifiedCandidateError): run_adaptive_cascade( ("單一完整段落",), 20, lambda candidate_chunks, seed: tuple(candidate_chunks), lambda trajectory, candidate_chunks, seed: TrajectoryGateResult( passed=False, candidate_results=( verify_candidate( CandidateObservation( candidate_chunks[0], candidate_chunks[0], 1.0, ) ), ), score=math.inf, rejection_reasons=("joined_output:semantic_gate",), chunk_artifacts=(ChunkCandidateArtifact(rms_db=-20.0),), ), max_candidates=2, sequence_final_verifier=lambda *args: callback_calls.append(args), ) assert callback_calls == [] result = run_adaptive_cascade( ("第一段", "第二段"), 30, lambda candidate_chunks, seed: tuple(candidate_chunks), lambda *args: _gate_result(True, 0.1), max_candidates=2, sequence_final_verifier=lambda *args: callback_calls.append(args), ) assert result.selection_mode == "whole_trajectory" assert callback_calls == [] def test_local_pass_join_fail_expands_and_selects_next_joined_safe_whole(): calls = [] def generator(chunks, seed): calls.append(seed) return tuple(f"seed-{seed}-{chunk}" for chunk in chunks) def verifier(trajectory, chunks, seed): observations = [] artifacts = [] for target in chunks: observation, artifact = _chunk_verification( target, target, embedding=[1.0, 0.0], rms_db=-20.0, ) observations.append(observation) artifacts.append(artifact) local = verify_trajectory(observations, chunk_artifacts=artifacts) joined_transcript = "".join(chunks) if seed == 11 else chunks[0] joined = _joined_verification("".join(chunks), joined_transcript) return qualify_trajectory_with_joined_output(local, joined) result = run_adaptive_cascade( ("第一段完整", "第二段完整"), 10, generator, verifier, max_candidates=2, ) assert calls == [10, 11] assert result.selection_mode == "whole_trajectory" assert result.candidate_index == 1 assert result.seed == 11 def test_all_local_pass_join_fail_can_only_return_as_sequence_dp(): chunks = ("第一段完整", "第二段完整") def generator(candidate_chunks, seed): return tuple(f"seed-{seed}-{chunk}" for chunk in candidate_chunks) def verifier(trajectory, candidate_chunks, seed): similarities = (0.9, 0.2) if seed == 20 else (0.2, 0.9) observations = [] artifacts = [] for target, similarity in zip(candidate_chunks, similarities, strict=True): observations.append( CandidateObservation( target_text=target, transcript_text=target, audio_duration_seconds=2.0, speaker_similarity=similarity, begin_speaker_similarity=similarity, end_speaker_similarity=similarity, ) ) artifacts.append( ChunkCandidateArtifact( speaker_embedding=np.array([1.0, 0.0], dtype=np.float32), rms_db=-20.0, ) ) local = verify_trajectory(observations, chunk_artifacts=artifacts) joined = _joined_verification("".join(candidate_chunks), candidate_chunks[0]) return qualify_trajectory_with_joined_output(local, joined) result = run_adaptive_cascade( chunks, 20, generator, verifier, max_candidates=2, ) assert result.selection_mode == "sequence_dp" assert result.candidate_index is None assert result.seed is None assert result.chunk_candidate_indices == (0, 1) assert result.attempted_seeds == (20, 21) def test_adaptive_cascade_returns_first_candidate_immediately(): generation_calls = [] verification_calls = [] def generator(chunks, seed): generation_calls.append((chunks, seed)) return tuple((chunk, seed) for chunk in chunks) def verifier(trajectory, chunks, seed): verification_calls.append((trajectory, chunks, seed)) assert all(chunk_seed == seed for _, chunk_seed in trajectory) return _gate_result(True, 0.2) result = run_adaptive_cascade( ["第一段", "第二段"], 12345, generator, verifier, ) assert isinstance(result, CascadeResult) assert result.seed == 12345 assert result.candidate_index == 0 assert result.attempted_seeds == (12345,) assert generation_calls == [(("第一段", "第二段"), 12345)] assert len(verification_calls) == 1 def test_marginal_stage_one_speaker_pass_expands_to_preferred_candidate(): calls = [] def generator(chunks, seed): calls.append(seed) return (chunks[0], seed) def verifier(trajectory, chunks, seed): if seed == 100: return _whole_speaker_verification( similarity=0.20, boundary_drop=0.04, ) if seed == 105: return _whole_speaker_verification( similarity=0.35, boundary_drop=0.02, ) return _whole_speaker_verification( similarity=0.30, boundary_drop=0.02, passed=False, ) result = run_adaptive_cascade( ["完整而且穩定的候選內容"], 100, generator, verifier, max_candidates=10, ) assert calls == list(range(100, 110)) assert result.selection_mode == "whole_trajectory" assert result.candidate_index == 5 assert result.seed == 105 def test_final_stage_can_return_marginal_hard_pass_to_preserve_availability(): calls = [] def generator(chunks, seed): calls.append(seed) return (chunks[0], seed) def verifier(trajectory, chunks, seed): return _whole_speaker_verification( similarity=0.20, boundary_drop=0.08, passed=seed == 200, ) result = run_adaptive_cascade( ["完整而且穩定的候選內容"], 200, generator, verifier, max_candidates=2, ) assert calls == [200, 201] assert result.seed == 200 assert result.candidate_index == 0 assert result.verification.passed def test_adaptive_cascade_expands_to_five_and_selects_lowest_verified_score(): calls = [] scores = { 100: None, 101: 0.4, 102: None, 103: 0.1, 104: 0.1, } def generator(chunks, seed): calls.append((chunks, seed)) return {"chunks": tuple((chunk, seed) for chunk in chunks), "seed": seed} def verifier(trajectory, chunks, seed): assert trajectory["seed"] == seed assert all(chunk_seed == seed for _, chunk_seed in trajectory["chunks"]) score = scores[seed] return _gate_result(score is not None, score or 0.0) result = run_adaptive_cascade( ["第一段", "第二段"], 100, generator, verifier, ) assert [seed for _, seed in calls] == [100, 101, 102, 103, 104] assert result.seed == 103 assert result.candidate_index == 3 assert result.verification.score == 0.1 assert result.attempted_seeds == (100, 101, 102, 103, 104) def test_adaptive_cascade_uses_ten_only_when_stage_five_has_no_verified_candidate(): calls = [] def generator(chunks, seed): calls.append(seed) return (chunks, seed) def verifier(trajectory, chunks, seed): return _gate_result(seed in {107, 109}, score=float(110 - seed)) result = run_adaptive_cascade( ["完整內容"], 100, generator, verifier, max_candidates=10, ) assert result.seed == 109 assert result.candidate_index == 9 assert result.attempted_seeds == tuple(range(100, 110)) assert calls == list(range(100, 110)) def test_sequence_fallback_can_take_a_prefix_from_a_and_suffix_from_b(): chunks = ("第一段完整", "第二段完整") trajectories = { 10: ("audio-a0", "audio-a1"), 11: ("audio-b0", "audio-b1"), } def generator(candidate_chunks, seed): assert candidate_chunks == chunks return trajectories[seed] def verifier(trajectory, candidate_chunks, seed): transcripts = ( candidate_chunks if seed == 999 else ( (candidate_chunks[0], "錯誤內容") if seed == 10 else ("錯誤內容", candidate_chunks[1]) ) ) observations = [] artifacts = [] for target, transcript in zip(candidate_chunks, transcripts, strict=True): observation, artifact = _chunk_verification( target, transcript, embedding=[1.0, 0.0], rms_db=-20.0, ) observations.append(observation) artifacts.append(artifact) return verify_trajectory(observations, chunk_artifacts=artifacts) result = run_adaptive_cascade( chunks, 10, generator, verifier, max_candidates=2, ) assert result.selection_mode == "sequence_dp" assert result.seed is None assert result.candidate_index is None assert result.trajectory == ("audio-a0", "audio-b1") assert result.chunk_candidate_indices == (0, 1) assert result.chunk_seeds == (10, 11) assert result.attempted_seeds == (10, 11) assert result.verification.passed assert len(result.verification.candidate_results) == 2 assert len(result.verification.chunk_artifacts) == 2 def test_same_seed_whole_trajectory_wins_before_sequence_fallback(): chunks = ("第一段完整", "第二段完整") def generator(candidate_chunks, seed): return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) def verifier(trajectory, candidate_chunks, seed): transcripts = ( (candidate_chunks[0], "錯誤內容") if seed == 20 else candidate_chunks ) observations = [] artifacts = [] for target, transcript in zip(candidate_chunks, transcripts, strict=True): observation, artifact = _chunk_verification( target, transcript, embedding=[1.0, 0.0], rms_db=-20.0, ) observations.append(observation) artifacts.append(artifact) return verify_trajectory(observations, chunk_artifacts=artifacts) result = run_adaptive_cascade( chunks, 20, generator, verifier, max_candidates=2, ) assert result.selection_mode == "whole_trajectory" assert result.seed == 21 assert result.candidate_index == 1 assert result.chunk_candidate_indices == (1, 1) assert result.chunk_seeds == (21, 21) assert result.trajectory == ("seed-21-chunk-0", "seed-21-chunk-1") def test_sequence_fallback_rejects_unsafe_transition_edge(): chunks = ("第一段完整", "第二段完整") def generator(candidate_chunks, seed): return tuple(f"seed-{seed}-chunk-{index}" for index in range(len(candidate_chunks))) def verifier(trajectory, candidate_chunks, seed): transcripts = ( (candidate_chunks[0], "錯誤內容") if seed == 30 else ("錯誤內容", candidate_chunks[1]) ) embeddings = ( ([1.0, 0.0], [1.0, 0.0]) if seed == 30 else ([1.0, 0.0], [1.0, 0.0, 0.0]) ) observations = [] artifacts = [] for target, transcript, embedding in zip( candidate_chunks, transcripts, embeddings, strict=True, ): observation, artifact = _chunk_verification( target, transcript, embedding=embedding, rms_db=-20.0, ) observations.append(observation) artifacts.append(artifact) return verify_trajectory(observations, chunk_artifacts=artifacts) with pytest.raises(NoQualifiedCandidateError, match="after 2 candidates"): run_adaptive_cascade( chunks, 30, generator, verifier, max_candidates=2, ) def test_transition_score_requires_safe_evidence_and_penalizes_rms_delta(): first_observation, first_artifact = _chunk_verification( "第一段完整", "第一段完整", embedding=[1.0, 0.0], rms_db=-20.0, ) second_observation, second_artifact = _chunk_verification( "第二段完整", "第二段完整", embedding=[1.0, 0.0], rms_db=-16.0, ) first = verify_trajectory([first_observation], chunk_artifacts=[first_artifact]) second = verify_trajectory([second_observation], chunk_artifacts=[second_artifact]) score = candidate_chunk_transition_score( first.candidate_results[0], first.chunk_artifacts[0], second.candidate_results[0], second.chunk_artifacts[0], ) assert score == pytest.approx(0.20) unsafe = ChunkCandidateArtifact(speaker_embedding=None, rms_db=-16.0) assert math.isinf( candidate_chunk_transition_score( first.candidate_results[0], first.chunk_artifacts[0], second.candidate_results[0], unsafe, ) ) @pytest.mark.parametrize( ("chunk_count", "expected_limit"), [(1, 10), (2, 10), (3, 6), (4, 5), (5, 4), (6, 3), (10, 2), (20, 1)], ) def test_candidate_limit_obeys_twenty_generated_chunk_budget(chunk_count, expected_limit): limit = candidate_limit_for_chunk_budget(chunk_count) assert limit == expected_limit assert limit * chunk_count <= 20 def test_candidate_limit_rejects_a_trajectory_that_already_exceeds_budget(): with pytest.raises(ValueError, match="one trajectory exceeds"): candidate_limit_for_chunk_budget(21) with pytest.raises(ValueError, match="between 1 and 10"): candidate_limit_for_chunk_budget(1, max_candidates=11) def test_360_character_profiles_stay_inside_generated_chunk_budget(): profiles = [ "甲" * 360, ("甲" * 11 + "。") * 30, ("甲" * 40 + "。") * 8 + "乙" * 32, ] for text in profiles: chunks = split_text_for_tts(text, max_chars=80, min_chunk_chars=12) chunks = split_leading_clause( chunks[0], search_chars=40, min_chunk_chars=12, ) + chunks[1:] limit = candidate_limit_for_chunk_budget(len(chunks)) assert limit * len(chunks) <= 20 def test_candidate_asr_value_error_rejects_seed_and_cascade_continues(): calls = [] def generator(chunks, seed): calls.append(seed) return _tone(seconds=0.25) def verifier(trajectory, chunks, seed): prepared = prepare_candidate_audio( trajectory, 16_000, transcriber=( (lambda *_: (_ for _ in ()).throw(ValueError("candidate ASR failure"))) if seed == 50 else (lambda *_: chunks[0]) ), ) if prepared is None: return _gate_result(False) return verify_trajectory( [ CandidateObservation( target_text=chunks[0], transcript_text=prepared.transcript_text, audio_duration_seconds=prepared.duration_seconds, ) ] ) result = run_adaptive_cascade( ["完整內容"], 50, generator, verifier, max_candidates=2, ) assert result.seed == 51 assert result.candidate_index == 1 assert calls == [50, 51] def test_adaptive_cascade_raises_without_qualified_candidate_and_never_falls_back(): generated = [] def generator(chunks, seed): generated.append(seed) return (chunks, seed) with pytest.raises(NoQualifiedCandidateError, match="after 5 candidates"): run_adaptive_cascade( ["必須完整"], 7, generator, lambda trajectory, chunks, seed: _gate_result(False), ) assert generated == [7, 8, 9, 10, 11] def test_adaptive_cascade_validates_contract_and_rejects_non_gate_verifier(): generator = lambda chunks, seed: (chunks, seed) with pytest.raises(ValueError, match="start with exactly one"): run_adaptive_cascade( ["內容"], 1, generator, lambda *args: _gate_result(True), initial_candidates=2, ) with pytest.raises(ValueError, match="between 1 and 10"): run_adaptive_cascade( ["內容"], 1, generator, lambda *args: _gate_result(True), max_candidates=11, ) with pytest.raises(TypeError, match="TrajectoryGateResult"): run_adaptive_cascade(["內容"], 1, generator, lambda *args: True)