kcode-k1-audio-classification-projection-state 0.2.0

Persisted status types and validation for K1 audio classification
Documentation
pub use kcode_k1_audio_classification_format::{
    ExecutedAnalysis, FragmentStageV1, PersonId, SpeakerLabelV1, TxId,
};
use serde::{Deserialize, Serialize};

pub const MAX_ERRORS: usize = 5_000;
pub type FragmentId = TxId;

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum OverallState {
    Queued,
    Running,
    Failed,
    Completed,
    Confirmed,
    Discarded,
}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum StageState {
    Pending,
    Running,
    Succeeded,
    Failed,
}

#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum LlmJobState {
    Running,
    Succeeded,
    Failed,
}

#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct StageStatus {
    pub stage: FragmentStageV1,
    pub state: StageState,
}

#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct LlmJobStatus {
    pub attempt: u32,
    pub sequence: u64,
    pub stage: FragmentStageV1,
    pub name: String,
    pub state: LlmJobState,
}

#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct FragmentStatus {
    pub state: OverallState,
    pub queue: StageStatus,
    pub transcript: StageStatus,
    pub speaker_labels: StageStatus,
    pub speaker_features: StageStatus,
    pub structuring: StageStatus,
    pub label_confirmation: StageStatus,
    pub attempt_count: u32,
    pub jobs: Vec<LlmJobStatus>,
    #[serde(with = "optional_fragment_id")]
    pub interim_txid: Option<FragmentId>,
    pub analysis: Option<ExecutedAnalysis>,
    pub confirmed_labels: Vec<SpeakerLabelV1>,
    pub final_transcript: Option<String>,
    pub errors: Vec<String>,
    pub errors_truncated: bool,
}

pub fn append_error(status: &mut FragmentStatus, error: String) {
    if status.errors.len() < MAX_ERRORS {
        status.errors.push(error);
    } else {
        status.errors_truncated = true;
    }
}

pub fn validate_labels(
    status: &FragmentStatus,
    labels: &[SpeakerLabelV1],
) -> Result<FragmentId, String> {
    final_transcript(status, labels).map(|(interim, _)| interim)
}

pub fn final_transcript(
    status: &FragmentStatus,
    labels: &[SpeakerLabelV1],
) -> Result<(FragmentId, String), String> {
    if status.state != OverallState::Completed {
        return Err("label confirmation requires Completed".to_string());
    }
    let interim = status
        .interim_txid
        .ok_or_else(|| "completed status has no interim transaction ID".to_string())?;
    let analysis = status
        .analysis
        .as_ref()
        .ok_or_else(|| "completed status has no analysis".to_string())?;
    if labels.len() != analysis.envelope.analysis.speakers.len() {
        return Err("speaker labels are not one-to-one".to_string());
    }
    for (label, expected) in labels.iter().zip(&analysis.envelope.analysis.speakers) {
        if label.speaker != expected.speaker {
            return Err("speaker labels are not in exact analysis order".to_string());
        }
    }
    Ok((
        interim,
        replace_transcript(&analysis.envelope.analysis.transcript, labels),
    ))
}

pub fn actionable_state(state: OverallState) -> i64 {
    match state {
        OverallState::Queued => 1,
        OverallState::Running => 2,
        _ => 0,
    }
}

pub fn interrupted_stage(status: &FragmentStatus) -> FragmentStageV1 {
    for stage in [
        &status.structuring,
        &status.speaker_features,
        &status.speaker_labels,
        &status.transcript,
    ] {
        if stage.state == StageState::Running {
            return stage.stage;
        }
    }
    FragmentStageV1::Queue
}

