use std::fmt;
use serde::Deserialize;
use serde_json::{Value, json};
use crate::{
CandidateEvidence, Cefr, Cohort, DeleteOutcome, Error, FeatureRow, IdentifyEvidence,
IdentifyOutcome, ObservationKey, SpeechClassifier, TrainOutcome,
};
pub const IDENTIFY_TOOL: &str = "kcode-speech-classification/identify";
pub const TRAIN_TOOL: &str = "kcode-speech-classification/train";
pub const DELETE_TOOL: &str = "kcode-speech-classification/delete";
pub const KTOOLS: [&str; 3] = [IDENTIFY_TOOL, TRAIN_TOOL, DELETE_TOOL];
#[derive(Debug)]
pub enum KtoolError {
InvalidArguments {
tool: &'static str,
source: serde_json::Error,
},
UnsupportedTool(String),
Classifier(Error),
}
impl fmt::Display for KtoolError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidArguments { tool, .. } => {
write!(formatter, "decoding {tool} arguments")
}
Self::UnsupportedTool(tool) => {
write!(formatter, "unsupported speech-classification Ktool {tool}")
}
Self::Classifier(error) => error.fmt(formatter),
}
}
}
impl std::error::Error for KtoolError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::InvalidArguments { source, .. } => Some(source),
Self::UnsupportedTool(_) => None,
Self::Classifier(error) => Some(error),
}
}
}
impl From<Error> for KtoolError {
fn from(value: Error) -> Self {
Self::Classifier(value)
}
}
#[derive(Debug)]
pub struct KtoolCall {
operation: KtoolOperation,
}
#[derive(Debug)]
enum KtoolOperation {
Identify(IdentifyRequest),
Train(TrainRequest),
Delete(ObservationKey),
}
#[derive(Debug)]
struct IdentifyRequest {
key: ObservationKey,
cohort: Cohort,
row: FeatureRow,
threshold: f64,
}
#[derive(Debug)]
struct TrainRequest {
key: ObservationKey,
cohort: Cohort,
row: FeatureRow,
speaker_id: String,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct IdentifyArguments {
key: ObservationKeyInput,
cohort: CohortInput,
row: FeatureRowInput,
threshold: f64,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct TrainArguments {
key: ObservationKeyInput,
cohort: CohortInput,
row: FeatureRowInput,
speaker_id: String,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct DeleteArguments {
key: ObservationKeyInput,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct ObservationKeyInput {
object_id: String,
piece_index: u32,
}
impl From<ObservationKeyInput> for ObservationKey {
fn from(value: ObservationKeyInput) -> Self {
Self {
object_id: value.object_id,
piece_index: value.piece_index,
}
}
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct CohortInput {
provider: String,
model: String,
prompt_version: String,
schema_version: String,
primary_language: String,
}
impl From<CohortInput> for Cohort {
fn from(value: CohortInput) -> Self {
Self {
provider: value.provider,
model: value.model,
prompt_version: value.prompt_version,
schema_version: value.schema_version,
primary_language: value.primary_language,
}
}
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct FeatureRowInput {
accent_variety: String,
perceived_age: f64,
vocal_gender_presentation: f64,
median_f0_hz: f64,
formant_dispersion_hz: f64,
vai: f64,
hypernasality: f64,
creaky_phonation_percent: f64,
rhotic_realization: String,
word_initial_stressed_prevocalic_t_vot_ms: f64,
breathiness: f64,
roughness: f64,
f0_pitch_span_semitones: f64,
articulation_rate_syllables_per_second: f64,
npvi_v: f64,
cefr: Cefr,
foreign_accentedness: f64,
unstressed_vowel_reduction_percent: f64,
lateral_realization: String,
filled_pauses_per_100_words: f64,
s_realization: String,
lexical_stress_accuracy_percent: f64,
monophthongization_percent: f64,
consonant_cluster_reduction_percent: f64,
}
impl From<FeatureRowInput> for FeatureRow {
fn from(value: FeatureRowInput) -> Self {
Self {
accent_variety: value.accent_variety,
perceived_age: value.perceived_age,
vocal_gender_presentation: value.vocal_gender_presentation,
median_f0_hz: value.median_f0_hz,
formant_dispersion_hz: value.formant_dispersion_hz,
vai: value.vai,
hypernasality: value.hypernasality,
creaky_phonation_percent: value.creaky_phonation_percent,
rhotic_realization: value.rhotic_realization,
word_initial_stressed_prevocalic_t_vot_ms: value
.word_initial_stressed_prevocalic_t_vot_ms,
breathiness: value.breathiness,
roughness: value.roughness,
f0_pitch_span_semitones: value.f0_pitch_span_semitones,
articulation_rate_syllables_per_second: value.articulation_rate_syllables_per_second,
npvi_v: value.npvi_v,
cefr: value.cefr,
foreign_accentedness: value.foreign_accentedness,
unstressed_vowel_reduction_percent: value.unstressed_vowel_reduction_percent,
lateral_realization: value.lateral_realization,
filled_pauses_per_100_words: value.filled_pauses_per_100_words,
s_realization: value.s_realization,
lexical_stress_accuracy_percent: value.lexical_stress_accuracy_percent,
monophthongization_percent: value.monophthongization_percent,
consonant_cluster_reduction_percent: value.consonant_cluster_reduction_percent,
}
}
}
impl SpeechClassifier {
pub fn execute_ktool(&self, call: KtoolCall) -> Result<String, KtoolError> {
match call.operation {
KtoolOperation::Identify(request) => self
.identify(request.key, request.cohort, request.row, request.threshold)
.map(render_identify_outcome)
.map_err(KtoolError::Classifier),
KtoolOperation::Train(request) => self
.train(request.key, request.cohort, request.row, request.speaker_id)
.map(render_train_outcome)
.map_err(KtoolError::Classifier),
KtoolOperation::Delete(key) => self
.delete(key)
.map(render_delete_outcome)
.map_err(KtoolError::Classifier),
}
}
}
pub fn decode_ktool(tool: &str, arguments: &Value) -> Result<KtoolCall, KtoolError> {
let operation = match tool {
IDENTIFY_TOOL => KtoolOperation::Identify(identify_request(arguments)?),
TRAIN_TOOL => KtoolOperation::Train(train_request(arguments)?),
DELETE_TOOL => KtoolOperation::Delete(delete_request(arguments)?),
_ => return Err(KtoolError::UnsupportedTool(tool.to_owned())),
};
Ok(KtoolCall { operation })
}
fn identify_request(arguments: &Value) -> Result<IdentifyRequest, KtoolError> {
let arguments =
serde_json::from_value::<IdentifyArguments>(arguments.clone()).map_err(|source| {
KtoolError::InvalidArguments {
tool: IDENTIFY_TOOL,
source,
}
})?;
Ok(IdentifyRequest {
key: arguments.key.into(),
cohort: arguments.cohort.into(),
row: arguments.row.into(),
threshold: arguments.threshold,
})
}
fn train_request(arguments: &Value) -> Result<TrainRequest, KtoolError> {
let arguments =
serde_json::from_value::<TrainArguments>(arguments.clone()).map_err(|source| {
KtoolError::InvalidArguments {
tool: TRAIN_TOOL,
source,
}
})?;
Ok(TrainRequest {
key: arguments.key.into(),
cohort: arguments.cohort.into(),
row: arguments.row.into(),
speaker_id: arguments.speaker_id,
})
}
fn delete_request(arguments: &Value) -> Result<ObservationKey, KtoolError> {
serde_json::from_value::<DeleteArguments>(arguments.clone())
.map(|arguments| arguments.key.into())
.map_err(|source| KtoolError::InvalidArguments {
tool: DELETE_TOOL,
source,
})
}
fn render_identify_outcome(outcome: IdentifyOutcome) -> String {
let retained = outcome.speaker_id.is_some();
render_json(json!({
"operation":"identify",
"speakerId":outcome.speaker_id,
"retained":retained,
"evidence":outcome.evidence.map(evidence_json),
}))
}
fn render_train_outcome(outcome: TrainOutcome) -> String {
let outcome = match outcome {
TrainOutcome::Added => "added",
TrainOutcome::Unchanged => "unchanged",
TrainOutcome::Corrected => "corrected",
};
render_json(json!({"operation":"train", "outcome":outcome}))
}
fn render_delete_outcome(outcome: DeleteOutcome) -> String {
let outcome = match outcome {
DeleteOutcome::Deleted => "deleted",
DeleteOutcome::NotFound => "not_found",
};
render_json(json!({"operation":"delete", "outcome":outcome}))
}
fn evidence_json(evidence: IdentifyEvidence) -> Value {
json!({
"best":candidate_json(evidence.best),
"runnerUp":evidence.runner_up.map(candidate_json),
"backgroundPopulationCost":evidence.background_population_cost,
"absoluteGap":evidence.absolute_gap,
"runnerUpGap":evidence.runner_up_gap,
"confidenceScore":evidence.confidence_score,
})
}
fn candidate_json(candidate: CandidateEvidence) -> Value {
json!({"speakerId":candidate.speaker_id, "cost":candidate.cost})
}
fn render_json(value: Value) -> String {
serde_json::to_string_pretty(&value)
.expect("speaker-classification outcomes contain finite JSON")
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT_PATH: AtomicU64 = AtomicU64::new(0);
fn row() -> Value {
json!({
"accentVariety":"General American",
"perceivedAge":36.0,
"vocalGenderPresentation":45.0,
"medianF0Hz":155.0,
"formantDispersionHz":1100.0,
"vai":1.1,
"hypernasality":0.5,
"creakyPhonationPercent":8.0,
"rhoticRealization":"rhotic",
"wordInitialStressedPrevocalicTVotMs":62.0,
"breathiness":20.0,
"roughness":10.0,
"f0PitchSpanSemitones":9.0,
"articulationRateSyllablesPerSecond":4.2,
"npviV":48.0,
"cefr":"C1",
"foreignAccentedness":2.0,
"unstressedVowelReductionPercent":72.0,
"lateralRealization":"alveolar",
"filledPausesPer100Words":2.5,
"sRealization":"alveolar",
"lexicalStressAccuracyPercent":92.0,
"monophthongizationPercent":4.0,
"consonantClusterReductionPercent":3.0
})
}
fn speaker_row(age: f64, accent: &str) -> Value {
let mut row = row();
row["perceivedAge"] = json!(age);
row["medianF0Hz"] = json!(120.0 + age);
row["accentVariety"] = json!(accent);
row["rhoticRealization"] = json!(format!("{accent}-rhotic"));
row
}
fn cohort() -> Value {
json!({
"provider":"google",
"model":"gemini-example",
"promptVersion":"speaker-features-1",
"schemaVersion":"features-1",
"primaryLanguage":"eng"
})
}
fn key() -> Value {
json!({"objectId":"AAECAwQF", "pieceIndex":3})
}
fn execute(classifier: &SpeechClassifier, tool: &str, arguments: Value) -> String {
let call = decode_ktool(tool, &arguments).unwrap();
classifier.execute_ktool(call).unwrap()
}
#[test]
fn camel_case_tool_contract_maps_every_classifier_field() {
let identify = identify_request(&json!({
"key":key(),
"cohort":cohort(),
"row":row(),
"threshold":2.5
}))
.unwrap();
assert_eq!(identify.key.object_id, "AAECAwQF");
assert_eq!(identify.key.piece_index, 3);
assert_eq!(identify.cohort.prompt_version, "speaker-features-1");
assert_eq!(identify.row.cefr, Cefr::C1);
assert_eq!(identify.row.consonant_cluster_reduction_percent, 3.0);
assert_eq!(identify.threshold, 2.5);
let train = train_request(&json!({
"key":key(),
"cohort":cohort(),
"row":row(),
"speakerId":"kennedy"
}))
.unwrap();
assert_eq!(train.speaker_id, "kennedy");
let delete = delete_request(&json!({"key":key()})).unwrap();
assert_eq!(delete.object_id, "AAECAwQF");
assert_eq!(delete.piece_index, 3);
}
#[test]
fn tool_contract_rejects_unknown_fields_at_every_level() {
let error = delete_request(&json!({"key":key(), "speakerId":"unexpected"})).unwrap_err();
assert!(
error
.to_string()
.contains("kcode-speech-classification/delete")
);
let mut row = row();
row["unexpected"] = json!(true);
let error = identify_request(&json!({
"key":key(),
"cohort":cohort(),
"row":row,
"threshold":2.5
}))
.unwrap_err();
assert!(
error
.to_string()
.contains("kcode-speech-classification/identify")
);
}
#[test]
fn outcomes_are_rendered_as_stable_camel_case_json() {
let rendered = render_identify_outcome(IdentifyOutcome {
speaker_id: Some("speaker-a".into()),
evidence: Some(IdentifyEvidence {
best: CandidateEvidence {
speaker_id: "speaker-a".into(),
cost: 4.0,
},
runner_up: None,
background_population_cost: 8.0,
absolute_gap: 4.0,
runner_up_gap: None,
confidence_score: 4.0,
}),
});
let rendered: Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(rendered["speakerId"], "speaker-a");
assert_eq!(rendered["retained"], true);
assert_eq!(rendered["evidence"]["confidenceScore"], 4.0);
assert!(rendered["evidence"].get("confidence_score").is_none());
let rendered: Value =
serde_json::from_str(&render_train_outcome(TrainOutcome::Corrected)).unwrap();
assert_eq!(
rendered,
json!({"operation":"train", "outcome":"corrected"})
);
let rendered: Value =
serde_json::from_str(&render_delete_outcome(DeleteOutcome::NotFound)).unwrap();
assert_eq!(
rendered,
json!({"operation":"delete", "outcome":"not_found"})
);
}
#[test]
fn complete_tool_workflow_preserves_labels_identification_and_deletion() {
let path = std::env::temp_dir().join(format!(
"kennedy-speech-classification-tool-test-{}-{}.sqlite3",
std::process::id(),
NEXT_PATH.fetch_add(1, Ordering::Relaxed)
));
let classifier = SpeechClassifier::open(&path).unwrap();
for (object_id, age, accent, speaker_id) in [
("alpha-1", 30.0, "alpha", "speaker-a"),
("alpha-2", 32.0, "alpha", "speaker-a"),
("beta-1", 70.0, "beta", "speaker-b"),
("beta-2", 72.0, "beta", "speaker-b"),
] {
let rendered = execute(
&classifier,
TRAIN_TOOL,
json!({
"key":{"objectId":object_id, "pieceIndex":0},
"cohort":cohort(),
"row":speaker_row(age, accent),
"speakerId":speaker_id
}),
);
let rendered: Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(rendered["outcome"], "added");
}
let rendered = execute(
&classifier,
IDENTIFY_TOOL,
json!({
"key":{"objectId":"query", "pieceIndex":0},
"cohort":cohort(),
"row":speaker_row(31.0, "alpha"),
"threshold":-1_000_000.0
}),
);
let rendered: Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(rendered["speakerId"], "speaker-a");
assert_eq!(rendered["retained"], true);
let rendered = execute(
&classifier,
DELETE_TOOL,
json!({"key":{"objectId":"query", "pieceIndex":0}}),
);
let rendered: Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(rendered["outcome"], "deleted");
drop(classifier);
std::fs::remove_file(path).unwrap();
}
}