Skip to main content

kcode_k1_audio_classification_projection_state/
lib.rs

1pub use kcode_k1_audio_classification_format::{
2    ExecutedAnalysis, FragmentStageV1, SpeakerLabelV1, TxId,
3};
4use serde::{Deserialize, Serialize};
5
6pub const MAX_ERRORS: usize = 5_000;
7pub type FragmentId = TxId;
8
9#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
10pub enum OverallState {
11    Queued,
12    Running,
13    Failed,
14    Completed,
15    Confirmed,
16    Discarded,
17}
18
19#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
20pub enum StageState {
21    Pending,
22    Running,
23    Succeeded,
24    Failed,
25}
26
27#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
28pub enum LlmJobState {
29    Running,
30    Succeeded,
31    Failed,
32}
33
34#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
35pub struct StageStatus {
36    pub stage: FragmentStageV1,
37    pub state: StageState,
38}
39
40#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
41pub struct LlmJobStatus {
42    pub attempt: u32,
43    pub sequence: u64,
44    pub stage: FragmentStageV1,
45    pub name: String,
46    pub state: LlmJobState,
47}
48
49#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
50pub struct FragmentStatus {
51    pub state: OverallState,
52    pub queue: StageStatus,
53    pub transcript: StageStatus,
54    pub speaker_labels: StageStatus,
55    pub speaker_features: StageStatus,
56    pub structuring: StageStatus,
57    pub label_confirmation: StageStatus,
58    pub attempt_count: u32,
59    pub jobs: Vec<LlmJobStatus>,
60    #[serde(with = "optional_fragment_id")]
61    pub interim_txid: Option<FragmentId>,
62    pub analysis: Option<ExecutedAnalysis>,
63    pub confirmed_labels: Vec<SpeakerLabelV1>,
64    pub final_transcript: Option<String>,
65    pub errors: Vec<String>,
66    pub errors_truncated: bool,
67}
68
69pub fn append_error(status: &mut FragmentStatus, error: String) {
70    if status.errors.len() < MAX_ERRORS {
71        status.errors.push(error);
72    } else {
73        status.errors_truncated = true;
74    }
75}
76
77pub fn validate_labels(
78    status: &FragmentStatus,
79    labels: &[SpeakerLabelV1],
80) -> Result<FragmentId, String> {
81    final_transcript(status, labels).map(|(interim, _)| interim)
82}
83
84pub fn final_transcript(
85    status: &FragmentStatus,
86    labels: &[SpeakerLabelV1],
87) -> Result<(FragmentId, String), String> {
88    if status.state != OverallState::Completed {
89        return Err("label confirmation requires Completed".to_string());
90    }
91    let interim = status
92        .interim_txid
93        .ok_or_else(|| "completed status has no interim transaction ID".to_string())?;
94    let analysis = status
95        .analysis
96        .as_ref()
97        .ok_or_else(|| "completed status has no analysis".to_string())?;
98    if labels.len() != analysis.envelope.analysis.speakers.len() {
99        return Err("speaker labels are not one-to-one".to_string());
100    }
101    for (label, expected) in labels.iter().zip(&analysis.envelope.analysis.speakers) {
102        if label.speaker != expected.speaker {
103            return Err("speaker labels are not in exact analysis order".to_string());
104        }
105        if invalid_person_id(&label.person_id) {
106            return Err("person ID is blank or contains a line break".to_string());
107        }
108    }
109    Ok((
110        interim,
111        replace_transcript(&analysis.envelope.analysis.transcript, labels),
112    ))
113}
114
115pub fn actionable_state(state: OverallState) -> i64 {
116    match state {
117        OverallState::Queued => 1,
118        OverallState::Running => 2,
119        _ => 0,
120    }
121}
122
123pub fn interrupted_stage(status: &FragmentStatus) -> FragmentStageV1 {
124    for stage in [
125        &status.structuring,
126        &status.speaker_features,
127        &status.speaker_labels,
128        &status.transcript,
129    ] {
130        if stage.state == StageState::Running {
131            return stage.stage;
132        }
133    }
134    FragmentStageV1::Queue
135}
136
137pub fn validate_status(
138    status: &FragmentStatus,
139    stored_actionable_state: i64,
140) -> Result<(), String> {
141    let identities = [
142        (&status.queue, FragmentStageV1::Queue),
143        (&status.transcript, FragmentStageV1::Transcript),
144        (&status.speaker_labels, FragmentStageV1::SpeakerLabels),
145        (&status.speaker_features, FragmentStageV1::SpeakerFeatures),
146        (&status.structuring, FragmentStageV1::Structuring),
147        (
148            &status.label_confirmation,
149            FragmentStageV1::LabelConfirmation,
150        ),
151    ];
152    if identities
153        .iter()
154        .any(|(stored, expected)| stored.stage != *expected)
155    {
156        return Err("stored stage identity does not match its field".to_string());
157    }
158    if actionable_state(status.state) != stored_actionable_state {
159        return Err("stored actionable state does not match status".to_string());
160    }
161    if status.errors.len() > MAX_ERRORS {
162        return Err("stored errors exceed the retention bound".to_string());
163    }
164    let mut previous = None;
165    for job in &status.jobs {
166        if job.attempt == 0
167            || job.attempt > status.attempt_count
168            || !is_analysis_stage(&job.stage)
169            || job.name.trim().is_empty()
170            || previous.is_some_and(|value| value >= (job.attempt, job.sequence))
171        {
172            return Err("stored LLM jobs are invalid or unordered".to_string());
173        }
174        previous = Some((job.attempt, job.sequence));
175    }
176    if !matches!(
177        status.queue.state,
178        StageState::Succeeded | StageState::Failed
179    ) {
180        return Err("stored Queue stage is neither succeeded nor failed".to_string());
181    }
182    if matches!(
183        status.state,
184        OverallState::Completed | OverallState::Confirmed
185    ) && (status.interim_txid.is_none() || status.analysis.is_none())
186    {
187        return Err("stored completed status lacks its analysis".to_string());
188    }
189    if status.state == OverallState::Confirmed
190        && (status.final_transcript.is_none()
191            || status.label_confirmation.state != StageState::Succeeded)
192    {
193        return Err("stored confirmed status is incomplete".to_string());
194    }
195    if status
196        .confirmed_labels
197        .iter()
198        .any(|label| invalid_person_id(&label.person_id))
199    {
200        return Err("stored person ID is invalid".to_string());
201    }
202    Ok(())
203}
204
205fn replace_transcript(transcript: &str, labels: &[SpeakerLabelV1]) -> String {
206    let mut output = String::with_capacity(transcript.len());
207    for line in transcript.split_inclusive('\n') {
208        let mut replaced = false;
209        for prefix in ["[high] ", "[medium] ", "[low] "] {
210            if let Some(rest) = line.strip_prefix(prefix) {
211                for label in labels {
212                    let speaker = label.speaker.to_string();
213                    if let Some(tail) = rest.strip_prefix(&speaker)
214                        && (tail.starts_with(':') || tail.starts_with(" [overlap]:"))
215                    {
216                        output.push_str(prefix);
217                        output.push_str(&label.person_id);
218                        output.push_str(tail);
219                        replaced = true;
220                        break;
221                    }
222                }
223            }
224            if replaced {
225                break;
226            }
227        }
228        if !replaced {
229            output.push_str(line);
230        }
231    }
232    output
233}
234
235fn is_analysis_stage(stage: &FragmentStageV1) -> bool {
236    matches!(
237        stage,
238        FragmentStageV1::Transcript
239            | FragmentStageV1::SpeakerLabels
240            | FragmentStageV1::SpeakerFeatures
241            | FragmentStageV1::Structuring
242    )
243}
244
245fn invalid_person_id(value: &str) -> bool {
246    value.trim().is_empty() || value.contains('\r') || value.contains('\n')
247}
248
249mod optional_fragment_id {
250    use super::FragmentId;
251    use serde::{Deserialize, Deserializer, Serialize, Serializer};
252
253    pub fn serialize<S: Serializer>(
254        value: &Option<FragmentId>,
255        serializer: S,
256    ) -> Result<S::Ok, S::Error> {
257        value.map(FragmentId::into_bytes).serialize(serializer)
258    }
259
260    pub fn deserialize<'de, D: Deserializer<'de>>(
261        deserializer: D,
262    ) -> Result<Option<FragmentId>, D::Error> {
263        Option::<[u8; 12]>::deserialize(deserializer).map(|value| value.map(FragmentId::from_bytes))
264    }
265}
266
267#[cfg(test)]
268mod tests {
269    use super::*;
270    use kcode_speaker_v3_analysis::{
271        AnalysisEnvelope, FeatureVector24, GeminiCohort, LocalSpeakerLabel, OggAudioMetadata,
272        StructuredAnalysis, StructuredSpeaker, StructurerProvenance,
273    };
274
275    fn id(value: u8) -> FragmentId {
276        FragmentId::from_bytes([value; 12])
277    }
278
279    fn stage(stage: FragmentStageV1, state: StageState) -> StageStatus {
280        StageStatus { stage, state }
281    }
282
283    fn status() -> FragmentStatus {
284        FragmentStatus {
285            state: OverallState::Queued,
286            queue: stage(FragmentStageV1::Queue, StageState::Succeeded),
287            transcript: stage(FragmentStageV1::Transcript, StageState::Pending),
288            speaker_labels: stage(FragmentStageV1::SpeakerLabels, StageState::Pending),
289            speaker_features: stage(FragmentStageV1::SpeakerFeatures, StageState::Pending),
290            structuring: stage(FragmentStageV1::Structuring, StageState::Pending),
291            label_confirmation: stage(FragmentStageV1::LabelConfirmation, StageState::Pending),
292            attempt_count: 0,
293            jobs: Vec::new(),
294            interim_txid: None,
295            analysis: None,
296            confirmed_labels: Vec::new(),
297            final_transcript: None,
298            errors: Vec::new(),
299            errors_truncated: false,
300        }
301    }
302
303    fn analysis(transcript: &str) -> ExecutedAnalysis {
304        let mut ogg = vec![0; 29];
305        ogg[..4].copy_from_slice(b"OggS");
306        ogg[26] = 1;
307        ogg[27] = 1;
308        let speaker = |number| StructuredSpeaker {
309            speaker: LocalSpeakerLabel::new(number).unwrap(),
310            language: "en".into(),
311            features: FeatureVector24::default(),
312            features_usable_for_training: true,
313        };
314        let provenance = StructurerProvenance {
315            model_id: "model".into(),
316            prompt_revision: "prompt".into(),
317        };
318        ExecutedAnalysis {
319            envelope: AnalysisEnvelope {
320                audio: OggAudioMetadata::from_bytes(&ogg, 1, None).unwrap(),
321                analysis: StructuredAnalysis {
322                    transcript: transcript.into(),
323                    speakers: vec![speaker(1), speaker(2)],
324                },
325                gemini: GeminiCohort {
326                    model_id: "gemini".into(),
327                    transcript_prompt_revision: "t".into(),
328                    feature_prompt_revisions: ["1".into(), "2".into(), "3".into()],
329                    feature_schema_revision: "s".into(),
330                },
331                structurer: provenance.clone(),
332            },
333            label_extractor: provenance,
334        }
335    }
336
337    fn label(number: u32, person_id: &str) -> SpeakerLabelV1 {
338        SpeakerLabelV1 {
339            speaker: LocalSpeakerLabel::new(number).unwrap(),
340            person_id: person_id.into(),
341        }
342    }
343
344    #[test]
345    fn validation_checks_persisted_invariants() {
346        let mut value = status();
347        assert_eq!(actionable_state(value.state), 1);
348        assert_eq!(interrupted_stage(&value), FragmentStageV1::Queue);
349        assert!(validate_status(&value, 1).is_ok());
350        value.transcript.state = StageState::Running;
351        value.structuring.state = StageState::Running;
352        assert_eq!(interrupted_stage(&value), FragmentStageV1::Structuring);
353        value.transcript.stage = FragmentStageV1::Structuring;
354        assert!(validate_status(&value, 1).is_err());
355    }
356
357    #[test]
358    fn error_retention_keeps_the_oldest_bound() {
359        let mut value = status();
360        for index in 0..=MAX_ERRORS {
361            append_error(&mut value, index.to_string());
362        }
363        assert_eq!(value.errors.len(), MAX_ERRORS);
364        assert_eq!(value.errors.first().map(String::as_str), Some("0"));
365        assert_eq!(value.errors.last().map(String::as_str), Some("4999"));
366        assert!(value.errors_truncated);
367        assert!(validate_status(&value, 1).is_ok());
368    }
369
370    #[test]
371    fn label_validation_and_transcript_derivation_are_exact() {
372        let mut value = status();
373        value.state = OverallState::Completed;
374        value.interim_txid = Some(id(7));
375        value.analysis = Some(analysis(
376            "[high] Speaker 1: hi\n[medium] Speaker 2 [overlap]: yo\nplain Speaker 1: no\n",
377        ));
378        let labels = vec![label(1, "alice"), label(2, "bob")];
379        assert_eq!(validate_labels(&value, &labels), Ok(id(7)));
380        assert_eq!(
381            final_transcript(&value, &labels),
382            Ok((
383                id(7),
384                "[high] alice: hi\n[medium] bob [overlap]: yo\nplain Speaker 1: no\n".into()
385            ))
386        );
387        assert!(final_transcript(&value, &[label(2, "bob"), label(1, "alice")]).is_err());
388        assert!(final_transcript(&value, &[label(1, "\n"), label(2, "bob")]).is_err());
389    }
390}