pub fn validate_status(
    status: &FragmentStatus,
    stored_actionable_state: i64,
) -> Result<(), String> {
    let identities = [
        (&status.queue, FragmentStageV1::Queue),
        (&status.transcript, FragmentStageV1::Transcript),
        (&status.speaker_labels, FragmentStageV1::SpeakerLabels),
        (&status.speaker_features, FragmentStageV1::SpeakerFeatures),
        (&status.structuring, FragmentStageV1::Structuring),
        (
            &status.label_confirmation,
            FragmentStageV1::LabelConfirmation,
        ),
    ];
    if identities
        .iter()
        .any(|(stored, expected)| stored.stage != *expected)
    {
        return Err("stored stage identity does not match its field".to_string());
    }
    if actionable_state(status.state) != stored_actionable_state {
        return Err("stored actionable state does not match status".to_string());
    }
    if status.errors.len() > MAX_ERRORS {
        return Err("stored errors exceed the retention bound".to_string());
    }
    let mut previous = None;
    for job in &status.jobs {
        if job.attempt == 0
            || job.attempt > status.attempt_count
            || !is_analysis_stage(&job.stage)
            || job.name.trim().is_empty()
            || previous.is_some_and(|value| value >= (job.attempt, job.sequence))
        {
            return Err("stored LLM jobs are invalid or unordered".to_string());
        }
        previous = Some((job.attempt, job.sequence));
    }
    if !matches!(
        status.queue.state,
        StageState::Succeeded | StageState::Failed
    ) {
        return Err("stored Queue stage is neither succeeded nor failed".to_string());
    }
    if matches!(
        status.state,
        OverallState::Completed | OverallState::Confirmed
    ) && (status.interim_txid.is_none() || status.analysis.is_none())
    {
        return Err("stored completed status lacks its analysis".to_string());
    }
    if status.state == OverallState::Confirmed
        && (status.final_transcript.is_none()
            || status.label_confirmation.state != StageState::Succeeded)
    {
        return Err("stored confirmed status is incomplete".to_string());
    }
    Ok(())
}

fn replace_transcript(transcript: &str, labels: &[SpeakerLabelV1]) -> String {
    let mut output = String::with_capacity(transcript.len());
    for line in transcript.split_inclusive('\n') {
        let mut replaced = false;
        for prefix in ["[high] ", "[medium] ", "[low] "] {
            if let Some(rest) = line.strip_prefix(prefix) {
                for label in labels {
                    let speaker = label.speaker.to_string();
                    if let Some(tail) = rest.strip_prefix(&speaker)
                        && (tail.starts_with(':') || tail.starts_with(" [overlap]:"))
                    {
                        output.push_str(prefix);
                        match label.person_id {
                            Some(person_id) => output.push_str(&person_id.to_string()),
                            None => output.push_str("Unknown"),
                        }
                        output.push_str(tail);
                        replaced = true;
                        break;
                    }
                }
            }
            if replaced {
                break;
            }
        }
        if !replaced {
            output.push_str(line);
        }
    }
    output
}

fn is_analysis_stage(stage: &FragmentStageV1) -> bool {
    matches!(
        stage,
        FragmentStageV1::Transcript
            | FragmentStageV1::SpeakerLabels
            | FragmentStageV1::SpeakerFeatures
            | FragmentStageV1::Structuring
    )
}

mod optional_fragment_id {
    use super::FragmentId;
    use serde::{Deserialize, Deserializer, Serialize, Serializer};

    pub fn serialize<S: Serializer>(
        value: &Option<FragmentId>,
        serializer: S,
    ) -> Result<S::Ok, S::Error> {
        value.map(FragmentId::into_bytes).serialize(serializer)
    }

