Skip to main content

kcode_speech_classification/
ktool.rs

1//! Stable Kennedy Ktool contract for the in-process speaker classifier.
2
3use std::fmt;
4
5use serde::Deserialize;
6use serde_json::{Value, json};
7
8use crate::{
9    CandidateEvidence, Cefr, Cohort, DeleteOutcome, Error, FeatureRow, IdentifyEvidence,
10    IdentifyOutcome, ObservationKey, SpeechClassifier, TrainOutcome,
11};
12
13/// Model-facing identify operation name.
14pub const IDENTIFY_TOOL: &str = "kcode-speech-classification/identify";
15/// Model-facing train operation name.
16pub const TRAIN_TOOL: &str = "kcode-speech-classification/train";
17/// Model-facing delete operation name.
18pub const DELETE_TOOL: &str = "kcode-speech-classification/delete";
19/// Every model-facing operation owned by this classifier.
20pub const KTOOLS: [&str; 3] = [IDENTIFY_TOOL, TRAIN_TOOL, DELETE_TOOL];
21
22/// Failure from the complete in-process Ktool workflow.
23#[derive(Debug)]
24pub enum KtoolError {
25    /// The exact operation argument object could not be decoded.
26    InvalidArguments {
27        /// Stable Ktool operation name.
28        tool: &'static str,
29        /// Underlying JSON decoding error retained for diagnostics.
30        source: serde_json::Error,
31    },
32    /// The requested operation is not one of [`KTOOLS`].
33    UnsupportedTool(String),
34    /// The classifier rejected or could not persist the operation.
35    Classifier(Error),
36}
37
38impl fmt::Display for KtoolError {
39    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
40        match self {
41            Self::InvalidArguments { tool, .. } => {
42                write!(formatter, "decoding {tool} arguments")
43            }
44            Self::UnsupportedTool(tool) => {
45                write!(formatter, "unsupported speech-classification Ktool {tool}")
46            }
47            Self::Classifier(error) => error.fmt(formatter),
48        }
49    }
50}
51
52impl std::error::Error for KtoolError {
53    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
54        match self {
55            Self::InvalidArguments { source, .. } => Some(source),
56            Self::UnsupportedTool(_) => None,
57            Self::Classifier(error) => Some(error),
58        }
59    }
60}
61
62impl From<Error> for KtoolError {
63    fn from(value: Error) -> Self {
64        Self::Classifier(value)
65    }
66}
67
68/// One fully decoded model-facing classifier call.
69#[derive(Debug)]
70pub struct KtoolCall {
71    operation: KtoolOperation,
72}
73
74#[derive(Debug)]
75enum KtoolOperation {
76    Identify(IdentifyRequest),
77    Train(TrainRequest),
78    Delete(ObservationKey),
79}
80
81#[derive(Debug)]
82struct IdentifyRequest {
83    key: ObservationKey,
84    cohort: Cohort,
85    row: FeatureRow,
86    threshold: f64,
87}
88
89#[derive(Debug)]
90struct TrainRequest {
91    key: ObservationKey,
92    cohort: Cohort,
93    row: FeatureRow,
94    speaker_id: String,
95}
96
97#[derive(Deserialize)]
98#[serde(rename_all = "camelCase", deny_unknown_fields)]
99struct IdentifyArguments {
100    key: ObservationKeyInput,
101    cohort: CohortInput,
102    row: FeatureRowInput,
103    threshold: f64,
104}
105
106#[derive(Deserialize)]
107#[serde(rename_all = "camelCase", deny_unknown_fields)]
108struct TrainArguments {
109    key: ObservationKeyInput,
110    cohort: CohortInput,
111    row: FeatureRowInput,
112    speaker_id: String,
113}
114
115#[derive(Deserialize)]
116#[serde(rename_all = "camelCase", deny_unknown_fields)]
117struct DeleteArguments {
118    key: ObservationKeyInput,
119}
120
121#[derive(Deserialize)]
122#[serde(rename_all = "camelCase", deny_unknown_fields)]
123struct ObservationKeyInput {
124    object_id: String,
125    piece_index: u32,
126}
127
128impl From<ObservationKeyInput> for ObservationKey {
129    fn from(value: ObservationKeyInput) -> Self {
130        Self {
131            object_id: value.object_id,
132            piece_index: value.piece_index,
133        }
134    }
135}
136
137#[derive(Deserialize)]
138#[serde(rename_all = "camelCase", deny_unknown_fields)]
139struct CohortInput {
140    provider: String,
141    model: String,
142    prompt_version: String,
143    schema_version: String,
144    primary_language: String,
145}
146
147impl From<CohortInput> for Cohort {
148    fn from(value: CohortInput) -> Self {
149        Self {
150            provider: value.provider,
151            model: value.model,
152            prompt_version: value.prompt_version,
153            schema_version: value.schema_version,
154            primary_language: value.primary_language,
155        }
156    }
157}
158
159#[derive(Deserialize)]
160#[serde(rename_all = "camelCase", deny_unknown_fields)]
161struct FeatureRowInput {
162    accent_variety: String,
163    perceived_age: f64,
164    vocal_gender_presentation: f64,
165    median_f0_hz: f64,
166    formant_dispersion_hz: f64,
167    vai: f64,
168    hypernasality: f64,
169    creaky_phonation_percent: f64,
170    rhotic_realization: String,
171    word_initial_stressed_prevocalic_t_vot_ms: f64,
172    breathiness: f64,
173    roughness: f64,
174    f0_pitch_span_semitones: f64,
175    articulation_rate_syllables_per_second: f64,
176    npvi_v: f64,
177    cefr: Cefr,
178    foreign_accentedness: f64,
179    unstressed_vowel_reduction_percent: f64,
180    lateral_realization: String,
181    filled_pauses_per_100_words: f64,
182    s_realization: String,
183    lexical_stress_accuracy_percent: f64,
184    monophthongization_percent: f64,
185    consonant_cluster_reduction_percent: f64,
186}
187
188impl From<FeatureRowInput> for FeatureRow {
189    fn from(value: FeatureRowInput) -> Self {
190        Self {
191            accent_variety: value.accent_variety,
192            perceived_age: value.perceived_age,
193            vocal_gender_presentation: value.vocal_gender_presentation,
194            median_f0_hz: value.median_f0_hz,
195            formant_dispersion_hz: value.formant_dispersion_hz,
196            vai: value.vai,
197            hypernasality: value.hypernasality,
198            creaky_phonation_percent: value.creaky_phonation_percent,
199            rhotic_realization: value.rhotic_realization,
200            word_initial_stressed_prevocalic_t_vot_ms: value
201                .word_initial_stressed_prevocalic_t_vot_ms,
202            breathiness: value.breathiness,
203            roughness: value.roughness,
204            f0_pitch_span_semitones: value.f0_pitch_span_semitones,
205            articulation_rate_syllables_per_second: value.articulation_rate_syllables_per_second,
206            npvi_v: value.npvi_v,
207            cefr: value.cefr,
208            foreign_accentedness: value.foreign_accentedness,
209            unstressed_vowel_reduction_percent: value.unstressed_vowel_reduction_percent,
210            lateral_realization: value.lateral_realization,
211            filled_pauses_per_100_words: value.filled_pauses_per_100_words,
212            s_realization: value.s_realization,
213            lexical_stress_accuracy_percent: value.lexical_stress_accuracy_percent,
214            monophthongization_percent: value.monophthongization_percent,
215            consonant_cluster_reduction_percent: value.consonant_cluster_reduction_percent,
216        }
217    }
218}
219
220impl SpeechClassifier {
221    /// Validates, executes, and renders one decoded model-facing classifier operation.
222    pub fn execute_ktool(&self, call: KtoolCall) -> Result<String, KtoolError> {
223        match call.operation {
224            KtoolOperation::Identify(request) => self
225                .identify(request.key, request.cohort, request.row, request.threshold)
226                .map(render_identify_outcome)
227                .map_err(KtoolError::Classifier),
228            KtoolOperation::Train(request) => self
229                .train(request.key, request.cohort, request.row, request.speaker_id)
230                .map(render_train_outcome)
231                .map_err(KtoolError::Classifier),
232            KtoolOperation::Delete(key) => self
233                .delete(key)
234                .map(render_delete_outcome)
235                .map_err(KtoolError::Classifier),
236        }
237    }
238}
239
240/// Decodes one exact Ktool name and argument object without touching classifier state.
241pub fn decode_ktool(tool: &str, arguments: &Value) -> Result<KtoolCall, KtoolError> {
242    let operation = match tool {
243        IDENTIFY_TOOL => KtoolOperation::Identify(identify_request(arguments)?),
244        TRAIN_TOOL => KtoolOperation::Train(train_request(arguments)?),
245        DELETE_TOOL => KtoolOperation::Delete(delete_request(arguments)?),
246        _ => return Err(KtoolError::UnsupportedTool(tool.to_owned())),
247    };
248    Ok(KtoolCall { operation })
249}
250
251fn identify_request(arguments: &Value) -> Result<IdentifyRequest, KtoolError> {
252    let arguments =
253        serde_json::from_value::<IdentifyArguments>(arguments.clone()).map_err(|source| {
254            KtoolError::InvalidArguments {
255                tool: IDENTIFY_TOOL,
256                source,
257            }
258        })?;
259    Ok(IdentifyRequest {
260        key: arguments.key.into(),
261        cohort: arguments.cohort.into(),
262        row: arguments.row.into(),
263        threshold: arguments.threshold,
264    })
265}
266
267fn train_request(arguments: &Value) -> Result<TrainRequest, KtoolError> {
268    let arguments =
269        serde_json::from_value::<TrainArguments>(arguments.clone()).map_err(|source| {
270            KtoolError::InvalidArguments {
271                tool: TRAIN_TOOL,
272                source,
273            }
274        })?;
275    Ok(TrainRequest {
276        key: arguments.key.into(),
277        cohort: arguments.cohort.into(),
278        row: arguments.row.into(),
279        speaker_id: arguments.speaker_id,
280    })
281}
282
283fn delete_request(arguments: &Value) -> Result<ObservationKey, KtoolError> {
284    serde_json::from_value::<DeleteArguments>(arguments.clone())
285        .map(|arguments| arguments.key.into())
286        .map_err(|source| KtoolError::InvalidArguments {
287            tool: DELETE_TOOL,
288            source,
289        })
290}
291
292fn render_identify_outcome(outcome: IdentifyOutcome) -> String {
293    let retained = outcome.speaker_id.is_some();
294    render_json(json!({
295        "operation":"identify",
296        "speakerId":outcome.speaker_id,
297        "retained":retained,
298        "evidence":outcome.evidence.map(evidence_json),
299    }))
300}
301
302fn render_train_outcome(outcome: TrainOutcome) -> String {
303    let outcome = match outcome {
304        TrainOutcome::Added => "added",
305        TrainOutcome::Unchanged => "unchanged",
306        TrainOutcome::Corrected => "corrected",
307    };
308    render_json(json!({"operation":"train", "outcome":outcome}))
309}
310
311fn render_delete_outcome(outcome: DeleteOutcome) -> String {
312    let outcome = match outcome {
313        DeleteOutcome::Deleted => "deleted",
314        DeleteOutcome::NotFound => "not_found",
315    };
316    render_json(json!({"operation":"delete", "outcome":outcome}))
317}
318
319fn evidence_json(evidence: IdentifyEvidence) -> Value {
320    json!({
321        "best":candidate_json(evidence.best),
322        "runnerUp":evidence.runner_up.map(candidate_json),
323        "backgroundPopulationCost":evidence.background_population_cost,
324        "absoluteGap":evidence.absolute_gap,
325        "runnerUpGap":evidence.runner_up_gap,
326        "confidenceScore":evidence.confidence_score,
327    })
328}
329
330fn candidate_json(candidate: CandidateEvidence) -> Value {
331    json!({"speakerId":candidate.speaker_id, "cost":candidate.cost})
332}
333
334fn render_json(value: Value) -> String {
335    serde_json::to_string_pretty(&value)
336        .expect("speaker-classification outcomes contain finite JSON")
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342    use std::sync::atomic::{AtomicU64, Ordering};
343
344    static NEXT_PATH: AtomicU64 = AtomicU64::new(0);
345
346    fn row() -> Value {
347        json!({
348            "accentVariety":"General American",
349            "perceivedAge":36.0,
350            "vocalGenderPresentation":45.0,
351            "medianF0Hz":155.0,
352            "formantDispersionHz":1100.0,
353            "vai":1.1,
354            "hypernasality":0.5,
355            "creakyPhonationPercent":8.0,
356            "rhoticRealization":"rhotic",
357            "wordInitialStressedPrevocalicTVotMs":62.0,
358            "breathiness":20.0,
359            "roughness":10.0,
360            "f0PitchSpanSemitones":9.0,
361            "articulationRateSyllablesPerSecond":4.2,
362            "npviV":48.0,
363            "cefr":"C1",
364            "foreignAccentedness":2.0,
365            "unstressedVowelReductionPercent":72.0,
366            "lateralRealization":"alveolar",
367            "filledPausesPer100Words":2.5,
368            "sRealization":"alveolar",
369            "lexicalStressAccuracyPercent":92.0,
370            "monophthongizationPercent":4.0,
371            "consonantClusterReductionPercent":3.0
372        })
373    }
374
375    fn speaker_row(age: f64, accent: &str) -> Value {
376        let mut row = row();
377        row["perceivedAge"] = json!(age);
378        row["medianF0Hz"] = json!(120.0 + age);
379        row["accentVariety"] = json!(accent);
380        row["rhoticRealization"] = json!(format!("{accent}-rhotic"));
381        row
382    }
383
384    fn cohort() -> Value {
385        json!({
386            "provider":"google",
387            "model":"gemini-example",
388            "promptVersion":"speaker-features-1",
389            "schemaVersion":"features-1",
390            "primaryLanguage":"eng"
391        })
392    }
393
394    fn key() -> Value {
395        json!({"objectId":"AAECAwQF", "pieceIndex":3})
396    }
397
398    fn execute(classifier: &SpeechClassifier, tool: &str, arguments: Value) -> String {
399        let call = decode_ktool(tool, &arguments).unwrap();
400        classifier.execute_ktool(call).unwrap()
401    }
402
403    #[test]
404    fn camel_case_tool_contract_maps_every_classifier_field() {
405        let identify = identify_request(&json!({
406            "key":key(),
407            "cohort":cohort(),
408            "row":row(),
409            "threshold":2.5
410        }))
411        .unwrap();
412        assert_eq!(identify.key.object_id, "AAECAwQF");
413        assert_eq!(identify.key.piece_index, 3);
414        assert_eq!(identify.cohort.prompt_version, "speaker-features-1");
415        assert_eq!(identify.row.cefr, Cefr::C1);
416        assert_eq!(identify.row.consonant_cluster_reduction_percent, 3.0);
417        assert_eq!(identify.threshold, 2.5);
418
419        let train = train_request(&json!({
420            "key":key(),
421            "cohort":cohort(),
422            "row":row(),
423            "speakerId":"kennedy"
424        }))
425        .unwrap();
426        assert_eq!(train.speaker_id, "kennedy");
427
428        let delete = delete_request(&json!({"key":key()})).unwrap();
429        assert_eq!(delete.object_id, "AAECAwQF");
430        assert_eq!(delete.piece_index, 3);
431    }
432
433    #[test]
434    fn tool_contract_rejects_unknown_fields_at_every_level() {
435        let error = delete_request(&json!({"key":key(), "speakerId":"unexpected"})).unwrap_err();
436        assert!(
437            error
438                .to_string()
439                .contains("kcode-speech-classification/delete")
440        );
441
442        let mut row = row();
443        row["unexpected"] = json!(true);
444        let error = identify_request(&json!({
445            "key":key(),
446            "cohort":cohort(),
447            "row":row,
448            "threshold":2.5
449        }))
450        .unwrap_err();
451        assert!(
452            error
453                .to_string()
454                .contains("kcode-speech-classification/identify")
455        );
456    }
457
458    #[test]
459    fn outcomes_are_rendered_as_stable_camel_case_json() {
460        let rendered = render_identify_outcome(IdentifyOutcome {
461            speaker_id: Some("speaker-a".into()),
462            evidence: Some(IdentifyEvidence {
463                best: CandidateEvidence {
464                    speaker_id: "speaker-a".into(),
465                    cost: 4.0,
466                },
467                runner_up: None,
468                background_population_cost: 8.0,
469                absolute_gap: 4.0,
470                runner_up_gap: None,
471                confidence_score: 4.0,
472            }),
473        });
474        let rendered: Value = serde_json::from_str(&rendered).unwrap();
475        assert_eq!(rendered["speakerId"], "speaker-a");
476        assert_eq!(rendered["retained"], true);
477        assert_eq!(rendered["evidence"]["confidenceScore"], 4.0);
478        assert!(rendered["evidence"].get("confidence_score").is_none());
479
480        let rendered: Value =
481            serde_json::from_str(&render_train_outcome(TrainOutcome::Corrected)).unwrap();
482        assert_eq!(
483            rendered,
484            json!({"operation":"train", "outcome":"corrected"})
485        );
486        let rendered: Value =
487            serde_json::from_str(&render_delete_outcome(DeleteOutcome::NotFound)).unwrap();
488        assert_eq!(
489            rendered,
490            json!({"operation":"delete", "outcome":"not_found"})
491        );
492    }
493
494    #[test]
495    fn complete_tool_workflow_preserves_labels_identification_and_deletion() {
496        let path = std::env::temp_dir().join(format!(
497            "kennedy-speech-classification-tool-test-{}-{}.sqlite3",
498            std::process::id(),
499            NEXT_PATH.fetch_add(1, Ordering::Relaxed)
500        ));
501        let classifier = SpeechClassifier::open(&path).unwrap();
502
503        for (object_id, age, accent, speaker_id) in [
504            ("alpha-1", 30.0, "alpha", "speaker-a"),
505            ("alpha-2", 32.0, "alpha", "speaker-a"),
506            ("beta-1", 70.0, "beta", "speaker-b"),
507            ("beta-2", 72.0, "beta", "speaker-b"),
508        ] {
509            let rendered = execute(
510                &classifier,
511                TRAIN_TOOL,
512                json!({
513                    "key":{"objectId":object_id, "pieceIndex":0},
514                    "cohort":cohort(),
515                    "row":speaker_row(age, accent),
516                    "speakerId":speaker_id
517                }),
518            );
519            let rendered: Value = serde_json::from_str(&rendered).unwrap();
520            assert_eq!(rendered["outcome"], "added");
521        }
522
523        let rendered = execute(
524            &classifier,
525            IDENTIFY_TOOL,
526            json!({
527                "key":{"objectId":"query", "pieceIndex":0},
528                "cohort":cohort(),
529                "row":speaker_row(31.0, "alpha"),
530                "threshold":-1_000_000.0
531            }),
532        );
533        let rendered: Value = serde_json::from_str(&rendered).unwrap();
534        assert_eq!(rendered["speakerId"], "speaker-a");
535        assert_eq!(rendered["retained"], true);
536
537        let rendered = execute(
538            &classifier,
539            DELETE_TOOL,
540            json!({"key":{"objectId":"query", "pieceIndex":0}}),
541        );
542        let rendered: Value = serde_json::from_str(&rendered).unwrap();
543        assert_eq!(rendered["outcome"], "deleted");
544
545        drop(classifier);
546        std::fs::remove_file(path).unwrap();
547    }
548}