Skip to main content

kcode_k1_audio_classification_projection_state/
lib.rs

1pub use kcode_k1_audio_classification_format::{
2    ExecutedAnalysis, FragmentStageV1, PersonId, 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    }
106    Ok((
107        interim,
108        replace_transcript(&analysis.envelope.analysis.transcript, labels),
109    ))
110}
111
112pub fn actionable_state(state: OverallState) -> i64 {
113    match state {
114        OverallState::Queued => 1,
115        OverallState::Running => 2,
116        _ => 0,
117    }
118}
119
120pub fn interrupted_stage(status: &FragmentStatus) -> FragmentStageV1 {
121    for stage in [
122        &status.structuring,
123        &status.speaker_features,
124        &status.speaker_labels,
125        &status.transcript,
126    ] {
127        if stage.state == StageState::Running {
128            return stage.stage;
129        }
130    }
131    FragmentStageV1::Queue
132}
133
134pub fn validate_status(
135    status: &FragmentStatus,
136    stored_actionable_state: i64,
137) -> Result<(), String> {
138    let identities = [
139        (&status.queue, FragmentStageV1::Queue),
140        (&status.transcript, FragmentStageV1::Transcript),
141        (&status.speaker_labels, FragmentStageV1::SpeakerLabels),
142        (&status.speaker_features, FragmentStageV1::SpeakerFeatures),
143        (&status.structuring, FragmentStageV1::Structuring),
144        (
145            &status.label_confirmation,
146            FragmentStageV1::LabelConfirmation,
147        ),
148    ];
149    if identities
150        .iter()
151        .any(|(stored, expected)| stored.stage != *expected)
152    {
153        return Err("stored stage identity does not match its field".to_string());
154    }
155    if actionable_state(status.state) != stored_actionable_state {
156        return Err("stored actionable state does not match status".to_string());
157    }
158    if status.errors.len() > MAX_ERRORS {
159        return Err("stored errors exceed the retention bound".to_string());
160    }
161    let mut previous = None;
162    for job in &status.jobs {
163        if job.attempt == 0
164            || job.attempt > status.attempt_count
165            || !is_analysis_stage(&job.stage)
166            || job.name.trim().is_empty()
167            || previous.is_some_and(|value| value >= (job.attempt, job.sequence))
168        {
169            return Err("stored LLM jobs are invalid or unordered".to_string());
170        }
171        previous = Some((job.attempt, job.sequence));
172    }
173    if !matches!(
174        status.queue.state,
175        StageState::Succeeded | StageState::Failed
176    ) {
177        return Err("stored Queue stage is neither succeeded nor failed".to_string());
178    }
179    if matches!(
180        status.state,
181        OverallState::Completed | OverallState::Confirmed
182    ) && (status.interim_txid.is_none() || status.analysis.is_none())
183    {
184        return Err("stored completed status lacks its analysis".to_string());
185    }
186    if status.state == OverallState::Confirmed
187        && (status.final_transcript.is_none()
188            || status.label_confirmation.state != StageState::Succeeded)
189    {
190        return Err("stored confirmed status is incomplete".to_string());
191    }
192    Ok(())
193}
194
195fn replace_transcript(transcript: &str, labels: &[SpeakerLabelV1]) -> String {
196    let mut output = String::with_capacity(transcript.len());
197    for line in transcript.split_inclusive('\n') {
198        let mut replaced = false;
199        for prefix in ["[high] ", "[medium] ", "[low] "] {
200            if let Some(rest) = line.strip_prefix(prefix) {
201                for label in labels {
202                    let speaker = label.speaker.to_string();
203                    if let Some(tail) = rest.strip_prefix(&speaker)
204                        && (tail.starts_with(':') || tail.starts_with(" [overlap]:"))
205                    {
206                        output.push_str(prefix);
207                        match label.person_id {
208                            Some(person_id) => output.push_str(&person_id.to_string()),
209                            None => output.push_str("Unknown"),
210                        }
211                        output.push_str(tail);
212                        replaced = true;
213                        break;
214                    }
215                }
216            }
217            if replaced {
218                break;
219            }
220        }
221        if !replaced {
222            output.push_str(line);
223        }
224    }
225    output
226}
227
228fn is_analysis_stage(stage: &FragmentStageV1) -> bool {
229    matches!(
230        stage,
231        FragmentStageV1::Transcript
232            | FragmentStageV1::SpeakerLabels
233            | FragmentStageV1::SpeakerFeatures
234            | FragmentStageV1::Structuring
235    )
236}
237
238mod optional_fragment_id {
239    use super::FragmentId;
240    use serde::{Deserialize, Deserializer, Serialize, Serializer};
241
242    pub fn serialize<S: Serializer>(
243        value: &Option<FragmentId>,
244        serializer: S,
245    ) -> Result<S::Ok, S::Error> {
246        value.map(FragmentId::into_bytes).serialize(serializer)
247    }
248
249    pub fn deserialize<'de, D: Deserializer<'de>>(
250        deserializer: D,
251    ) -> Result<Option<FragmentId>, D::Error> {
252        Option::<[u8; 12]>::deserialize(deserializer).map(|value| value.map(FragmentId::from_bytes))
253    }
254}
255
256#[cfg(test)]
257mod tests {
258    use super::*;
259    use kcode_speaker_v3_analysis::{
260        AnalysisEnvelope, FeatureVector24, GeminiCohort, LocalSpeakerLabel, OggAudioMetadata,
261        StructuredAnalysis, StructuredSpeaker, StructurerProvenance,
262    };
263
264    fn id(value: u8) -> FragmentId {
265        FragmentId::from_bytes([value; 12])
266    }
267
268    fn person(value: u8) -> PersonId {
269        PersonId::from_tx_id(id(value))
270    }
271
272    fn stage(stage: FragmentStageV1, state: StageState) -> StageStatus {
273        StageStatus { stage, state }
274    }
275
276    fn status() -> FragmentStatus {
277        FragmentStatus {
278            state: OverallState::Queued,
279            queue: stage(FragmentStageV1::Queue, StageState::Succeeded),
280            transcript: stage(FragmentStageV1::Transcript, StageState::Pending),
281            speaker_labels: stage(FragmentStageV1::SpeakerLabels, StageState::Pending),
282            speaker_features: stage(FragmentStageV1::SpeakerFeatures, StageState::Pending),
283            structuring: stage(FragmentStageV1::Structuring, StageState::Pending),
284            label_confirmation: stage(FragmentStageV1::LabelConfirmation, StageState::Pending),
285            attempt_count: 0,
286            jobs: Vec::new(),
287            interim_txid: None,
288            analysis: None,
289            confirmed_labels: Vec::new(),
290            final_transcript: None,
291            errors: Vec::new(),
292            errors_truncated: false,
293        }
294    }
295
296    fn analysis(transcript: &str) -> ExecutedAnalysis {
297        let mut ogg = vec![0; 29];
298        ogg[..4].copy_from_slice(b"OggS");
299        ogg[26] = 1;
300        ogg[27] = 1;
301        let speaker = |number| StructuredSpeaker {
302            speaker: LocalSpeakerLabel::new(number).unwrap(),
303            language: "en".into(),
304            features: FeatureVector24::default(),
305            features_usable_for_training: true,
306        };
307        let provenance = StructurerProvenance {
308            model_id: "model".into(),
309            prompt_revision: "prompt".into(),
310        };
311        ExecutedAnalysis {
312            envelope: AnalysisEnvelope {
313                audio: OggAudioMetadata::from_bytes(&ogg, 1, None).unwrap(),
314                analysis: StructuredAnalysis {
315                    transcript: transcript.into(),
316                    speakers: vec![speaker(1), speaker(2)],
317                },
318                gemini: GeminiCohort {
319                    model_id: "gemini".into(),
320                    transcript_prompt_revision: "t".into(),
321                    feature_prompt_revisions: ["1".into(), "2".into(), "3".into()],
322                    feature_schema_revision: "s".into(),
323                },
324                structurer: provenance.clone(),
325            },
326            label_extractor: provenance,
327        }
328    }
329
330    fn label(number: u32, person_id: Option<PersonId>) -> SpeakerLabelV1 {
331        SpeakerLabelV1 {
332            speaker: LocalSpeakerLabel::new(number).unwrap(),
333            person_id,
334        }
335    }
336
337    #[test]
338    fn validation_checks_persisted_invariants() {
339        let mut value = status();
340        assert_eq!(actionable_state(value.state), 1);
341        assert_eq!(interrupted_stage(&value), FragmentStageV1::Queue);
342        assert!(validate_status(&value, 1).is_ok());
343        value.transcript.state = StageState::Running;
344        value.structuring.state = StageState::Running;
345        assert_eq!(interrupted_stage(&value), FragmentStageV1::Structuring);
346        value.transcript.stage = FragmentStageV1::Structuring;
347        assert!(validate_status(&value, 1).is_err());
348    }
349
350    #[test]
351    fn error_retention_keeps_the_oldest_bound() {
352        let mut value = status();
353        for index in 0..=MAX_ERRORS {
354            append_error(&mut value, index.to_string());
355        }
356        assert_eq!(value.errors.len(), MAX_ERRORS);
357        assert_eq!(value.errors.first().map(String::as_str), Some("0"));
358        assert_eq!(value.errors.last().map(String::as_str), Some("4999"));
359        assert!(value.errors_truncated);
360        assert!(validate_status(&value, 1).is_ok());
361    }
362
363    #[test]
364    fn label_validation_and_transcript_derivation_are_exact() {
365        let mut value = status();
366        value.state = OverallState::Completed;
367        value.interim_txid = Some(id(7));
368        value.analysis = Some(analysis(
369            "[high] Speaker 1: hi\n[medium] Speaker 2 [overlap]: yo\nplain Speaker 1: no\n[low] Speaker 1 [overlap]: end",
370        ));
371        let known = person(0xab);
372        let labels = vec![label(1, Some(known)), label(2, None)];
373        assert_eq!(validate_labels(&value, &labels), Ok(id(7)));
374        assert_eq!(
375            final_transcript(&value, &labels),
376            Ok((
377                id(7),
378                "[high] abababababababababababab: hi\n[medium] Unknown [overlap]: yo\nplain Speaker 1: no\n[low] abababababababababababab [overlap]: end".into()
379            ))
380        );
381        assert!(final_transcript(&value, &[label(2, None), label(1, Some(known))]).is_err());
382        assert!(final_transcript(&value, &[label(1, Some(known))]).is_err());
383    }
384}