1pub use kcode_k1_audio_classification_projection::{
2 ExecutedAnalysis, FragmentId, FragmentStageV1, FragmentStatus, LlmJobState, LlmJobStatus,
3 OverallState, SpeakerLabelV1, StageState, StageStatus,
4};
5
6use kcode_k1_audio_classification_projection::{Projection, ProjectionEffect};
7use kcode_k1_audio_fragment_runner as runner;
8use kcode_k1_audio_fragment_submit as fragment_submit;
9use kcode_k1_audio_fragment_transactions as transactions;
10use kcode_k1_objects::K1Objects;
11use kcode_k1_peering::K1Peering;
12use kcode_k1_txn_ordering::{K1TxnOrdering, Subsystem, SubsystemId, TxId};
13use kcode_speaker_v3_analysis::Analyzer;
14use std::collections::{HashMap, HashSet};
15use std::future::Future;
16use std::path::Path;
17#[cfg(test)]
18use std::path::PathBuf;
19use std::pin::Pin;
20use std::rc::Rc;
21use std::sync::atomic::{AtomicBool, AtomicU64, Ordering as AtomicOrdering};
22use std::sync::{Arc, Mutex};
23use std::thread::{self, JoinHandle};
24use tokio::sync::mpsc;
25use tokio::task::JoinHandle as LocalJoinHandle;
26
27const RESTART_ERROR: &str = "analysis interrupted by restart";
28const SUBSYSTEM_NAME: &str = "audio-classification";
29
30pub struct AudioClassification {
31 projection: Arc<Projection>,
32 _ordering: Arc<K1TxnOrdering>,
33 peering: Arc<K1Peering>,
34 objects: Arc<K1Objects>,
35 control: Arc<Control>,
36 health: Arc<Health>,
37 worker: Option<JoinHandle<()>>,
38}
39
40impl AudioClassification {
41 pub fn open(
42 root: &Path,
43 ordering: Arc<K1TxnOrdering>,
44 peering: Arc<K1Peering>,
45 objects: Arc<K1Objects>,
46 analyzer: Analyzer,
47 ) -> Result<Self, String> {
48 Self::open_with_engine(
49 root,
50 ordering,
51 peering,
52 objects,
53 Box::new(ProductionEngine { analyzer }),
54 )
55 }
56
57 fn open_with_engine(
58 root: &Path,
59 ordering: Arc<K1TxnOrdering>,
60 peering: Arc<K1Peering>,
61 objects: Arc<K1Objects>,
62 engine: Box<dyn Engine>,
63 ) -> Result<Self, String> {
64 let (projection, cursor) = Projection::open(root, &ordering)?;
65 let projection = Arc::new(projection);
66 let control = Arc::new(Control::default());
67 let health = Arc::new(Health::default());
68 let replaying = Arc::new(AtomicBool::new(true));
69 let replayed = Arc::new(AtomicU64::new(0));
70 let callback = Arc::new(Callback {
71 projection: projection.clone(),
72 control: control.clone(),
73 health: health.clone(),
74 replaying: replaying.clone(),
75 replayed: replayed.clone(),
76 });
77 let (sender, receiver) = mpsc::unbounded_channel();
78 control.install(sender.clone());
79 let worker = spawn_worker(
80 engine,
81 sender,
82 receiver,
83 control.clone(),
84 health.clone(),
85 peering.clone(),
86 objects.clone(),
87 )?;
88 let startup = (|| -> Result<(), String> {
89 ordering
90 .register_subsystem(subsystem_id()?, cursor, callback)
91 .map_err(|error| format!("register audio classification: {error}"))?;
92 replaying.store(false, AtomicOrdering::Release);
93 record_replay(root, replayed.load(AtomicOrdering::Acquire));
94 let interrupted = projection
95 .running()
96 .map_err(|error| format!("read interrupted analyses: {error}"))?;
97 for fragment in interrupted {
98 transactions::submit_failure(
99 &peering,
100 fragment.fragment_id,
101 fragment.stage,
102 None,
103 RESTART_ERROR.to_owned(),
104 )
105 .map_err(|error| format!("persist interrupted analysis: {error}"))?;
106 }
107 control.activate();
108 let queued = projection
109 .queued()
110 .map_err(|error| format!("read queued analyses: {error}"))?;
111 for fragment_id in queued {
112 control.start(fragment_id, 0)?;
113 }
114 health.ensure()
115 })();
116 if let Err(error) = startup {
117 control.shutdown();
118 let _ = worker.join();
119 return Err(error);
120 }
121 Ok(Self {
122 projection,
123 _ordering: ordering,
124 peering,
125 objects,
126 control,
127 health,
128 worker: Some(worker),
129 })
130 }
131
132 pub fn submit(&self, ogg_bytes: &[u8]) -> Result<FragmentId, String> {
133 self.health.ensure()?;
134 fragment_submit::submit(&self.objects, &self.peering, ogg_bytes)
135 .map_err(|error| self.submission_error(error))
136 }
137
138 pub fn status(&self, fragment_id: FragmentId) -> Result<Option<FragmentStatus>, String> {
139 self.health.ensure()?;
140 self.projection.status(fragment_id)
141 }
142
143 pub fn retry(&self, fragment_id: FragmentId) -> Result<(), String> {
144 self.health.ensure()?;
145 let status = self
146 .projection
147 .status(fragment_id)?
148 .ok_or_else(|| "unknown audio fragment".to_owned())?;
149 if status.state != OverallState::Failed {
150 return Err("retry requires Failed state".to_owned());
151 }
152 if !self
153 .control
154 .start(fragment_id, status.attempt_count)
155 .map_err(|error| self.worker_error(error))?
156 {
157 return Err("retry is already active".to_owned());
158 }
159 Ok(())
160 }
161
162 pub fn discard(&self, fragment_id: FragmentId) -> Result<(), String> {
163 self.health.ensure()?;
164 let status = self
165 .projection
166 .status(fragment_id)?
167 .ok_or_else(|| "unknown audio fragment".to_owned())?;
168 if status.state == OverallState::Discarded {
169 return Ok(());
170 }
171 if !self.control.reserve_discard(fragment_id) {
172 return Err("discard is already active".to_owned());
173 }
174 match transactions::submit_discard(&self.peering, fragment_id) {
175 Ok(_) => {
176 self.control.discard_done(fragment_id);
177 Ok(())
178 }
179 Err(error) => {
180 if !is_committed_error(&error) {
181 self.control.discard_done(fragment_id);
182 }
183 Err(self.submission_error(error))
184 }
185 }
186 }
187
188 pub fn submit_labels(
189 &self,
190 fragment_id: FragmentId,
191 labels: Vec<SpeakerLabelV1>,
192 ) -> Result<(), String> {
193 self.health.ensure()?;
194 let interim = self.projection.validate_labels(fragment_id, &labels)?;
195 if !self.control.reserve_labels(fragment_id) {
196 return Err("label submission is already active".to_owned());
197 }
198 match transactions::submit_label_confirmation(&self.peering, fragment_id, interim, labels) {
199 Ok(_) => {
200 self.control.labels_done(fragment_id);
201 Ok(())
202 }
203 Err(error) => {
204 if !is_committed_error(&error) {
205 self.control.labels_done(fragment_id);
206 }
207 Err(self.submission_error(error))
208 }
209 }
210 }
211
212 fn submission_error(&self, error: String) -> String {
213 if is_committed_error(&error) {
214 self.health
215 .fault("transaction commitment is ambiguous".to_owned());
216 self.control.shutdown();
217 }
218 error
219 }
220
221 fn worker_error(&self, error: String) -> String {
222 self.health.fault(error.clone());
223 self.control.shutdown();
224 error
225 }
226
227 #[cfg(test)]
228 fn inject_error_burst(&self, id: FragmentId, errors: Vec<String>) -> Result<(), String> {
229 self.projection.inject_errors(id, errors)
230 }
231}
232
233impl Drop for AudioClassification {
234 fn drop(&mut self) {
235 self.control.shutdown();
236 if let Some(worker) = self.worker.take() {
237 let _ = worker.join();
238 }
239 }
240}
241
242#[derive(Default)]
243struct Health {
244 reopen: AtomicBool,
245 diagnostic: Mutex<Option<String>>,
246}
247
248impl Health {
249 fn ensure(&self) -> Result<(), String> {
250 if !self.reopen.load(AtomicOrdering::Acquire) {
251 return Ok(());
252 }
253 let detail = lock(&self.diagnostic)
254 .clone()
255 .unwrap_or_else(|| "processing fault".to_owned());
256 Err(format!("audio classification requires reopen: {detail}"))
257 }
258
259 fn fault(&self, diagnostic: String) {
260 let mut current = lock(&self.diagnostic);
261 if current.is_none() {
262 *current = Some(diagnostic);
263 }
264 self.reopen.store(true, AtomicOrdering::Release);
265 }
266}
267
268#[derive(Default)]
269struct Control {
270 state: Mutex<ControlState>,
271 live: AtomicBool,
272}
273
274#[derive(Default)]
275struct ControlState {
276 sender: Option<mpsc::UnboundedSender<WorkerCommand>>,
277 starting: HashMap<FragmentId, StartReservation>,
278 labels: HashSet<FragmentId>,
279 discards: HashSet<FragmentId>,
280 next_generation: u64,
281}
282
283struct StartReservation {
284 generation: u64,
285 baseline_attempt: u32,
286}
287
288impl Control {
289 fn install(&self, sender: mpsc::UnboundedSender<WorkerCommand>) {
290 lock(&self.state).sender = Some(sender);
291 }
292
293 fn activate(&self) {
294 self.live.store(true, AtomicOrdering::Release);
295 }
296
297 fn start(&self, id: FragmentId, baseline_attempt: u32) -> Result<bool, String> {
298 if !self.live.load(AtomicOrdering::Acquire) {
299 return Err("audio classification worker is unavailable".to_owned());
300 }
301 let mut state = lock(&self.state);
302 if state.starting.contains_key(&id) {
303 return Ok(false);
304 }
305 state.next_generation = state
306 .next_generation
307 .checked_add(1)
308 .ok_or_else(|| "audio classification worker generation overflow".to_owned())?;
309 let generation = state.next_generation;
310 let sender = state
311 .sender
312 .clone()
313 .ok_or_else(|| "audio classification worker is unavailable".to_owned())?;
314 state.starting.insert(
315 id,
316 StartReservation {
317 generation,
318 baseline_attempt,
319 },
320 );
321 drop(state);
322 if sender.send(WorkerCommand::Start(id, generation)).is_err() {
323 self.finished(id, generation);
324 return Err("audio classification worker is unavailable".to_owned());
325 }
326 Ok(true)
327 }
328
329 fn progress_applied(&self, id: FragmentId, attempt_count: u32) {
330 let mut state = lock(&self.state);
331 if state
332 .starting
333 .get(&id)
334 .is_some_and(|value| attempt_count > value.baseline_attempt)
335 {
336 state.starting.remove(&id);
337 }
338 }
339
340 fn abort(&self, id: FragmentId) -> Result<(), String> {
341 let sender = {
342 let mut state = lock(&self.state);
343 state.starting.remove(&id);
344 state.labels.remove(&id);
345 state.discards.remove(&id);
346 state.sender.clone()
347 };
348 if self.live.load(AtomicOrdering::Acquire)
349 && sender.is_some_and(|sender| sender.send(WorkerCommand::Abort(id)).is_err())
350 {
351 return Err("audio classification worker is unavailable".to_owned());
352 }
353 Ok(())
354 }
355
356 fn finished(&self, id: FragmentId, generation: u64) {
357 let mut state = lock(&self.state);
358 if state
359 .starting
360 .get(&id)
361 .is_some_and(|value| value.generation == generation)
362 {
363 state.starting.remove(&id);
364 }
365 }
366
367 fn reserve_labels(&self, id: FragmentId) -> bool {
368 lock(&self.state).labels.insert(id)
369 }
370
371 fn labels_done(&self, id: FragmentId) {
372 lock(&self.state).labels.remove(&id);
373 }
374
375 fn reserve_discard(&self, id: FragmentId) -> bool {
376 lock(&self.state).discards.insert(id)
377 }
378
379 fn discard_done(&self, id: FragmentId) {
380 lock(&self.state).discards.remove(&id);
381 }
382
383 fn shutdown(&self) {
384 self.live.store(false, AtomicOrdering::Release);
385 let sender = lock(&self.state).sender.clone();
386 if let Some(sender) = sender {
387 let _ = sender.send(WorkerCommand::Stop);
388 }
389 }
390}
391
392struct Callback {
393 projection: Arc<Projection>,
394 control: Arc<Control>,
395 health: Arc<Health>,
396 replaying: Arc<AtomicBool>,
397 replayed: Arc<AtomicU64>,
398}
399
400impl Callback {
401 fn fault(&self, diagnostic: String) -> String {
402 self.health.fault(diagnostic.clone());
403 self.control.shutdown();
404 diagnostic
405 }
406
407 fn react(&self, fragment_id: FragmentId, effect: ProjectionEffect) -> Result<(), String> {
408 match effect {
409 ProjectionEffect::Start => {
410 if self.control.live.load(AtomicOrdering::Acquire) {
411 self.control.start(fragment_id, 0)?;
412 }
413 }
414 ProjectionEffect::Abort => self.control.abort(fragment_id)?,
415 ProjectionEffect::LabelsCommitted => self.control.labels_done(fragment_id),
416 ProjectionEffect::None => {
417 let status = self
418 .projection
419 .status(fragment_id)?
420 .ok_or_else(|| "applied event has no projected fragment".to_owned())?;
421 self.control
422 .progress_applied(fragment_id, status.attempt_count);
423 }
424 }
425 Ok(())
426 }
427}
428
429impl Subsystem for Callback {
430 fn submit_txn(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
431 let applied = self
432 .projection
433 .apply(id, payload)
434 .map_err(|error| self.fault(format!("apply audio classification event: {error}")))?;
435 if self.replaying.load(AtomicOrdering::Acquire) {
436 self.replayed.fetch_add(1, AtomicOrdering::AcqRel);
437 return Ok(());
438 }
439 self.react(applied.fragment_id, applied.effect)
440 .map_err(|error| self.fault(format!("apply audio classification effect: {error}")))
441 }
442
443 fn reorg(&self) -> Result<(), String> {
444 let result = self.projection.clear();
445 self.health
446 .fault("canonical reorganization requires reopen".to_owned());
447 self.control.shutdown();
448 result
449 }
450}
451
452type EngineFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + 'a>>;
453
454trait Engine: Send + 'static {
455 fn run<'a>(
456 &'a self,
457 peering: &'a K1Peering,
458 fragment_id: FragmentId,
459 ogg_bytes: &'a [u8],
460 ) -> EngineFuture<'a>;
461}
462
463struct ProductionEngine {
464 analyzer: Analyzer,
465}
466
467impl Engine for ProductionEngine {
468 fn run<'a>(
469 &'a self,
470 peering: &'a K1Peering,
471 fragment_id: FragmentId,
472 ogg_bytes: &'a [u8],
473 ) -> EngineFuture<'a> {
474 Box::pin(runner::run(&self.analyzer, peering, fragment_id, ogg_bytes))
475 }
476}
477
478enum WorkerCommand {
479 Start(FragmentId, u64),
480 Abort(FragmentId),
481 Finished(FragmentId, u64, Result<(), String>),
482 Stop,
483}
484
485fn spawn_worker(
486 engine: Box<dyn Engine>,
487 sender: mpsc::UnboundedSender<WorkerCommand>,
488 receiver: mpsc::UnboundedReceiver<WorkerCommand>,
489 control: Arc<Control>,
490 health: Arc<Health>,
491 peering: Arc<K1Peering>,
492 objects: Arc<K1Objects>,
493) -> Result<JoinHandle<()>, String> {
494 thread::Builder::new()
495 .name("k1-audio-classification".to_owned())
496 .spawn(move || worker_main(engine, sender, receiver, control, health, peering, objects))
497 .map_err(|error| format!("start audio classification worker: {error}"))
498}
499
500fn worker_main(
501 engine: Box<dyn Engine>,
502 sender: mpsc::UnboundedSender<WorkerCommand>,
503 mut receiver: mpsc::UnboundedReceiver<WorkerCommand>,
504 control: Arc<Control>,
505 health: Arc<Health>,
506 peering: Arc<K1Peering>,
507 objects: Arc<K1Objects>,
508) {
509 let runtime = match tokio::runtime::Builder::new_current_thread().build() {
510 Ok(runtime) => runtime,
511 Err(error) => {
512 health.fault(format!("create worker runtime: {error}"));
513 control.shutdown();
514 return;
515 }
516 };
517 let local = tokio::task::LocalSet::new();
518 runtime.block_on(local.run_until(async move {
519 let engine: Rc<dyn Engine> = Rc::from(engine);
520 let mut tasks: HashMap<FragmentId, (u64, LocalJoinHandle<()>)> = HashMap::new();
521 while let Some(command) = receiver.recv().await {
522 match command {
523 WorkerCommand::Start(id, generation) => {
524 if let Some((old_generation, old_task)) = tasks.remove(&id) {
525 old_task.abort();
526 control.finished(id, old_generation);
527 }
528 let task_engine = engine.clone();
529 let task_peering = peering.clone();
530 let task_objects = objects.clone();
531 let task_sender = sender.clone();
532 let task = tokio::task::spawn_local(async move {
533 let result =
534 run_fragment(task_engine.as_ref(), &task_peering, &task_objects, id)
535 .await;
536 let _ = task_sender.send(WorkerCommand::Finished(id, generation, result));
537 });
538 tasks.insert(id, (generation, task));
539 }
540 WorkerCommand::Abort(id) => {
541 if let Some((generation, task)) = tasks.remove(&id) {
542 task.abort();
543 control.finished(id, generation);
544 }
545 }
546 WorkerCommand::Finished(id, generation, result) => {
547 if tasks
548 .get(&id)
549 .is_some_and(|(current, _)| *current == generation)
550 {
551 tasks.remove(&id);
552 control.finished(id, generation);
553 }
554 if let Err(error) = result {
555 health.fault(format!("runner persistence failure: {error}"));
556 control.shutdown();
557 break;
558 }
559 }
560 WorkerCommand::Stop => break,
561 }
562 }
563 for (id, (generation, task)) in tasks {
564 task.abort();
565 control.finished(id, generation);
566 }
567 }));
568}
569
570async fn run_fragment(
571 engine: &dyn Engine,
572 peering: &K1Peering,
573 objects: &K1Objects,
574 id: FragmentId,
575) -> Result<(), String> {
576 let object = match objects.load(id) {
577 Ok(Some(object)) if object.file_type == "audio/ogg" => object,
578 Ok(Some(_)) => {
579 transactions::submit_failure(
580 peering,
581 id,
582 FragmentStageV1::Queue,
583 None,
584 "audio Object is not audio/ogg".to_owned(),
585 )?;
586 return Ok(());
587 }
588 Ok(None) => {
589 transactions::submit_failure(
590 peering,
591 id,
592 FragmentStageV1::Queue,
593 None,
594 "audio Object is unavailable".to_owned(),
595 )?;
596 return Ok(());
597 }
598 Err(error) => {
599 transactions::submit_failure(
600 peering,
601 id,
602 FragmentStageV1::Queue,
603 None,
604 format!("load audio Object: {error}"),
605 )?;
606 return Ok(());
607 }
608 };
609 engine.run(peering, id, &object.data).await
610}
611
612fn subsystem_id() -> Result<SubsystemId, String> {
613 SubsystemId::from_str(SUBSYSTEM_NAME)
614}
615
616fn is_committed_error(error: &str) -> bool {
617 error.to_ascii_lowercase().contains("committed")
618}
619
620fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
621 mutex.lock().unwrap_or_else(|error| error.into_inner())
622}
623
624#[cfg(test)]
625fn record_replay(root: &Path, count: u64) {
626 lock(replay_counts()).insert(root.to_path_buf(), count);
627}
628
629#[cfg(not(test))]
630fn record_replay(_root: &Path, _count: u64) {}
631
632#[cfg(test)]
633fn replay_counts() -> &'static Mutex<HashMap<PathBuf, u64>> {
634 use std::sync::OnceLock;
635 static COUNTS: OnceLock<Mutex<HashMap<PathBuf, u64>>> = OnceLock::new();
636 COUNTS.get_or_init(|| Mutex::new(HashMap::new()))
637}
638
639#[cfg(test)]
640mod tests {
641 use super::*;
642 use kcode_k1_audio_classification_testkit::{
643 AdapterError, AdapterFactory, ClassificationAdapter, OpenRequest,
644 SUCCESS_INTERIM_TRANSCRIPT, ScriptedOutcome, TestId, TestJob, TestJobState, TestState,
645 TestStatus, run_all,
646 };
647 use kcode_k1_audio_fragment_transactions::ProgressUpdateV1;
648 use kcode_k1_txn_ordering::REGISTER_AT_TIP;
649 use kcode_speaker_v3_analysis::{
650 AnalysisEnvelope, FeatureVector24, GeminiCohort, OggAudioMetadata, StructuredAnalysis,
651 StructuredSpeaker, StructurerProvenance,
652 };
653 use std::collections::VecDeque;
654 use tokio::sync::oneshot;
655
656 const PROJECTION_DATABASE_FILE: &str = "audio-classification.sqlite3";
657
658 struct ScriptedEngine {
659 outcomes: Mutex<VecDeque<ScriptedOutcome>>,
660 }
661
662 impl Engine for ScriptedEngine {
663 fn run<'a>(
664 &'a self,
665 peering: &'a K1Peering,
666 id: FragmentId,
667 bytes: &'a [u8],
668 ) -> EngineFuture<'a> {
669 Box::pin(async move {
670 let Some(mut outcome) = lock(&self.outcomes).pop_front() else {
671 return Ok(());
672 };
673 transactions::submit_progress(
674 peering,
675 id,
676 ProgressUpdateV1::LlmJobStarted {
677 sequence: 1,
678 stage: FragmentStageV1::Transcript,
679 name: "transcript".to_owned(),
680 },
681 )?;
682 loop {
683 match outcome {
684 ScriptedOutcome::Blocked {
685 gate,
686 outcome: next,
687 } => {
688 let (sender, receiver) = oneshot::channel();
689 thread::Builder::new()
690 .name("audio-test-gate".to_owned())
691 .spawn(move || {
692 gate.wait_for_release();
693 let _ = sender.send(());
694 })
695 .map_err(|error| error.to_string())?;
696 receiver
697 .await
698 .map_err(|_| "test gate cancelled".to_owned())?;
699 outcome = *next;
700 }
701 ScriptedOutcome::Success => {
702 transactions::submit_progress(
703 peering,
704 id,
705 ProgressUpdateV1::LlmJobSucceeded { sequence: 1 },
706 )?;
707 transactions::submit_transcription_complete(
708 peering,
709 id,
710 success_analysis(bytes)?,
711 )?;
712 return Ok(());
713 }
714 ScriptedOutcome::Failure { message } => {
715 transactions::submit_failure(
716 peering,
717 id,
718 FragmentStageV1::Transcript,
719 Some(1),
720 message,
721 )?;
722 return Ok(());
723 }
724 }
725 }
726 })
727 }
728 }
729
730 fn success_analysis(bytes: &[u8]) -> Result<ExecutedAnalysis, String> {
731 let speaker_one = "Speaker 1".parse().map_err(|error| format!("{error}"))?;
732 let speaker_two = "Speaker 2".parse().map_err(|error| format!("{error}"))?;
733 let provenance = StructurerProvenance {
734 model_id: "test-model".to_owned(),
735 prompt_revision: "test-prompt".to_owned(),
736 };
737 Ok(ExecutedAnalysis {
738 envelope: AnalysisEnvelope {
739 audio: OggAudioMetadata::from_ogg_bytes(bytes)
740 .map_err(|error| error.to_string())?,
741 analysis: StructuredAnalysis {
742 transcript: SUCCESS_INTERIM_TRANSCRIPT.to_owned(),
743 speakers: vec![
744 StructuredSpeaker {
745 speaker: speaker_one,
746 language: "English".to_owned(),
747 features: FeatureVector24::default(),
748 features_usable_for_training: true,
749 },
750 StructuredSpeaker {
751 speaker: speaker_two,
752 language: "English".to_owned(),
753 features: FeatureVector24::default(),
754 features_usable_for_training: true,
755 },
756 ],
757 },
758 gemini: GeminiCohort {
759 model_id: "test-gemini".to_owned(),
760 transcript_prompt_revision: "test-transcript".to_owned(),
761 feature_prompt_revisions: std::array::from_fn(|_| "test-feature".to_owned()),
762 feature_schema_revision: "test-schema".to_owned(),
763 },
764 structurer: provenance.clone(),
765 },
766 label_extractor: provenance,
767 })
768 }
769
770 struct TestFactory;
771
772 struct TestAdapter {
773 facade: AudioClassification,
774 }
775
776 impl AdapterFactory for TestFactory {
777 type Adapter = TestAdapter;
778
779 fn open(&self, request: OpenRequest) -> Result<Self::Adapter, AdapterError> {
780 let ordering = Arc::new(
781 K1TxnOrdering::open(&request.root.join("ordering")).map_err(adapter_error)?,
782 );
783 let peering = Arc::new(
784 K1Peering::open(&request.root.join("peering"), ordering.clone())
785 .map_err(adapter_error)?,
786 );
787 let objects = Arc::new(
788 K1Objects::open(ordering.clone(), peering.clone()).map_err(adapter_error)?,
789 );
790 let facade = AudioClassification::open_with_engine(
791 &request.root.join("projection"),
792 ordering,
793 peering,
794 objects,
795 Box::new(ScriptedEngine {
796 outcomes: Mutex::new(request.outcomes.into()),
797 }),
798 )
799 .map_err(adapter_error)?;
800 Ok(TestAdapter { facade })
801 }
802
803 fn append_queued_while_closed(
804 &self,
805 root: &Path,
806 bytes: &[u8],
807 ) -> Result<TestId, AdapterError> {
808 let ordering =
809 Arc::new(K1TxnOrdering::open(&root.join("ordering")).map_err(adapter_error)?);
810 let peering = Arc::new(
811 K1Peering::open(&root.join("peering"), ordering.clone()).map_err(adapter_error)?,
812 );
813 let objects = Arc::new(
814 K1Objects::open(ordering.clone(), peering.clone()).map_err(adapter_error)?,
815 );
816 ordering
817 .register_subsystem(
818 subsystem_id().map_err(adapter_error)?,
819 Some(REGISTER_AT_TIP),
820 Arc::new(NoopSubsystem),
821 )
822 .map_err(adapter_error)?;
823 fragment_submit::submit(&objects, &peering, bytes)
824 .map(to_test_id)
825 .map_err(adapter_error)
826 }
827
828 fn corrupt_projection(&self, root: &Path) -> Result<(), AdapterError> {
829 std::fs::write(
830 root.join("projection").join(PROJECTION_DATABASE_FILE),
831 b"corrupt projection",
832 )
833 .map_err(|error| adapter_error(error.to_string()))
834 }
835
836 fn startup_replay_count(&self, root: &Path) -> Result<u64, AdapterError> {
837 lock(replay_counts())
838 .get(&root.join("projection"))
839 .copied()
840 .ok_or_else(|| adapter_error("startup replay count is unavailable".to_owned()))
841 }
842 }
843
844 impl ClassificationAdapter for TestAdapter {
845 fn submit(&self, bytes: &[u8]) -> Result<TestId, AdapterError> {
846 self.facade
847 .submit(bytes)
848 .map(to_test_id)
849 .map_err(adapter_error)
850 }
851
852 fn status(&self, id: TestId) -> Result<TestStatus, AdapterError> {
853 let status = self
854 .facade
855 .status(from_test_id(id))
856 .map_err(adapter_error)?
857 .ok_or_else(|| adapter_error("status is unavailable".to_owned()))?;
858 Ok(TestStatus {
859 state: match status.state {
860 OverallState::Queued => TestState::Queued,
861 OverallState::Running => TestState::Running,
862 OverallState::Failed => TestState::Failed,
863 OverallState::Completed => TestState::Completed,
864 OverallState::Confirmed => TestState::Confirmed,
865 OverallState::Discarded => TestState::Discarded,
866 },
867 attempt_count: status.attempt_count,
868 jobs: status
869 .jobs
870 .into_iter()
871 .map(|job| TestJob {
872 attempt: job.attempt,
873 sequence: job.sequence,
874 state: match job.state {
875 LlmJobState::Running => TestJobState::Running,
876 LlmJobState::Succeeded => TestJobState::Succeeded,
877 LlmJobState::Failed => TestJobState::Failed,
878 },
879 })
880 .collect(),
881 interim_transcript: status
882 .analysis
883 .as_ref()
884 .map(|analysis| analysis.envelope.analysis.transcript.clone()),
885 final_transcript: status.final_transcript,
886 labels: status
887 .confirmed_labels
888 .into_iter()
889 .map(|label| label.person_id)
890 .collect(),
891 errors: status.errors,
892 errors_truncated: status.errors_truncated,
893 })
894 }
895
896 fn retry(&self, id: TestId) -> Result<(), AdapterError> {
897 self.facade.retry(from_test_id(id)).map_err(adapter_error)
898 }
899
900 fn discard(&self, id: TestId) -> Result<(), AdapterError> {
901 self.facade.discard(from_test_id(id)).map_err(adapter_error)
902 }
903
904 fn submit_labels(&self, id: TestId, labels: Vec<String>) -> Result<(), AdapterError> {
905 let fragment_id = from_test_id(id);
906 let status = self
907 .facade
908 .status(fragment_id)
909 .map_err(adapter_error)?
910 .ok_or_else(|| adapter_error("status is unavailable".to_owned()))?;
911 let speakers = status
912 .analysis
913 .ok_or_else(|| adapter_error("analysis is unavailable".to_owned()))?
914 .envelope
915 .analysis
916 .speakers;
917 if speakers.len() != labels.len() {
918 return Err(adapter_error("label count mismatch".to_owned()));
919 }
920 let labels = speakers
921 .into_iter()
922 .zip(labels)
923 .map(|(speaker, person_id)| SpeakerLabelV1 {
924 speaker: speaker.speaker,
925 person_id,
926 })
927 .collect();
928 self.facade
929 .submit_labels(fragment_id, labels)
930 .map_err(adapter_error)
931 }
932
933 fn inject_error_burst(&self, id: TestId, errors: Vec<String>) -> Result<(), AdapterError> {
934 self.facade
935 .inject_error_burst(from_test_id(id), errors)
936 .map_err(adapter_error)
937 }
938 }
939
940 struct NoopSubsystem;
941
942 impl Subsystem for NoopSubsystem {
943 fn submit_txn(&self, _id: TxId, _payload: &[u8]) -> Result<(), String> {
944 Ok(())
945 }
946
947 fn reorg(&self) -> Result<(), String> {
948 Ok(())
949 }
950 }
951
952 fn to_test_id(id: FragmentId) -> TestId {
953 id.into_bytes()
954 }
955
956 fn from_test_id(id: TestId) -> FragmentId {
957 FragmentId::from_bytes(id)
958 }
959
960 fn adapter_error(message: String) -> AdapterError {
961 Box::new(std::io::Error::other(message))
962 }
963
964 #[test]
965 fn published_conformance() {
966 run_all(&TestFactory).expect("audio classification conformance");
967 }
968}