1use 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
13pub const IDENTIFY_TOOL: &str = "kcode-speech-classification/identify";
15pub const TRAIN_TOOL: &str = "kcode-speech-classification/train";
17pub const DELETE_TOOL: &str = "kcode-speech-classification/delete";
19pub const KTOOLS: [&str; 3] = [IDENTIFY_TOOL, TRAIN_TOOL, DELETE_TOOL];
21
22#[derive(Debug)]
24pub enum KtoolError {
25 InvalidArguments {
27 tool: &'static str,
29 source: serde_json::Error,
31 },
32 UnsupportedTool(String),
34 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#[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 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
240pub 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}