    pub fn deserialize<'de, D: Deserializer<'de>>(
        deserializer: D,
    ) -> Result<Option<FragmentId>, D::Error> {
        Option::<[u8; 12]>::deserialize(deserializer).map(|value| value.map(FragmentId::from_bytes))
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use kcode_speaker_v3_analysis::{
        AnalysisEnvelope, FeatureVector24, GeminiCohort, LocalSpeakerLabel, OggAudioMetadata,
        StructuredAnalysis, StructuredSpeaker, StructurerProvenance,
    };

    fn id(value: u8) -> FragmentId {
        FragmentId::from_bytes([value; 12])
    }

    fn person(value: u8) -> PersonId {
        PersonId::from_tx_id(id(value))
    }

    fn stage(stage: FragmentStageV1, state: StageState) -> StageStatus {
        StageStatus { stage, state }
    }

    fn status() -> FragmentStatus {
        FragmentStatus {
            state: OverallState::Queued,
            queue: stage(FragmentStageV1::Queue, StageState::Succeeded),
            transcript: stage(FragmentStageV1::Transcript, StageState::Pending),
            speaker_labels: stage(FragmentStageV1::SpeakerLabels, StageState::Pending),
            speaker_features: stage(FragmentStageV1::SpeakerFeatures, StageState::Pending),
            structuring: stage(FragmentStageV1::Structuring, StageState::Pending),
            label_confirmation: stage(FragmentStageV1::LabelConfirmation, StageState::Pending),
            attempt_count: 0,
            jobs: Vec::new(),
            interim_txid: None,
            analysis: None,
            confirmed_labels: Vec::new(),
            final_transcript: None,
            errors: Vec::new(),
            errors_truncated: false,
        }
    }

    fn analysis(transcript: &str) -> ExecutedAnalysis {
        let mut ogg = vec![0; 29];
        ogg[..4].copy_from_slice(b"OggS");
        ogg[26] = 1;
        ogg[27] = 1;
        let speaker = |number| StructuredSpeaker {
            speaker: LocalSpeakerLabel::new(number).unwrap(),
            language: "en".into(),
            features: FeatureVector24::default(),
            features_usable_for_training: true,
        };
        let provenance = StructurerProvenance {
            model_id: "model".into(),
            prompt_revision: "prompt".into(),
        };
        ExecutedAnalysis {
            envelope: AnalysisEnvelope {
                audio: OggAudioMetadata::from_bytes(&ogg, 1, None).unwrap(),
                analysis: StructuredAnalysis {
                    transcript: transcript.into(),
                    speakers: vec![speaker(1), speaker(2)],
                },
                gemini: GeminiCohort {
                    model_id: "gemini".into(),
                    transcript_prompt_revision: "t".into(),
                    feature_prompt_revisions: ["1".into(), "2".into(), "3".into()],
                    feature_schema_revision: "s".into(),
                },
                structurer: provenance.clone(),
            },
            label_extractor: provenance,
        }
    }

    fn label(number: u32, person_id: Option<PersonId>) -> SpeakerLabelV1 {
        SpeakerLabelV1 {
            speaker: LocalSpeakerLabel::new(number).unwrap(),
            person_id,
        }
    }

    #[test]
    fn validation_checks_persisted_invariants() {
        let mut value = status();
        assert_eq!(actionable_state(value.state), 1);
        assert_eq!(interrupted_stage(&value), FragmentStageV1::Queue);
        assert!(validate_status(&value, 1).is_ok());
        value.transcript.state = StageState::Running;
        value.structuring.state = StageState::Running;
        assert_eq!(interrupted_stage(&value), FragmentStageV1::Structuring);
        value.transcript.stage = FragmentStageV1::Structuring;
        assert!(validate_status(&value, 1).is_err());
    }

    #[test]
    fn error_retention_keeps_the_oldest_bound() {
        let mut value = status();
        for index in 0..=MAX_ERRORS {
            append_error(&mut value, index.to_string());
        }
        assert_eq!(value.errors.len(), MAX_ERRORS);
        assert_eq!(value.errors.first().map(String::as_str), Some("0"));
        assert_eq!(value.errors.last().map(String::as_str), Some("4999"));
        assert!(value.errors_truncated);
        assert!(validate_status(&value, 1).is_ok());
    }

    #[test]
    fn label_validation_and_transcript_derivation_are_exact() {
        let mut value = status();
        value.state = OverallState::Completed;
        value.interim_txid = Some(id(7));
        value.analysis = Some(analysis(
            "[high] Speaker 1: hi\n[medium] Speaker 2 [overlap]: yo\nplain Speaker 1: no\n[low] Speaker 1 [overlap]: end",
        ));
        let known = person(0xab);
        let labels = vec![label(1, Some(known)), label(2, None)];
        assert_eq!(validate_labels(&value, &labels), Ok(id(7)));
        assert_eq!(
            final_transcript(&value, &labels),
            Ok((
                id(7),
                "[high] abababababababababababab: hi\n[medium] Unknown [overlap]: yo\nplain Speaker 1: no\n[low] abababababababababababab [overlap]: end".into()
            ))
        );
        assert!(final_transcript(&value, &[label(2, None), label(1, Some(known))]).is_err());
        assert!(final_transcript(&value, &[label(1, Some(known))]).is_err());
    }
}