Skip to main content

kcode_k1_audio_classification_projection_fold/
lib.rs

1use kcode_k1_audio_classification_event_format::{
2    AttemptFinalAnalysisV1, AudioClassificationEventV3 as V5, AudioClassificationEventV6 as V6,
3    decode_event, decode_event_v6, decode_final_event_v7,
4};
5use kcode_k1_audio_classification_projection_artifacts::AttemptArtifacts;
6use kcode_k1_audio_classification_projection_state as state;
7use kcode_k1_audio_classification_projection_v5_fold::fold_v5 as fold_decoded_v5;
8use kcode_k1_audio_classification_projection_v6_fold::fold_v6 as fold_decoded_v6;
9
10pub use kcode_k1_audio_classification_projection_artifacts::{
11    LocalSpeakerLabel, SpeakerFeatureEvidence, TxId,
12};
13pub use kcode_k1_audio_classification_projection_state::{
14    ExecutedAnalysis, FragmentId, FragmentStageV1, FragmentStatus, LlmJobState, LlmJobStatus,
15    OverallState, SpeakerLabelV1, StageState, StageStatus,
16};
17pub use kcode_k1_audio_classification_projection_v5_fold::ProjectionEffect;
18
19#[derive(Clone)]
20pub struct FragmentProjection {
21    status: FragmentStatus,
22    artifacts: AttemptArtifacts,
23}
24
25#[derive(Clone, Debug, PartialEq)]
26pub struct ResumePlan {
27    pub fragment_id: FragmentId,
28    pub transcript: Option<String>,
29    pub speaker_labels: Option<Vec<LocalSpeakerLabel>>,
30    pub speaker_features: Vec<SpeakerFeatureEvidence>,
31    pub final_analysis: Option<ExecutedAnalysis>,
32    pub active_attempt: Option<TxId>,
33    pub interrupted: bool,
34}
35
36#[derive(Clone)]
37pub struct FoldResult {
38    pub fragment_id: FragmentId,
39    pub effect: ProjectionEffect,
40    pub replacement: Option<FragmentProjection>,
41}
42
43impl FragmentProjection {
44    pub fn fragment_id(&self) -> FragmentId {
45        self.artifacts.fragment_id()
46    }
47
48    pub fn visible_status(&self) -> FragmentStatus {
49        let mut status = self.status.clone();
50        overlay_stages(&mut status, &self.artifacts);
51        status.state = projected_state(&status, &self.artifacts);
52        status
53    }
54
55    pub fn resume_plan(&self) -> ResumePlan {
56        let interrupted = !self.artifacts.complete() && self.status.state == OverallState::Running;
57        ResumePlan {
58            fragment_id: self.fragment_id(),
59            transcript: self.artifacts.transcript().map(str::to_owned),
60            speaker_labels: self.artifacts.speaker_labels().map(<[_]>::to_vec),
61            speaker_features: self.artifacts.speaker_features().to_vec(),
62            final_analysis: self.artifacts.final_analysis().cloned(),
63            active_attempt: self.artifacts.active_attempt().filter(|_| interrupted),
64            interrupted,
65        }
66    }
67
68    pub fn needs_work(&self) -> bool {
69        let state = self.visible_status().state;
70        !self.artifacts.complete() && matches!(state, OverallState::Queued | OverallState::Running)
71    }
72
73    pub fn validate_labels(&self, labels: &[SpeakerLabelV1]) -> Result<TxId, String> {
74        let mut status = self.status.clone();
75        if status.analysis.is_some() && status.interim_txid.is_some() {
76            status.state = OverallState::Completed;
77        }
78        state::validate_labels(&status, labels)
79    }
80
81    #[cfg(feature = "testkit")]
82    pub fn inject_errors(&mut self, errors: Vec<String>) {
83        for error in errors {
84            state::append_error(&mut self.status, error);
85        }
86    }
87}
88
89pub fn fold(
90    current: Option<&FragmentProjection>,
91    callback_txid: TxId,
92    payload: &[u8],
93) -> Result<FoldResult, String> {
94    match payload.first() {
95        Some(6) => fold_v6(
96            current,
97            callback_txid,
98            &decode_event_v6(payload).map_err(message)?,
99        ),
100        Some(7) => fold_v7(
101            current,
102            callback_txid,
103            &decode_final_event_v7(payload).map_err(message)?,
104        ),
105        _ => fold_v5(
106            current,
107            callback_txid,
108            &decode_event(payload).map_err(message)?,
109        ),
110    }
111}
112
113fn fold_v5(
114    current: Option<&FragmentProjection>,
115    callback: TxId,
116    event: &V5,
117) -> Result<FoldResult, String> {
118    let update = fold_decoded_v5(
119        current.map(|value| (&value.status, &value.artifacts)),
120        callback,
121        event,
122    )?;
123    let replacement = update
124        .replacement
125        .map(|value| projection(value.status, value.artifacts));
126    Ok(FoldResult {
127        fragment_id: update.fragment_id,
128        effect: update.effect,
129        replacement,
130    })
131}
132
133fn fold_v6(
134    current: Option<&FragmentProjection>,
135    callback: TxId,
136    event: &V6,
137) -> Result<FoldResult, String> {
138    let fragment = v6_fragment(event);
139    let current = current.ok_or_else(|| "event references an unknown fragment".to_string())?;
140    require(fragment == current.fragment_id(), "V6 fragment mismatch")?;
141    let update = fold_decoded_v6(&current.status, &current.artifacts, callback, event)?;
142    let replacement = update
143        .replacement
144        .map(|value| projection(value.status, value.artifacts));
145    Ok(none_effect(fragment, replacement))
146}
147
148fn fold_v7(
149    current: Option<&FragmentProjection>,
150    callback: TxId,
151    event: &AttemptFinalAnalysisV1,
152) -> Result<FoldResult, String> {
153    let fragment = event.fragment_id;
154    let current = current.ok_or_else(|| "event references an unknown fragment".to_string())?;
155    require(fragment == current.fragment_id(), "V7 fragment mismatch")?;
156    if current.status.state != OverallState::Running
157        || current.artifacts.active_attempt() != Some(event.attempt_txid)
158    {
159        return Ok(none_effect(fragment, None));
160    }
161    let update = current.artifacts.fold_v7(callback, event)?;
162    let Some(artifacts) = update.replacement else {
163        return Ok(none_effect(fragment, None));
164    };
165    let analysis = artifacts
166        .final_analysis()
167        .cloned()
168        .ok_or("final analysis is absent after V7")?;
169    let mut status = current.status.clone();
170    status.analysis = Some(analysis);
171    status.interim_txid = Some(callback);
172    status.transcript.state = StageState::Succeeded;
173    status.speaker_labels.state = StageState::Succeeded;
174    status.speaker_features.state = StageState::Succeeded;
175    status.structuring.state = StageState::Succeeded;
176    status.state = projected_state(&status, &artifacts);
177    Ok(none_effect(fragment, Some(projection(status, artifacts))))
178}
179
180fn overlay_stages(status: &mut FragmentStatus, artifacts: &AttemptArtifacts) {
181    let terminal = matches!(
182        status.state,
183        OverallState::Completed | OverallState::Confirmed
184    );
185    for (stage, artifact_stage) in [
186        (&mut status.transcript, FragmentStageV1::Transcript),
187        (&mut status.speaker_labels, FragmentStageV1::SpeakerLabels),
188        (
189            &mut status.speaker_features,
190            FragmentStageV1::SpeakerFeatures,
191        ),
192        (&mut status.structuring, FragmentStageV1::Structuring),
193    ] {
194        if artifacts.stage_present(artifact_stage) {
195            stage.state = StageState::Succeeded;
196        } else if terminal {
197            stage.state = StageState::Pending;
198        }
199    }
200}
201
202fn projected_state(status: &FragmentStatus, artifacts: &AttemptArtifacts) -> OverallState {
203    let confirmed = artifacts.complete()
204        && status.label_confirmation.state == StageState::Succeeded
205        && status.final_transcript.is_some();
206    match status.state {
207        OverallState::Discarded => OverallState::Discarded,
208        _ if confirmed => OverallState::Confirmed,
209        _ if artifacts.complete() => OverallState::Completed,
210        OverallState::Failed => OverallState::Failed,
211        OverallState::Running => OverallState::Running,
212        _ => OverallState::Queued,
213    }
214}
215
216fn v6_fragment(event: &V6) -> FragmentId {
217    match event {
218        V6::AttemptStarted(value) => value.fragment_id,
219        V6::Progress(value) => value.fragment_id,
220        V6::GeminiTranscript(value) => value.fragment_id,
221        V6::TerraSpeakerLabels(value) => value.fragment_id,
222        V6::GeminiFeatureBundle(value) => value.fragment_id,
223        V6::Failed(value) => value.fragment_id,
224    }
225}
226
227fn projection(status: FragmentStatus, artifacts: AttemptArtifacts) -> FragmentProjection {
228    FragmentProjection { status, artifacts }
229}
230
231fn message(error: impl ToString) -> String {
232    error.to_string()
233}
234
235fn require(condition: bool, message: &str) -> Result<(), String> {
236    condition.then_some(()).ok_or_else(|| message.to_string())
237}
238
239fn none_effect(fragment_id: FragmentId, replacement: Option<FragmentProjection>) -> FoldResult {
240    FoldResult {
241        fragment_id,
242        effect: ProjectionEffect::None,
243        replacement,
244    }
245}
246
247#[cfg(test)]
248mod tests;