Skip to main content

kcode_k1_audio_classification_driver/
lib.rs

1pub use kcode_k1_audio_fragment_runner::FragmentId;
2
3use kcode_k1_audio_fragment_runner as runner;
4use kcode_k1_audio_fragment_transactions::{self as transactions, FragmentStageV1};
5use kcode_k1_objects::K1Objects;
6use kcode_k1_peering::K1Peering;
7use kcode_speaker_v3_analysis::Analyzer;
8use std::collections::{HashMap, VecDeque};
9use std::future::Future;
10use std::pin::Pin;
11use std::sync::atomic::{AtomicBool, Ordering};
12use std::sync::{Arc, Mutex};
13use std::thread;
14
15pub const MAX_ACTIVE_ATTEMPTS: usize = 8;
16pub type EngineFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + 'a>>;
17
18pub trait FragmentEngine: Send + Sync + 'static {
19    fn run<'a>(
20        &'a self,
21        fragment_id: FragmentId,
22        ogg_bytes: &'a [u8],
23        is_active: &'a (dyn Fn() -> bool + Send + Sync),
24    ) -> EngineFuture<'a>;
25}
26
27#[derive(Clone, Copy, Debug, Eq, PartialEq)]
28pub enum StartOutcome {
29    Started,
30    Pending,
31    AlreadyActive,
32}
33
34pub struct AudioClassificationDriver {
35    inner: Arc<Inner>,
36}
37
38impl AudioClassificationDriver {
39    pub fn open(peering: Arc<K1Peering>, objects: Arc<K1Objects>, analyzer: Analyzer) -> Self {
40        let engine: Arc<dyn FragmentEngine> = Arc::new(RunnerEngine {
41            peering: peering.clone(),
42            analyzer: Arc::new(analyzer),
43        });
44        Self::with_engine(peering, objects, engine)
45    }
46
47    pub fn with_engine(
48        peering: Arc<K1Peering>,
49        objects: Arc<K1Objects>,
50        engine: Arc<dyn FragmentEngine>,
51    ) -> Self {
52        Self::with_resources(Arc::new(LiveResources { peering, objects }), engine)
53    }
54
55    pub fn start(&self, id: FragmentId) -> Result<StartOutcome, String> {
56        let (outcome, lane) = self.inner.reserve(id)?;
57        lane.map_or(Ok(()), |lane| self.inner.launch(lane))?;
58        Ok(outcome)
59    }
60
61    pub fn abort(&self, id: FragmentId) {
62        let mut state = lock(&self.inner.state);
63        if let Some(entry) = state.current.remove(&id) {
64            entry.active.store(false, Ordering::Release);
65        }
66        state.pending.retain(|(pending, _)| *pending != id);
67    }
68
69    pub fn ensure_healthy(&self) -> Result<(), String> {
70        match &lock(&self.inner.state).fault {
71            Some(error) => Err(format!("audio classification driver is unhealthy: {error}")),
72            None => Ok(()),
73        }
74    }
75
76    pub fn shutdown(&self) {
77        lock(&self.inner.state).stop(None);
78    }
79
80    fn with_resources(
81        resources: Arc<dyn FragmentResources>,
82        engine: Arc<dyn FragmentEngine>,
83    ) -> Self {
84        Self {
85            inner: Arc::new(Inner {
86                resources,
87                engine,
88                state: Mutex::new(State {
89                    accepting: true,
90                    ..State::default()
91                }),
92            }),
93        }
94    }
95}
96
97impl Drop for AudioClassificationDriver {
98    fn drop(&mut self) {
99        self.shutdown();
100    }
101}
102
103struct RunnerEngine {
104    peering: Arc<K1Peering>,
105    analyzer: Arc<Analyzer>,
106}
107
108impl FragmentEngine for RunnerEngine {
109    fn run<'a>(
110        &'a self,
111        id: FragmentId,
112        bytes: &'a [u8],
113        active: &'a (dyn Fn() -> bool + Send + Sync),
114    ) -> EngineFuture<'a> {
115        Box::pin(runner::run_while_active(
116            &self.analyzer,
117            &self.peering,
118            id,
119            bytes,
120            active,
121        ))
122    }
123}
124
125trait FragmentResources: Send + Sync + 'static {
126    fn load(&self, id: FragmentId) -> Result<Option<(String, Vec<u8>)>, String>;
127    fn fail_queue(&self, id: FragmentId, error: String) -> Result<(), String>;
128}
129
130struct LiveResources {
131    peering: Arc<K1Peering>,
132    objects: Arc<K1Objects>,
133}
134
135impl FragmentResources for LiveResources {
136    fn load(&self, id: FragmentId) -> Result<Option<(String, Vec<u8>)>, String> {
137        self.objects
138            .load(id)
139            .map(|object| object.map(|object| (object.file_type, object.data)))
140    }
141
142    fn fail_queue(&self, id: FragmentId, error: String) -> Result<(), String> {
143        transactions::submit_failure(&self.peering, id, FragmentStageV1::Queue, None, error)
144            .map(|_| ())
145    }
146}
147
148#[derive(Default)]
149struct State {
150    accepting: bool,
151    fault: Option<String>,
152    next_generation: u64,
153    lanes: usize,
154    current: HashMap<FragmentId, Entry>,
155    pending: VecDeque<(FragmentId, u64)>,
156}
157
158struct Entry {
159    generation: u64,
160    active: Arc<AtomicBool>,
161}
162
163struct Lane {
164    id: FragmentId,
165    generation: u64,
166    active: Arc<AtomicBool>,
167}
168
169impl State {
170    fn stop(&mut self, fault: Option<String>) {
171        self.accepting = false;
172        self.pending.clear();
173        for entry in self.current.drain().map(|(_, entry)| entry) {
174            entry.active.store(false, Ordering::Release);
175        }
176        self.fault = self.fault.take().or(fault);
177    }
178
179    fn next_lane(&mut self) -> Option<Lane> {
180        if !self.accepting || self.lanes == MAX_ACTIVE_ATTEMPTS {
181            return None;
182        }
183        let (id, generation) = self.pending.pop_front()?;
184        let entry = self
185            .current
186            .get(&id)
187            .filter(|entry| entry.generation == generation)?;
188        self.lanes += 1;
189        Some(Lane {
190            id,
191            generation,
192            active: entry.active.clone(),
193        })
194    }
195}
196
197struct Inner {
198    resources: Arc<dyn FragmentResources>,
199    engine: Arc<dyn FragmentEngine>,
200    state: Mutex<State>,
201}
202
203impl Inner {
204    fn reserve(&self, id: FragmentId) -> Result<(StartOutcome, Option<Lane>), String> {
205        let mut state = lock(&self.state);
206        if let Some(error) = &state.fault {
207            return Err(format!("audio classification driver is unhealthy: {error}"));
208        }
209        if !state.accepting {
210            return Err("audio classification driver is shut down".into());
211        }
212        if state.current.contains_key(&id) {
213            return Ok((StartOutcome::AlreadyActive, None));
214        }
215        state.next_generation = state
216            .next_generation
217            .checked_add(1)
218            .ok_or_else(|| "audio classification generation overflow".to_owned())?;
219        let generation = state.next_generation;
220        let active = Arc::new(AtomicBool::new(true));
221        state.current.insert(
222            id,
223            Entry {
224                generation,
225                active: active.clone(),
226            },
227        );
228        if state.lanes == MAX_ACTIVE_ATTEMPTS {
229            state.pending.push_back((id, generation));
230            return Ok((StartOutcome::Pending, None));
231        }
232        state.lanes += 1;
233        Ok((
234            StartOutcome::Started,
235            Some(Lane {
236                id,
237                generation,
238                active,
239            }),
240        ))
241    }
242
243    fn launch(self: &Arc<Self>, lane: Lane) -> Result<(), String> {
244        let id = lane.id;
245        let generation = lane.generation;
246        let inner = self.clone();
247        if let Err(error) = thread::Builder::new()
248            .name("k1-audio-fragment".to_owned())
249            .spawn(move || inner.run_lane(lane))
250        {
251            let error = format!("start audio classification lane: {error}");
252            self.finished(id, generation, Err(error.clone()));
253            return Err(error);
254        }
255        Ok(())
256    }
257
258    fn run_lane(self: Arc<Self>, lane: Lane) {
259        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
260            tokio::runtime::Builder::new_current_thread()
261                .enable_all()
262                .build()
263                .map_err(|error| format!("create audio classification runtime: {error}"))
264                .and_then(|runtime| {
265                    runtime.block_on(run_attempt(
266                        self.resources.as_ref(),
267                        self.engine.as_ref(),
268                        lane.id,
269                        lane.active.clone(),
270                    ))
271                })
272        }))
273        .unwrap_or_else(|_| Err("audio classification lane panicked".into()));
274        self.finished(lane.id, lane.generation, result);
275    }
276
277    fn finished(self: &Arc<Self>, id: FragmentId, generation: u64, result: Result<(), String>) {
278        let next = {
279            let mut state = lock(&self.state);
280            state.lanes = state.lanes.saturating_sub(1);
281            let active = state.current.get(&id).and_then(|entry| {
282                (entry.generation == generation).then(|| entry.active.load(Ordering::Acquire))
283            });
284            if active.is_some() {
285                state.current.remove(&id);
286            }
287            if let Some(error) = result.err().filter(|_| active == Some(true)) {
288                state.stop(Some(format!("runner persistence failure: {error}")));
289                None
290            } else {
291                state.next_lane()
292            }
293        };
294        if let Some(lane) = next {
295            let _ = self.launch(lane);
296        }
297    }
298}
299
300async fn run_attempt(
301    resources: &dyn FragmentResources,
302    engine: &dyn FragmentEngine,
303    id: FragmentId,
304    active: Arc<AtomicBool>,
305) -> Result<(), String> {
306    let object = match resources.load(id) {
307        Ok(Some(object)) if object.0 == "audio/ogg" => object,
308        Ok(Some(_)) => return fail_if_active(resources, id, &active, "wrong Object media type"),
309        Ok(None) => return fail_if_active(resources, id, &active, "audio Object is unavailable"),
310        Err(error) => {
311            let error = format!("load audio Object: {error}");
312            return fail_if_active(resources, id, &active, &error);
313        }
314    };
315    let is_active = || active.load(Ordering::Acquire);
316    if is_active() {
317        engine.run(id, &object.1, &is_active).await?;
318    }
319    Ok(())
320}
321
322fn fail_if_active(
323    resources: &dyn FragmentResources,
324    id: FragmentId,
325    active: &AtomicBool,
326    error: &str,
327) -> Result<(), String> {
328    if active.load(Ordering::Acquire) {
329        resources.fail_queue(id, error.to_owned())?;
330    }
331    Ok(())
332}
333
334fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
335    mutex.lock().unwrap_or_else(|error| error.into_inner())
336}
337
338#[cfg(test)]
339mod tests {
340    use super::*;
341    use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
342    use std::time::{Duration, Instant};
343
344    type Action = (Option<Arc<AtomicBool>>, Result<(), String>);
345
346    #[derive(Clone)]
347    enum Load {
348        Missing,
349        Wrong,
350        Error,
351        Wait(Arc<AtomicBool>),
352    }
353
354    #[derive(Default)]
355    struct Harness {
356        actions: Mutex<HashMap<FragmentId, VecDeque<Action>>>,
357        loads: Mutex<HashMap<FragmentId, Load>>,
358        load_calls: AtomicUsize,
359        failures: AtomicUsize,
360        polls: AtomicUsize,
361    }
362
363    impl FragmentEngine for Harness {
364        fn run<'a>(
365            &'a self,
366            id: FragmentId,
367            _bytes: &'a [u8],
368            _active: &'a (dyn Fn() -> bool + Send + Sync),
369        ) -> EngineFuture<'a> {
370            if id == FragmentId::from_bytes([0; 12]) {
371                return Box::pin(async {
372                    let mut command =
373                        tokio::process::Command::new(std::env::current_exe().unwrap());
374                    command.arg("--list").stdout(std::process::Stdio::null());
375                    tokio::time::timeout(Duration::from_secs(1), command.status())
376                        .await
377                        .unwrap()
378                        .map(|_| ())
379                        .map_err(|error| error.to_string())
380                });
381            }
382            let (gate, result) = lock(&self.actions)
383                .get_mut(&id)
384                .and_then(VecDeque::pop_front)
385                .unwrap_or((None, Ok(())));
386            Box::pin(async move {
387                self.polls.fetch_add(1, AtomicOrdering::SeqCst);
388                if let Some(gate) = gate {
389                    wait_gate(&gate);
390                }
391                result.inspect_err(|error| assert_ne!(error, "panic"))
392            })
393        }
394    }
395
396    impl FragmentResources for Harness {
397        fn load(&self, id: FragmentId) -> Result<Option<(String, Vec<u8>)>, String> {
398            self.load_calls.fetch_add(1, AtomicOrdering::SeqCst);
399            let load = lock(&self.loads).get(&id).cloned();
400            if let Some(Load::Wait(gate)) = &load {
401                wait_gate(gate);
402            }
403            match load {
404                None | Some(Load::Wait(_)) => Ok(Some(("audio/ogg".into(), vec![1]))),
405                Some(Load::Missing) => Ok(None),
406                Some(Load::Wrong) => Ok(Some(("text/plain".into(), vec![1]))),
407                Some(Load::Error) => Err("read failed".into()),
408            }
409        }
410
411        fn fail_queue(&self, _id: FragmentId, _error: String) -> Result<(), String> {
412            self.failures.fetch_add(1, AtomicOrdering::SeqCst);
413            Ok(())
414        }
415    }
416
417    fn wait_gate(gate: &AtomicBool) {
418        while !gate.load(Ordering::Acquire) {
419            thread::yield_now();
420        }
421    }
422
423    fn wait_for(condition: impl Fn() -> bool) {
424        let deadline = Instant::now() + Duration::from_secs(2);
425        while !condition() && Instant::now() < deadline {
426            thread::sleep(Duration::from_millis(2));
427        }
428        assert!(condition());
429    }
430
431    #[test]
432    fn runtime_isolation_admission_generation_failures_and_shutdown() {
433        let id = |value| FragmentId::from_bytes([value; 12]);
434        let gate = || Arc::new(AtomicBool::new(false));
435        let harness = Arc::new(Harness::default());
436        let driver = AudioClassificationDriver::with_resources(harness.clone(), harness.clone());
437        assert_eq!(driver.start(id(0)), Ok(StartOutcome::Started));
438        wait_for(|| lock(&driver.inner.state).lanes == 0);
439        let (first, rest) = (gate(), gate());
440        lock(&harness.actions).insert(id(1), vec![(Some(first.clone()), Ok(()))].into());
441        for value in 2..=9 {
442            lock(&harness.actions).insert(id(value), vec![(Some(rest.clone()), Ok(()))].into());
443        }
444        assert!((1..=8).all(|value| driver.start(id(value)) == Ok(StartOutcome::Started)));
445        assert_eq!(driver.start(id(9)), Ok(StartOutcome::Pending));
446        wait_for(|| harness.polls.load(AtomicOrdering::SeqCst) == 8);
447        driver.abort(id(1));
448        thread::sleep(Duration::from_millis(10));
449        assert_eq!(harness.polls.load(AtomicOrdering::SeqCst), 8);
450        first.store(true, Ordering::Release);
451        wait_for(|| harness.polls.load(AtomicOrdering::SeqCst) == 9);
452        rest.store(true, Ordering::Release);
453
454        let harness = Arc::new(Harness::default());
455        let driver = AudioClassificationDriver::with_resources(harness.clone(), harness.clone());
456        let (old, new) = (gate(), gate());
457        lock(&harness.actions).insert(
458            id(1),
459            vec![
460                (Some(old.clone()), Err("panic".into())),
461                (Some(new.clone()), Ok(())),
462            ]
463            .into(),
464        );
465        assert_eq!(driver.start(id(1)), Ok(StartOutcome::Started));
466        wait_for(|| harness.polls.load(AtomicOrdering::SeqCst) == 1);
467        driver.abort(id(1));
468        assert_eq!(driver.start(id(1)), Ok(StartOutcome::Started));
469        wait_for(|| harness.polls.load(AtomicOrdering::SeqCst) == 2);
470        old.store(true, Ordering::Release);
471        wait_for(|| lock(&driver.inner.state).lanes == 1);
472        assert_eq!(driver.ensure_healthy(), Ok(()));
473        assert_eq!(driver.start(id(1)), Ok(StartOutcome::AlreadyActive));
474        new.store(true, Ordering::Release);
475        wait_for(|| lock(&driver.inner.state).lanes == 0);
476        lock(&harness.actions).insert(id(2), vec![(None, Err("panic".into()))].into());
477        assert_eq!(driver.start(id(2)), Ok(StartOutcome::Started));
478        wait_for(|| driver.ensure_healthy().is_err());
479        assert_eq!(lock(&driver.inner.state).lanes, 0);
480        assert!(driver.start(id(3)).is_err());
481
482        let harness = Arc::new(Harness::default());
483        let driver = AudioClassificationDriver::with_resources(harness.clone(), harness.clone());
484        for (value, load) in [(1, Load::Missing), (2, Load::Wrong), (3, Load::Error)] {
485            lock(&harness.loads).insert(id(value), load);
486            assert_eq!(driver.start(id(value)), Ok(StartOutcome::Started));
487        }
488        wait_for(|| harness.failures.load(AtomicOrdering::SeqCst) == 3);
489        assert_eq!(harness.polls.load(AtomicOrdering::SeqCst), 0);
490        let load_gate = gate();
491        lock(&harness.loads).insert(id(4), Load::Wait(load_gate.clone()));
492        assert_eq!(driver.start(id(4)), Ok(StartOutcome::Started));
493        wait_for(|| harness.load_calls.load(AtomicOrdering::SeqCst) == 4);
494        assert_eq!(driver.start(id(5)), Ok(StartOutcome::Started));
495        wait_for(|| harness.polls.load(AtomicOrdering::SeqCst) == 1);
496        let started = Instant::now();
497        driver.shutdown();
498        assert!(started.elapsed() < Duration::from_millis(100));
499        assert!(driver.start(id(6)).is_err());
500        load_gate.store(true, Ordering::Release);
501        wait_for(|| lock(&driver.inner.state).lanes == 0);
502    }
503}