use std::path::{Path, PathBuf};
use std::sync::Arc;
use kcode_k1_audio_classification::{
AudioClassification, FragmentId, FragmentStatus, OverallState,
};
use kcode_k1_full_audio_domain::{
ManifestEntry, decode_manifest, encode_manifest, stitch_transcripts, validate_fragment_geometry,
};
use kcode_k1_objects::{K1Objects, Object};
const MANIFEST_FILE_TYPE: &str = "k1-full-audio-manifest-v1";
pub type FullAudioId = kcode_k1_objects::TxId;
pub struct K1FullAudio {
ffmpeg_path: PathBuf,
objects: Arc<K1Objects>,
classification: Arc<AudioClassification>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum FullAudioState {
Processing,
AwaitingLabels,
NeedsAttention,
Complete,
}
#[derive(Clone, Debug, PartialEq)]
pub struct FullAudioFragmentStatus {
pub fragment_id: FragmentId,
pub start_sample_48k: u64,
pub end_sample_48k: u64,
pub status: FragmentStatus,
}
#[derive(Clone, Debug, PartialEq)]
pub struct FullAudioStatus {
pub state: FullAudioState,
pub fragments: Vec<FullAudioFragmentStatus>,
pub final_transcript: Option<String>,
}
struct PendingFragment {
start_sample_48k: u64,
end_sample_48k: u64,
ogg_bytes: Vec<u8>,
}
impl K1FullAudio {
pub fn open(
ffmpeg_path: impl AsRef<Path>,
objects: Arc<K1Objects>,
classification: Arc<AudioClassification>,
) -> Result<Self, String> {
Ok(Self {
ffmpeg_path: validate_ffmpeg_path(ffmpeg_path.as_ref())?,
objects,
classification,
})
}
pub fn submit(&self, audio: &[u8]) -> Result<FullAudioId, String> {
submit_orchestration(
audio,
|input| kcode_audio_to_ogg_opus::convert_to_ogg_opus(&self.ffmpeg_path, input),
|ogg_bytes| {
kcode_ogg_opus_fragments::split_ogg_opus(ogg_bytes).map(|fragments| {
fragments
.into_iter()
.map(|fragment| PendingFragment {
start_sample_48k: fragment.start_sample_48k,
end_sample_48k: fragment.end_sample_48k,
ogg_bytes: fragment.ogg_bytes,
})
.collect()
})
},
|ogg_bytes| self.classification.submit(ogg_bytes),
|manifest| self.objects.save("", MANIFEST_FILE_TYPE, "", manifest),
)
}
pub fn status(&self, id: FullAudioId) -> Result<FullAudioStatus, String> {
status_orchestration(
id,
|manifest_id| self.objects.load(manifest_id),
|fragment_id| self.classification.status(fragment_id),
)
}
}
fn validate_ffmpeg_path(path: &Path) -> Result<PathBuf, String> {
if !path.is_absolute() {
return Err("ffmpeg path must be absolute".to_owned());
}
Ok(path.to_path_buf())
}
fn submit_orchestration<Convert, Split, Submit, Save>(
audio: &[u8],
convert: Convert,
split: Split,
mut submit: Submit,
save: Save,
) -> Result<FullAudioId, String>
where
Convert: FnOnce(&[u8]) -> Result<Vec<u8>, String>,
Split: FnOnce(&[u8]) -> Result<Vec<PendingFragment>, String>,
Submit: FnMut(&[u8]) -> Result<FragmentId, String>,
Save: FnOnce(&[u8]) -> Result<FullAudioId, String>,
{
let ogg_bytes = convert(audio)?;
let fragments = split(&ogg_bytes)?;
validate_fragment_geometry(
fragments
.iter()
.map(|fragment| (fragment.start_sample_48k, fragment.end_sample_48k)),
)?;
let mut entries = Vec::new();
entries
.try_reserve_exact(fragments.len())
.map_err(|error| format!("allocate manifest entries: {error}"))?;
for fragment in fragments {
let fragment_id = submit(&fragment.ogg_bytes)?;
entries.push(ManifestEntry {
fragment_id,
start_sample_48k: fragment.start_sample_48k,
end_sample_48k: fragment.end_sample_48k,
});
}
let manifest = encode_manifest(&entries)?;
save(&manifest)
}
fn status_orchestration<Load, Status>(
id: FullAudioId,
mut load: Load,
mut status: Status,
) -> Result<FullAudioStatus, String>
where
Load: FnMut(FullAudioId) -> Result<Option<Object>, String>,
Status: FnMut(FragmentId) -> Result<Option<FragmentStatus>, String>,
{
let object = load(id)?.ok_or_else(|| "unknown full-audio ID".to_owned())?;
if object.file_type != MANIFEST_FILE_TYPE
|| !object.filename.is_empty()
|| !object.description.is_empty()
{
return Err("object is not a full-audio manifest".to_owned());
}
let entries = decode_manifest(&object.data)?;
let mut fragments = Vec::new();
fragments
.try_reserve_exact(entries.len())
.map_err(|error| format!("allocate fragment statuses: {error}"))?;
for entry in entries {
let fragment_status = status(entry.fragment_id)?
.ok_or_else(|| format!("unknown audio fragment {}", entry.fragment_id))?;
fragments.push(FullAudioFragmentStatus {
fragment_id: entry.fragment_id,
start_sample_48k: entry.start_sample_48k,
end_sample_48k: entry.end_sample_48k,
status: fragment_status,
});
}
let state = derive_state(&fragments)?;
let final_transcript = if state == FullAudioState::Complete {
let transcripts = fragments
.iter()
.filter_map(|fragment| fragment.status.final_transcript.as_deref())
.collect::<Vec<_>>();
Some(stitch_transcripts(&transcripts))
} else {
None
};
Ok(FullAudioStatus {
state,
fragments,
final_transcript,
})
}
fn derive_state(fragments: &[FullAudioFragmentStatus]) -> Result<FullAudioState, String> {
if fragments.is_empty() {
return Err("manifest has no fragments".to_owned());
}
for fragment in fragments {
let confirmed = fragment.status.state == OverallState::Confirmed;
let has_transcript = fragment.status.final_transcript.is_some();
if confirmed && !has_transcript {
return Err("confirmed fragment has no final transcript".to_owned());
}
if !confirmed && has_transcript {
return Err("non-confirmed fragment has a final transcript".to_owned());
}
}
if has_state(fragments, &[OverallState::Failed, OverallState::Discarded]) {
return Ok(FullAudioState::NeedsAttention);
}
if has_state(fragments, &[OverallState::Queued, OverallState::Running]) {
return Ok(FullAudioState::Processing);
}
if fragments
.iter()
.all(|fragment| fragment.status.state == OverallState::Confirmed)
{
return Ok(FullAudioState::Complete);
}
if fragments.iter().all(|fragment| {
[OverallState::Completed, OverallState::Confirmed].contains(&fragment.status.state)
}) {
return Ok(FullAudioState::AwaitingLabels);
}
Err("classification returned an inconsistent state".to_owned())
}
fn has_state(fragments: &[FullAudioFragmentStatus], states: &[OverallState]) -> bool {
fragments
.iter()
.any(|fragment| states.contains(&fragment.status.state))
}
#[cfg(test)]
mod tests {
use super::*;
use FullAudioState::{AwaitingLabels, NeedsAttention, Processing};
use OverallState::{Completed, Confirmed, Failed, Queued, Running};
use kcode_k1_audio_classification::{FragmentStageV1, StageState, StageStatus};
use std::cell::{Cell, RefCell};
use std::sync::mpsc;
use std::thread;
use std::time::Duration;
fn id(value: u8) -> FragmentId {
FragmentId::from_bytes([value; 12])
}
fn entry(value: u8, start: u64, end: u64) -> ManifestEntry {
ManifestEntry {
fragment_id: id(value),
start_sample_48k: start,
end_sample_48k: end,
}
}
fn pending(value: u8, start: u64, end: u64) -> PendingFragment {
PendingFragment {
start_sample_48k: start,
end_sample_48k: end,
ogg_bytes: vec![value],
}
}
fn stage(stage: FragmentStageV1) -> StageStatus {
StageStatus {
stage,
state: StageState::Pending,
}
}
fn classification_status(state: OverallState, transcript: Option<&str>) -> FragmentStatus {
FragmentStatus {
state,
queue: stage(FragmentStageV1::Queue),
transcript: stage(FragmentStageV1::Transcript),
speaker_labels: stage(FragmentStageV1::SpeakerLabels),
speaker_features: stage(FragmentStageV1::SpeakerFeatures),
structuring: stage(FragmentStageV1::Structuring),
label_confirmation: stage(FragmentStageV1::LabelConfirmation),
attempt_count: 0,
jobs: Vec::new(),
interim_txid: None,
analysis: None,
confirmed_labels: Vec::new(),
final_transcript: transcript.map(str::to_owned),
errors: Vec::new(),
errors_truncated: false,
}
}
fn object_with_data(data: Vec<u8>) -> Object {
Object {
filename: String::new(),
file_type: MANIFEST_FILE_TYPE.to_owned(),
description: String::new(),
data,
}
}
fn status_result(statuses: &[FragmentStatus]) -> Result<FullAudioStatus, String> {
let entries = statuses
.iter()
.enumerate()
.map(|(index, _)| {
let value = u8::try_from(index + 1).expect("test fragment ID");
let start = u64::try_from(index).expect("test index") * 80;
entry(value, start, start + 100)
})
.collect::<Vec<_>>();
let object = object_with_data(encode_manifest(&entries).expect("manifest"));
status_orchestration(
id(90),
|_| Ok(Some(object.clone())),
|fragment_id| {
Ok(entries
.iter()
.position(|entry| entry.fragment_id == fragment_id)
.map(|index| statuses[index].clone()))
},
)
}
#[test]
fn submission_contract_is_preserved() {
let events = RefCell::new(Vec::new());
let returned = submit_orchestration(
b"input",
|audio| {
assert_eq!(audio, b"input");
events.borrow_mut().push("convert");
Ok(b"ogg".to_vec())
},
|ogg| {
assert_eq!(ogg, b"ogg");
events.borrow_mut().push("split");
Ok(vec![pending(10, 0, 100), pending(20, 80, 180)])
},
|fragment| {
let (event, id) = if fragment[0] == 10 {
("submit-10", id(1))
} else {
("submit-20", id(2))
};
events.borrow_mut().push(event);
Ok(id)
},
|manifest| {
events.borrow_mut().push("save");
assert_eq!(
decode_manifest(manifest).expect("saved manifest"),
vec![entry(1, 0, 100), entry(2, 80, 180)]
);
Ok(id(9))
},
)
.expect("submit");
assert_eq!(returned, id(9));
assert_eq!(
events.borrow().as_slice(),
["convert", "split", "submit-10", "submit-20", "save"]
);
let partial_submissions = Cell::new(0);
let partial_saves = Cell::new(0);
let partial = submit_orchestration(
b"input",
|_| Ok(vec![1]),
|_| Ok(vec![pending(1, 0, 100), pending(2, 80, 180)]),
|_| {
partial_submissions.set(partial_submissions.get() + 1);
if partial_submissions.get() == 2 {
Err("second submission failed".to_owned())
} else {
Ok(id(1))
}
},
|_| {
partial_saves.set(partial_saves.get() + 1);
Ok(id(9))
},
);
assert!(partial.is_err());
assert_eq!((partial_submissions.get(), partial_saves.get()), (2, 0));
let overlong_submissions = Cell::new(0);
let overlong_saves = Cell::new(0);
let overlong = submit_orchestration(
b"input",
|_| Ok(vec![1]),
|_| Ok(vec![pending(1, 0, 7_200_001)]),
|_| {
overlong_submissions.set(overlong_submissions.get() + 1);
Ok(id(1))
},
|_| {
overlong_saves.set(overlong_saves.get() + 1);
Ok(id(9))
},
)
.unwrap_err();
assert_eq!(overlong, "invalid fragment interval");
assert_eq!((overlong_submissions.get(), overlong_saves.get()), (0, 0));
}
#[test]
fn status_contract_is_preserved() {
for (first, second, transcript, expected) in [
(Failed, Queued, None, NeedsAttention),
(Completed, Running, None, Processing),
(Completed, Confirmed, Some("ready"), AwaitingLabels),
] {
let status = status_result(&[
classification_status(first, None),
classification_status(second, transcript),
])
.expect("derived state");
assert_eq!((status.state, status.final_transcript), (expected, None));
}
let complete = status_result(&[
classification_status(Confirmed, Some("[high] A: one two three four")),
classification_status(Confirmed, Some("[medium] B: two three four five")),
])
.expect("complete");
assert_eq!(complete.state, FullAudioState::Complete);
assert_eq!(
complete.final_transcript.as_deref(),
Some("[high] A: one two three four\n[medium] B: five")
);
assert_eq!(complete.fragments[0].fragment_id, id(1));
assert_eq!(complete.fragments[1].start_sample_48k, 80);
for status in [
classification_status(Confirmed, None),
classification_status(Completed, Some("not confirmed")),
] {
assert!(status_result(&[status]).is_err());
}
let entries = [entry(1, 0, 100)];
let valid = object_with_data(encode_manifest(&entries).expect("manifest"));
let mut wrong_type = valid.clone();
wrong_type.file_type = "wrong".to_owned();
let mut wrong_filename = valid.clone();
wrong_filename.filename = "named".to_owned();
for object in [
None,
Some(valid),
Some(wrong_type),
Some(wrong_filename),
Some(object_with_data(vec![1])),
] {
assert!(status_orchestration(id(90), |_| Ok(object.clone()), |_| Ok(None)).is_err());
}
let statuses = [
classification_status(Confirmed, Some("Alpha, beta gamma delta")),
classification_status(Confirmed, Some("BETA GAMMA DELTA epsilon")),
];
let first = status_result(&statuses).expect("first status");
assert_eq!(first, status_result(&statuses).expect("second status"));
assert_eq!(
first.final_transcript.as_deref(),
Some("Alpha, beta gamma delta\nepsilon")
);
}
#[test]
fn independent_calls_and_path_validation_are_preserved() {
let (entered_sender, entered_receiver) = mpsc::sync_channel(1);
let (release_sender, release_receiver) = mpsc::channel();
let blocked = thread::spawn(move || {
submit_orchestration(
b"input",
|_| Ok(vec![1]),
|_| Ok(vec![pending(1, 0, 100)]),
|_| {
entered_sender.send(()).expect("signal blocked submit");
release_receiver.recv().expect("release blocked submit");
Ok(id(1))
},
|_| Ok(id(10)),
)
});
entered_receiver
.recv_timeout(Duration::from_secs(2))
.expect("blocked operation entered");
let independent = submit_orchestration(
b"input",
|_| Ok(vec![1]),
|_| Ok(vec![pending(2, 0, 100)]),
|_| Ok(id(2)),
|_| Ok(id(20)),
);
assert_eq!(independent.expect("independent submit"), id(20));
release_sender.send(()).expect("release blocked operation");
assert_eq!(
blocked.join().expect("join").expect("blocked submit"),
id(10)
);
assert_eq!(
validate_ffmpeg_path(Path::new("/definitely/not/present")).expect("absolute path"),
PathBuf::from("/definitely/not/present")
);
assert!(validate_ffmpeg_path(Path::new("relative/ffmpeg")).is_err());
}
}