Skip to main content

kcode_k1_full_audio/
lib.rs

1use std::path::{Path, PathBuf};
2use std::sync::Arc;
3
4use kcode_k1_audio_classification::{
5    AudioClassification, FragmentId, FragmentStatus, OverallState,
6};
7use kcode_k1_full_audio_domain::{
8    ManifestEntry, decode_manifest, encode_manifest, stitch_transcripts, validate_fragment_geometry,
9};
10use kcode_k1_objects::{K1Objects, Object};
11
12const MANIFEST_FILE_TYPE: &str = "k1-full-audio-manifest-v1";
13
14pub type FullAudioId = kcode_k1_objects::TxId;
15
16pub struct K1FullAudio {
17    ffmpeg_path: PathBuf,
18    objects: Arc<K1Objects>,
19    classification: Arc<AudioClassification>,
20}
21
22#[derive(Clone, Debug, Eq, PartialEq)]
23pub enum FullAudioState {
24    Processing,
25    AwaitingLabels,
26    NeedsAttention,
27    Complete,
28}
29
30#[derive(Clone, Debug, PartialEq)]
31pub struct FullAudioFragmentStatus {
32    pub fragment_id: FragmentId,
33    pub start_sample_48k: u64,
34    pub end_sample_48k: u64,
35    pub status: FragmentStatus,
36}
37
38#[derive(Clone, Debug, PartialEq)]
39pub struct FullAudioStatus {
40    pub state: FullAudioState,
41    pub fragments: Vec<FullAudioFragmentStatus>,
42    pub final_transcript: Option<String>,
43}
44
45struct PendingFragment {
46    start_sample_48k: u64,
47    end_sample_48k: u64,
48    ogg_bytes: Vec<u8>,
49}
50
51impl K1FullAudio {
52    pub fn open(
53        ffmpeg_path: impl AsRef<Path>,
54        objects: Arc<K1Objects>,
55        classification: Arc<AudioClassification>,
56    ) -> Result<Self, String> {
57        Ok(Self {
58            ffmpeg_path: validate_ffmpeg_path(ffmpeg_path.as_ref())?,
59            objects,
60            classification,
61        })
62    }
63
64    pub fn submit(&self, audio: &[u8]) -> Result<FullAudioId, String> {
65        submit_orchestration(
66            audio,
67            |input| kcode_audio_to_ogg_opus::convert_to_ogg_opus(&self.ffmpeg_path, input),
68            |ogg_bytes| {
69                kcode_ogg_opus_fragments::split_ogg_opus(ogg_bytes).map(|fragments| {
70                    fragments
71                        .into_iter()
72                        .map(|fragment| PendingFragment {
73                            start_sample_48k: fragment.start_sample_48k,
74                            end_sample_48k: fragment.end_sample_48k,
75                            ogg_bytes: fragment.ogg_bytes,
76                        })
77                        .collect()
78                })
79            },
80            |ogg_bytes| self.classification.submit(ogg_bytes),
81            |manifest| self.objects.save("", MANIFEST_FILE_TYPE, "", manifest),
82        )
83    }
84
85    pub fn status(&self, id: FullAudioId) -> Result<FullAudioStatus, String> {
86        status_orchestration(
87            id,
88            |manifest_id| self.objects.load(manifest_id),
89            |fragment_id| self.classification.status(fragment_id),
90        )
91    }
92}
93
94fn validate_ffmpeg_path(path: &Path) -> Result<PathBuf, String> {
95    if !path.is_absolute() {
96        return Err("ffmpeg path must be absolute".to_owned());
97    }
98    Ok(path.to_path_buf())
99}
100
101fn submit_orchestration<Convert, Split, Submit, Save>(
102    audio: &[u8],
103    convert: Convert,
104    split: Split,
105    mut submit: Submit,
106    save: Save,
107) -> Result<FullAudioId, String>
108where
109    Convert: FnOnce(&[u8]) -> Result<Vec<u8>, String>,
110    Split: FnOnce(&[u8]) -> Result<Vec<PendingFragment>, String>,
111    Submit: FnMut(&[u8]) -> Result<FragmentId, String>,
112    Save: FnOnce(&[u8]) -> Result<FullAudioId, String>,
113{
114    let ogg_bytes = convert(audio)?;
115    let fragments = split(&ogg_bytes)?;
116    validate_fragment_geometry(
117        fragments
118            .iter()
119            .map(|fragment| (fragment.start_sample_48k, fragment.end_sample_48k)),
120    )?;
121    let mut entries = Vec::new();
122    entries
123        .try_reserve_exact(fragments.len())
124        .map_err(|error| format!("allocate manifest entries: {error}"))?;
125    for fragment in fragments {
126        let fragment_id = submit(&fragment.ogg_bytes)?;
127        entries.push(ManifestEntry {
128            fragment_id,
129            start_sample_48k: fragment.start_sample_48k,
130            end_sample_48k: fragment.end_sample_48k,
131        });
132    }
133    let manifest = encode_manifest(&entries)?;
134    save(&manifest)
135}
136
137fn status_orchestration<Load, Status>(
138    id: FullAudioId,
139    mut load: Load,
140    mut status: Status,
141) -> Result<FullAudioStatus, String>
142where
143    Load: FnMut(FullAudioId) -> Result<Option<Object>, String>,
144    Status: FnMut(FragmentId) -> Result<Option<FragmentStatus>, String>,
145{
146    let object = load(id)?.ok_or_else(|| "unknown full-audio ID".to_owned())?;
147    if object.file_type != MANIFEST_FILE_TYPE
148        || !object.filename.is_empty()
149        || !object.description.is_empty()
150    {
151        return Err("object is not a full-audio manifest".to_owned());
152    }
153    let entries = decode_manifest(&object.data)?;
154    let mut fragments = Vec::new();
155    fragments
156        .try_reserve_exact(entries.len())
157        .map_err(|error| format!("allocate fragment statuses: {error}"))?;
158    for entry in entries {
159        let fragment_status = status(entry.fragment_id)?
160            .ok_or_else(|| format!("unknown audio fragment {}", entry.fragment_id))?;
161        fragments.push(FullAudioFragmentStatus {
162            fragment_id: entry.fragment_id,
163            start_sample_48k: entry.start_sample_48k,
164            end_sample_48k: entry.end_sample_48k,
165            status: fragment_status,
166        });
167    }
168    let state = derive_state(&fragments)?;
169    let final_transcript = if state == FullAudioState::Complete {
170        let transcripts = fragments
171            .iter()
172            .filter_map(|fragment| fragment.status.final_transcript.as_deref())
173            .collect::<Vec<_>>();
174        Some(stitch_transcripts(&transcripts))
175    } else {
176        None
177    };
178    Ok(FullAudioStatus {
179        state,
180        fragments,
181        final_transcript,
182    })
183}
184
185fn derive_state(fragments: &[FullAudioFragmentStatus]) -> Result<FullAudioState, String> {
186    if fragments.is_empty() {
187        return Err("manifest has no fragments".to_owned());
188    }
189    for fragment in fragments {
190        let confirmed = fragment.status.state == OverallState::Confirmed;
191        let has_transcript = fragment.status.final_transcript.is_some();
192        if confirmed && !has_transcript {
193            return Err("confirmed fragment has no final transcript".to_owned());
194        }
195        if !confirmed && has_transcript {
196            return Err("non-confirmed fragment has a final transcript".to_owned());
197        }
198    }
199    if has_state(fragments, &[OverallState::Failed, OverallState::Discarded]) {
200        return Ok(FullAudioState::NeedsAttention);
201    }
202    if has_state(fragments, &[OverallState::Queued, OverallState::Running]) {
203        return Ok(FullAudioState::Processing);
204    }
205    if fragments
206        .iter()
207        .all(|fragment| fragment.status.state == OverallState::Confirmed)
208    {
209        return Ok(FullAudioState::Complete);
210    }
211    if fragments.iter().all(|fragment| {
212        [OverallState::Completed, OverallState::Confirmed].contains(&fragment.status.state)
213    }) {
214        return Ok(FullAudioState::AwaitingLabels);
215    }
216    Err("classification returned an inconsistent state".to_owned())
217}
218
219fn has_state(fragments: &[FullAudioFragmentStatus], states: &[OverallState]) -> bool {
220    fragments
221        .iter()
222        .any(|fragment| states.contains(&fragment.status.state))
223}
224
225#[cfg(test)]
226mod tests;