kcode-k1-full-audio 0.3.2

Durable orchestration of full audio into classified overlapping K1 fragments
Documentation
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;