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}