1use 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 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
166pub 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}