Skip to main content

kcode_speaker_system/
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, Cohort, DeleteOutcome, Error, FeatureRow, IdentifyEvidence, IdentifyOutcome,
10    ObservationKey, SpeechClassifier, TrainOutcome,
11};
12
13pub const IDENTIFY_TOOL: &str = "kcode-speaker-system/identify";
14pub const TRAIN_TOOL: &str = "kcode-speaker-system/train";
15pub const DELETE_TOOL: &str = "kcode-speaker-system/delete";
16pub const KTOOLS: [&str; 3] = [IDENTIFY_TOOL, TRAIN_TOOL, DELETE_TOOL];
17
18#[derive(Debug)]
19pub enum KtoolError {
20    InvalidArguments {
21        tool: &'static str,
22        source: serde_json::Error,
23    },
24    UnsupportedTool(String),
25    Classifier(Error),
26}
27
28impl fmt::Display for KtoolError {
29    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
30        match self {
31            Self::InvalidArguments { tool, .. } => write!(formatter, "decoding {tool} arguments"),
32            Self::UnsupportedTool(tool) => {
33                write!(formatter, "unsupported speaker-system Ktool {tool}")
34            }
35            Self::Classifier(error) => error.fmt(formatter),
36        }
37    }
38}
39
40impl std::error::Error for KtoolError {
41    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
42        match self {
43            Self::InvalidArguments { source, .. } => Some(source),
44            Self::UnsupportedTool(_) => None,
45            Self::Classifier(error) => Some(error),
46        }
47    }
48}
49
50impl From<Error> for KtoolError {
51    fn from(value: Error) -> Self {
52        Self::Classifier(value)
53    }
54}
55
56#[derive(Debug)]
57pub struct KtoolCall {
58    operation: KtoolOperation,
59}
60
61#[derive(Debug)]
62enum KtoolOperation {
63    Identify(IdentifyRequest),
64    Train(TrainRequest),
65    Delete(ObservationKey),
66}
67
68#[derive(Debug)]
69struct IdentifyRequest {
70    key: ObservationKey,
71    cohort: Cohort,
72    row: FeatureRow,
73    threshold: f64,
74}
75
76#[derive(Debug)]
77struct TrainRequest {
78    key: ObservationKey,
79    cohort: Cohort,
80    row: FeatureRow,
81    speaker_id: String,
82}
83
84#[derive(Deserialize)]
85#[serde(rename_all = "camelCase", deny_unknown_fields)]
86struct IdentifyArguments {
87    key: ObservationKeyInput,
88    cohort: CohortInput,
89    row: FeatureRow,
90    threshold: f64,
91}
92
93#[derive(Deserialize)]
94#[serde(rename_all = "camelCase", deny_unknown_fields)]
95struct TrainArguments {
96    key: ObservationKeyInput,
97    cohort: CohortInput,
98    row: FeatureRow,
99    speaker_id: String,
100}
101
102#[derive(Deserialize)]
103#[serde(rename_all = "camelCase", deny_unknown_fields)]
104struct DeleteArguments {
105    key: ObservationKeyInput,
106}
107
108#[derive(Deserialize)]
109#[serde(rename_all = "camelCase", deny_unknown_fields)]
110struct ObservationKeyInput {
111    object_id: String,
112    piece_index: u32,
113}
114
115impl From<ObservationKeyInput> for ObservationKey {
116    fn from(value: ObservationKeyInput) -> Self {
117        Self {
118            object_id: value.object_id,
119            piece_index: value.piece_index,
120        }
121    }
122}
123
124#[derive(Deserialize)]
125#[serde(rename_all = "camelCase", deny_unknown_fields)]
126struct CohortInput {
127    provider: String,
128    model: String,
129    prompt_version: String,
130    schema_version: String,
131    primary_language: String,
132}
133
134impl From<CohortInput> for Cohort {
135    fn from(value: CohortInput) -> Self {
136        Self {
137            provider: value.provider,
138            model: value.model,
139            prompt_version: value.prompt_version,
140            schema_version: value.schema_version,
141            primary_language: value.primary_language,
142        }
143    }
144}
145
146impl SpeechClassifier {
147    /// Executes and renders one already decoded classifier operation.
148    pub fn execute_ktool(&self, call: KtoolCall) -> Result<String, KtoolError> {
149        match call.operation {
150            KtoolOperation::Identify(request) => self
151                .identify(request.key, request.cohort, request.row, request.threshold)
152                .map(render_identify_outcome)
153                .map_err(KtoolError::Classifier),
154            KtoolOperation::Train(request) => self
155                .train(request.key, request.cohort, request.row, request.speaker_id)
156                .map(render_train_outcome)
157                .map_err(KtoolError::Classifier),
158            KtoolOperation::Delete(key) => self
159                .delete(key)
160                .map(render_delete_outcome)
161                .map_err(KtoolError::Classifier),
162        }
163    }
164}
165
166/// Decodes one exact Ktool name and argument object without touching classifier state.
167pub fn decode_ktool(tool: &str, arguments: &Value) -> Result<KtoolCall, KtoolError> {
168    let operation = match tool {
169        IDENTIFY_TOOL => KtoolOperation::Identify(identify_request(arguments)?),
170        TRAIN_TOOL => KtoolOperation::Train(train_request(arguments)?),
171        DELETE_TOOL => KtoolOperation::Delete(delete_request(arguments)?),
172        _ => return Err(KtoolError::UnsupportedTool(tool.to_owned())),
173    };
174    Ok(KtoolCall { operation })
175}
176
177fn identify_request(arguments: &Value) -> Result<IdentifyRequest, KtoolError> {
178    let arguments =
179        serde_json::from_value::<IdentifyArguments>(arguments.clone()).map_err(|source| {
180            KtoolError::InvalidArguments {
181                tool: IDENTIFY_TOOL,
182                source,
183            }
184        })?;
185    Ok(IdentifyRequest {
186        key: arguments.key.into(),
187        cohort: arguments.cohort.into(),
188        row: arguments.row,
189        threshold: arguments.threshold,
190    })
191}
192
193fn train_request(arguments: &Value) -> Result<TrainRequest, KtoolError> {
194    let arguments =
195        serde_json::from_value::<TrainArguments>(arguments.clone()).map_err(|source| {
196            KtoolError::InvalidArguments {
197                tool: TRAIN_TOOL,
198                source,
199            }
200        })?;
201    Ok(TrainRequest {
202        key: arguments.key.into(),
203        cohort: arguments.cohort.into(),
204        row: arguments.row,
205        speaker_id: arguments.speaker_id,
206    })
207}
208
209fn delete_request(arguments: &Value) -> Result<ObservationKey, KtoolError> {
210    serde_json::from_value::<DeleteArguments>(arguments.clone())
211        .map(|arguments| arguments.key.into())
212        .map_err(|source| KtoolError::InvalidArguments {
213            tool: DELETE_TOOL,
214            source,
215        })
216}
217
218fn render_identify_outcome(outcome: IdentifyOutcome) -> String {
219    let retained = outcome.speaker_id.is_some();
220    render_json(json!({
221        "operation":"identify",
222        "speakerId":outcome.speaker_id,
223        "retained":retained,
224        "evidence":outcome.evidence.map(evidence_json),
225    }))
226}
227
228fn render_train_outcome(outcome: TrainOutcome) -> String {
229    let outcome = match outcome {
230        TrainOutcome::Added => "added",
231        TrainOutcome::Unchanged => "unchanged",
232        TrainOutcome::Corrected => "corrected",
233    };
234    render_json(json!({"operation":"train", "outcome":outcome}))
235}
236
237fn render_delete_outcome(outcome: DeleteOutcome) -> String {
238    let outcome = match outcome {
239        DeleteOutcome::Deleted => "deleted",
240        DeleteOutcome::NotFound => "not_found",
241    };
242    render_json(json!({"operation":"delete", "outcome":outcome}))
243}
244
245fn evidence_json(evidence: IdentifyEvidence) -> Value {
246    json!({
247        "best":candidate_json(evidence.best),
248        "runnerUp":evidence.runner_up.map(candidate_json),
249        "backgroundPopulationCost":evidence.background_population_cost,
250        "absoluteGap":evidence.absolute_gap,
251        "runnerUpGap":evidence.runner_up_gap,
252        "confidenceScore":evidence.confidence_score,
253    })
254}
255
256fn candidate_json(candidate: CandidateEvidence) -> Value {
257    json!({"speakerId":candidate.speaker_id, "cost":candidate.cost})
258}
259
260fn render_json(value: Value) -> String {
261    serde_json::to_string_pretty(&value).expect("classifier outcomes contain finite JSON")
262}
263
264#[cfg(test)]
265mod tests {
266    use super::*;
267
268    fn cohort() -> Value {
269        json!({
270            "provider":"google",
271            "model":"gemini-example",
272            "promptVersion":"speaker-24-1",
273            "schemaVersion":"speaker-24-1",
274            "primaryLanguage":"eng"
275        })
276    }
277
278    fn row(seed: u8) -> Value {
279        Value::Array(
280            (0..crate::FEATURE_COUNT)
281                .map(|index| json!((usize::from(seed) + index) % 100))
282                .collect(),
283        )
284    }
285
286    fn key() -> Value {
287        json!({"objectId":"AAECAwQF", "pieceIndex":3})
288    }
289
290    #[test]
291    fn strict_contract_decodes_all_three_operations() {
292        assert!(
293            decode_ktool(
294                IDENTIFY_TOOL,
295                &json!({"key":key(), "cohort":cohort(), "row":row(1), "threshold":2.5})
296            )
297            .is_ok()
298        );
299        assert!(
300            decode_ktool(
301                TRAIN_TOOL,
302                &json!({
303                    "key":key(),
304                    "cohort":cohort(),
305                    "row":row(1),
306                    "speakerId":"Full Name"
307                })
308            )
309            .is_ok()
310        );
311        assert!(decode_ktool(DELETE_TOOL, &json!({"key":key()})).is_ok());
312        assert!(decode_ktool("kcode-speaker-system/set-sample-state", &json!({})).is_err());
313    }
314
315    #[test]
316    fn strict_contract_rejects_unknown_fields_and_bad_vectors() {
317        assert!(
318            decode_ktool(DELETE_TOOL, &json!({"key":key(), "speakerId":"unexpected"})).is_err()
319        );
320        assert!(
321            decode_ktool(
322                TRAIN_TOOL,
323                &json!({
324                    "key":key(),
325                    "cohort":cohort(),
326                    "row":[1, 2],
327                    "speakerId":"Full Name"
328                })
329            )
330            .is_err()
331        );
332    }
333
334    #[test]
335    fn outcomes_remain_stable_camel_case_json() {
336        let rendered: Value = serde_json::from_str(&render_identify_outcome(IdentifyOutcome {
337            speaker_id: Some("Full Name".to_owned()),
338            evidence: Some(IdentifyEvidence {
339                best: CandidateEvidence {
340                    speaker_id: "Full Name".to_owned(),
341                    cost: -4.0,
342                },
343                runner_up: None,
344                background_population_cost: 0.0,
345                absolute_gap: 4.0,
346                runner_up_gap: None,
347                confidence_score: 4.0,
348            }),
349        }))
350        .unwrap();
351        assert_eq!(rendered["speakerId"], "Full Name");
352        assert_eq!(rendered["retained"], true);
353        assert_eq!(rendered["evidence"]["confidenceScore"], 4.0);
354        assert_eq!(
355            serde_json::from_str::<Value>(&render_train_outcome(TrainOutcome::Corrected)).unwrap(),
356            json!({"operation":"train", "outcome":"corrected"})
357        );
358        assert_eq!(
359            serde_json::from_str::<Value>(&render_delete_outcome(DeleteOutcome::NotFound)).unwrap(),
360            json!({"operation":"delete", "outcome":"not_found"})
361        );
362    }
363}