1use crate::{
2 CommitReceipt, Decision, FEATURE_COUNT, FeatureVector, Identification, Key, SampleState,
3 SampleStateRequest, SpeakerSystem, SystemError,
4};
5use serde::{Deserialize, Serialize, de::DeserializeOwned};
6use std::{error::Error, fmt};
7
8pub const IDENTIFY_KTOOL: &str = "kcode-speaker-system/identify";
9pub const SET_SAMPLE_STATE_KTOOL: &str = "kcode-speaker-system/set-sample-state";
10
11#[derive(Clone, Copy, Debug, Eq, PartialEq)]
12pub struct KtoolSpec {
13 pub name: &'static str,
14 pub description: &'static str,
15 pub input_schema: &'static str,
16}
17
18pub const KTOOLS: &[KtoolSpec] = &[
19 KtoolSpec {
20 name: IDENTIFY_KTOOL,
21 description: "Identify one frozen 24-rating speaker profile using the explicitly loaded immutable model and return decision and LLR evidence.",
22 input_schema: r#"{"cohortId":"gemini-speaker-24-normalized/1","ratings":[24 integer values in 0..=100]}"#,
23 },
24 KtoolSpec {
25 name: SET_SAMPLE_STATE_KTOOL,
26 description: "Change one successor sample to active, confirmed with a speaker ID, or retracted using provenance-bearing event and sample IDs.",
27 input_schema: r#"{"eventId":"validated key","sampleId":"store-issued sample ID","reason":"nonblank reason","state":{"status":"active"|"confirmed"|"retracted","speakerId":"required only for confirmed"}}"#,
28 },
29];
30
31#[derive(Debug)]
32pub enum KtoolError {
33 UnknownTool,
34 InvalidArguments(String),
35 ModelUnavailable,
36 Execution(SystemError),
37 Serialization(String),
38}
39
40impl KtoolError {
41 pub fn code(&self) -> &'static str {
42 match self {
43 Self::UnknownTool => "unknown_ktool",
44 Self::InvalidArguments(_) => "invalid_arguments",
45 Self::ModelUnavailable => "model_unavailable",
46 Self::Execution(_) | Self::Serialization(_) => "execution_failed",
47 }
48 }
49}
50
51impl fmt::Display for KtoolError {
52 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
53 match self {
54 Self::UnknownTool => formatter.write_str("unknown_ktool"),
55 Self::InvalidArguments(message) => {
56 write!(formatter, "invalid_arguments: {message}")
57 }
58 Self::ModelUnavailable => formatter.write_str("model_unavailable"),
59 Self::Execution(error) => write!(formatter, "execution_failed: {error}"),
60 Self::Serialization(message) => {
61 write!(formatter, "execution_failed: {message}")
62 }
63 }
64 }
65}
66
67impl Error for KtoolError {
68 fn source(&self) -> Option<&(dyn Error + 'static)> {
69 match self {
70 Self::Execution(error) => Some(error),
71 Self::UnknownTool
72 | Self::InvalidArguments(_)
73 | Self::ModelUnavailable
74 | Self::Serialization(_) => None,
75 }
76 }
77}
78
79pub fn execute(system: &SpeakerSystem, name: &str, arguments: &str) -> Result<String, KtoolError> {
80 match name {
81 IDENTIFY_KTOOL => execute_identify(system, arguments),
82 SET_SAMPLE_STATE_KTOOL => execute_set_sample_state(system, arguments),
83 _ => Err(KtoolError::UnknownTool),
84 }
85}
86
87#[derive(Deserialize)]
88#[serde(rename_all = "camelCase", deny_unknown_fields)]
89struct IdentifyArguments {
90 cohort_id: Key,
91 ratings: [u8; FEATURE_COUNT],
92}
93
94#[derive(Serialize)]
95#[serde(rename_all = "camelCase")]
96struct IdentifyResponse {
97 decision: DecisionResponse,
98 best: CandidateResponse,
99 runner_up: Option<CandidateResponse>,
100 absolute_pass: bool,
101 margin_pass: bool,
102}
103
104#[derive(Serialize)]
105#[serde(tag = "status", rename_all = "snake_case")]
106enum DecisionResponse {
107 Known {
108 #[serde(rename = "speakerId")]
109 speaker_id: Key,
110 },
111 Unknown,
112}
113
114#[derive(Serialize)]
115#[serde(rename_all = "camelCase")]
116struct CandidateResponse {
117 speaker_id: Key,
118 llr: f64,
119}
120
121fn execute_identify(system: &SpeakerSystem, arguments: &str) -> Result<String, KtoolError> {
122 let arguments: IdentifyArguments = decode_arguments(arguments)?;
123 if &arguments.cohort_id != crate::cohort_id() {
124 return Err(KtoolError::InvalidArguments(
125 "cohortId must equal gemini-speaker-24-normalized/1".into(),
126 ));
127 }
128 let features = FeatureVector::new(arguments.ratings).map_err(|error| {
129 KtoolError::InvalidArguments(format!("ratings must be integers in 0..=100: {error}"))
130 })?;
131 let identification = system
132 .identify(&arguments.cohort_id, &features)
133 .map_err(map_identify_error)?;
134 encode_response(identify_response(identification))
135}
136
137fn identify_response(identification: Identification) -> IdentifyResponse {
138 let decision = match identification.decision {
139 Decision::Known { speaker_id } => DecisionResponse::Known { speaker_id },
140 Decision::Unknown => DecisionResponse::Unknown,
141 };
142 IdentifyResponse {
143 decision,
144 best: CandidateResponse {
145 speaker_id: identification.best.speaker_id,
146 llr: identification.best.llr,
147 },
148 runner_up: identification.runner_up.map(|candidate| CandidateResponse {
149 speaker_id: candidate.speaker_id,
150 llr: candidate.llr,
151 }),
152 absolute_pass: identification.absolute_pass,
153 margin_pass: identification.margin_pass,
154 }
155}
156
157fn map_identify_error(error: SystemError) -> KtoolError {
158 match error {
159 SystemError::ModelUnavailable => KtoolError::ModelUnavailable,
160 other => KtoolError::Execution(other),
161 }
162}
163
164#[derive(Deserialize)]
165#[serde(rename_all = "camelCase", deny_unknown_fields)]
166struct SetSampleStateArguments {
167 event_id: Key,
168 sample_id: Key,
169 reason: String,
170 state: StateArgument,
171}
172
173#[derive(Deserialize)]
174#[serde(untagged)]
175enum StateArgument {
176 Active(ActiveState),
177 Confirmed(ConfirmedState),
178 Retracted(RetractedState),
179}
180
181#[derive(Deserialize)]
182#[serde(deny_unknown_fields)]
183struct ActiveState {
184 status: ActiveStatus,
185}
186
187#[derive(Deserialize)]
188#[serde(deny_unknown_fields)]
189struct ConfirmedState {
190 status: ConfirmedStatus,
191 #[serde(rename = "speakerId")]
192 speaker_id: Key,
193}
194
195#[derive(Deserialize)]
196#[serde(deny_unknown_fields)]
197struct RetractedState {
198 status: RetractedStatus,
199}
200
201#[derive(Deserialize)]
202enum ActiveStatus {
203 #[serde(rename = "active")]
204 Active,
205}
206
207#[derive(Deserialize)]
208enum ConfirmedStatus {
209 #[serde(rename = "confirmed")]
210 Confirmed,
211}
212
213#[derive(Deserialize)]
214enum RetractedStatus {
215 #[serde(rename = "retracted")]
216 Retracted,
217}
218
219#[derive(Serialize)]
220#[serde(rename_all = "camelCase")]
221struct StateResponse {
222 revision: u64,
223 applied: u64,
224}
225
226fn execute_set_sample_state(system: &SpeakerSystem, arguments: &str) -> Result<String, KtoolError> {
227 let arguments: SetSampleStateArguments = decode_arguments(arguments)?;
228 if arguments.reason.trim().is_empty() {
229 return Err(KtoolError::InvalidArguments(
230 "reason must be nonblank".into(),
231 ));
232 }
233
234 let state = match arguments.state {
235 StateArgument::Active(state) => {
236 let ActiveStatus::Active = state.status;
237 SampleState::Unlabeled
238 }
239 StateArgument::Confirmed(state) => {
240 let ConfirmedStatus::Confirmed = state.status;
241 SampleState::Confirmed {
242 speaker_id: state.speaker_id,
243 }
244 }
245 StateArgument::Retracted(state) => {
246 let RetractedStatus::Retracted = state.status;
247 SampleState::Retracted
248 }
249 };
250
251 let receipt = system
252 .change_sample_state(SampleStateRequest {
253 event_id: arguments.event_id,
254 sample_id: arguments.sample_id,
255 state,
256 reason: arguments.reason,
257 })
258 .map_err(KtoolError::Execution)?;
259 encode_receipt(receipt)
260}
261
262fn decode_arguments<T: DeserializeOwned>(arguments: &str) -> Result<T, KtoolError> {
263 serde_json::from_str(arguments).map_err(|error| KtoolError::InvalidArguments(error.to_string()))
264}
265
266fn encode_receipt(receipt: CommitReceipt) -> Result<String, KtoolError> {
267 encode_response(StateResponse {
268 revision: receipt.revision,
269 applied: receipt.applied,
270 })
271}
272
273fn encode_response<T: Serialize>(response: T) -> Result<String, KtoolError> {
274 serde_json::to_string(&response).map_err(|error| KtoolError::Serialization(error.to_string()))
275}
276
277#[cfg(test)]
278mod tests {
279 use super::*;
280 use crate::{
281 AttemptSelectionRequest, NormalizedAttempt, ObjectId, RecordingKind, SegmentBinding,
282 SourceRegistration, cohort_id, open,
283 };
284 use serde_json::{Value, json};
285 use std::{
286 fs,
287 path::PathBuf,
288 sync::atomic::{AtomicU64, Ordering},
289 };
290
291 static NEXT_PATH: AtomicU64 = AtomicU64::new(0);
292
293 fn key(value: &str) -> Key {
294 Key::parse(value).unwrap()
295 }
296
297 fn object(value: &str) -> ObjectId {
298 ObjectId::parse(value).unwrap()
299 }
300
301 fn path() -> PathBuf {
302 std::env::temp_dir().join(format!(
303 "kcode-speaker-system-ktool-{}-{}.db",
304 std::process::id(),
305 NEXT_PATH.fetch_add(1, Ordering::Relaxed)
306 ))
307 }
308
309 fn ratings() -> Vec<u8> {
310 (10..34).collect()
311 }
312
313 fn registered_sample(system: &SpeakerSystem) -> Key {
314 system
315 .register_source(SourceRegistration {
316 source_object: object("SOURCE31"),
317 source_duration_ms: 100_000,
318 group_id: key("group/source/31"),
319 recording_kind: RecordingKind::VoiceNote,
320 segments: vec![SegmentBinding {
321 event_id: key("event/register/31"),
322 clip_object: object("CLIP0031"),
323 }],
324 })
325 .unwrap();
326
327 let feature_list = ratings()
328 .into_iter()
329 .map(|rating| rating.to_string())
330 .collect::<Vec<_>>()
331 .join(",");
332 let normalized_response = format!(
333 r#"{{"status":"scored","speakers":[{{"speakerOrdinal":0,"primaryLanguage":"en-US","closestDialect":"General American English","usableSpeechMs":20000,"features":[{feature_list}]}}],"additionalSpeakers":[]}}"#
334 );
335 system
336 .record_normalized_attempt(NormalizedAttempt {
337 event_id: key("event/attempt/31"),
338 attempt_id: key("attempt/31"),
339 cohort_id: cohort_id().clone(),
340 source_object: object("SOURCE31"),
341 clip_object: object("CLIP0031"),
342 segment_ordinal: 0,
343 provider_result_object: object("RESULT31"),
344 recording_quality: Some(90),
345 normalized_response: &normalized_response,
346 })
347 .unwrap();
348 system
349 .select_attempt(AttemptSelectionRequest {
350 event_id: key("event/select/31"),
351 attempt_id: key("attempt/31"),
352 cohort_id: cohort_id().clone(),
353 clip_object: object("CLIP0031"),
354 reason: "selected complete successor sample".into(),
355 })
356 .unwrap();
357
358 system
359 .attempt(&key("attempt/31"))
360 .unwrap()
361 .unwrap()
362 .sample_ids[0]
363 .clone()
364 }
365
366 #[test]
367 fn identify_json_is_strict_and_model_unavailable_is_explicit() {
368 let database = path();
369 let system = open(&database).unwrap();
370 let valid = json!({
371 "cohortId": "gemini-speaker-24-normalized/1",
372 "ratings": ratings(),
373 });
374
375 let error = execute(&system, IDENTIFY_KTOOL, &valid.to_string()).unwrap_err();
376 assert!(matches!(error, KtoolError::ModelUnavailable));
377 assert_eq!(error.code(), "model_unavailable");
378 assert_eq!(error.to_string(), "model_unavailable");
379
380 let unknown = json!({
381 "cohortId": "gemini-speaker-24-normalized/1",
382 "ratings": ratings(),
383 "speaker": "legacy",
384 });
385 assert!(matches!(
386 execute(&system, IDENTIFY_KTOOL, &unknown.to_string()),
387 Err(KtoolError::InvalidArguments(_))
388 ));
389
390 let legacy = json!({
391 "cohort": "legacy",
392 "features": ratings(),
393 });
394 assert!(matches!(
395 execute(&system, IDENTIFY_KTOOL, &legacy.to_string()),
396 Err(KtoolError::InvalidArguments(_))
397 ));
398
399 let mut short = ratings();
400 short.pop();
401 let wrong_count = json!({
402 "cohortId": "gemini-speaker-24-normalized/1",
403 "ratings": short,
404 });
405 assert!(matches!(
406 execute(&system, IDENTIFY_KTOOL, &wrong_count.to_string()),
407 Err(KtoolError::InvalidArguments(_))
408 ));
409
410 let mut out_of_range = ratings();
411 out_of_range[0] = 101;
412 let wrong_range = json!({
413 "cohortId": "gemini-speaker-24-normalized/1",
414 "ratings": out_of_range,
415 });
416 assert!(matches!(
417 execute(&system, IDENTIFY_KTOOL, &wrong_range.to_string()),
418 Err(KtoolError::InvalidArguments(_))
419 ));
420
421 let wrong_cohort = json!({
422 "cohortId": "gemini-speaker-24-freeform/1",
423 "ratings": ratings(),
424 });
425 assert!(matches!(
426 execute(&system, IDENTIFY_KTOOL, &wrong_cohort.to_string()),
427 Err(KtoolError::InvalidArguments(_))
428 ));
429
430 drop(system);
431 fs::remove_file(database).unwrap();
432 }
433
434 #[test]
435 fn state_json_rejects_legacy_before_mutation_and_changes_successor_state() {
436 let database = path();
437 let system = open(&database).unwrap();
438 let sample_id = registered_sample(&system);
439 let event_id = "event/state/ktool/31";
440
441 let legacy = json!({
442 "eventId": event_id,
443 "sampleId": sample_id,
444 "reason": "legacy operation must fail",
445 "action": "train",
446 "speakerId": "speaker/alice",
447 });
448 assert!(matches!(
449 execute(&system, SET_SAMPLE_STATE_KTOOL, &legacy.to_string()),
450 Err(KtoolError::InvalidArguments(_))
451 ));
452
453 let confirmed = json!({
454 "eventId": event_id,
455 "sampleId": sample_id,
456 "reason": "confirmed from successor sample provenance",
457 "state": {
458 "status": "confirmed",
459 "speakerId": "speaker/alice",
460 },
461 });
462 let response = execute(&system, SET_SAMPLE_STATE_KTOOL, &confirmed.to_string()).unwrap();
463 let response: Value = serde_json::from_str(&response).unwrap();
464 assert_eq!(response["applied"], 1);
465 assert_eq!(
466 kcode_speaker_dataset::active_rows(&system.dataset(cohort_id()).unwrap()).len(),
467 1
468 );
469
470 let active = json!({
471 "eventId": "event/state/ktool/32",
472 "sampleId": sample_id,
473 "reason": "return sample to active unconfirmed review",
474 "state": {
475 "status": "active",
476 },
477 });
478 execute(&system, SET_SAMPLE_STATE_KTOOL, &active.to_string()).unwrap();
479 assert!(matches!(
480 system.dataset(cohort_id()),
481 Err(crate::SystemError::Dataset(
482 crate::DatasetError::EmptyActive
483 ))
484 ));
485
486 let retracted = json!({
487 "eventId": "event/state/ktool/33",
488 "sampleId": sample_id,
489 "reason": "retract successor sample",
490 "state": {
491 "status": "retracted",
492 },
493 });
494 execute(&system, SET_SAMPLE_STATE_KTOOL, &retracted.to_string()).unwrap();
495 assert!(matches!(
496 system.dataset(cohort_id()),
497 Err(crate::SystemError::Dataset(
498 crate::DatasetError::EmptyActive
499 ))
500 ));
501
502 let speaker_on_active = json!({
503 "eventId": "event/state/ktool/34",
504 "sampleId": sample_id,
505 "reason": "unknown nested field must fail",
506 "state": {
507 "status": "active",
508 "speakerId": "speaker/alice",
509 },
510 });
511 assert!(matches!(
512 execute(
513 &system,
514 SET_SAMPLE_STATE_KTOOL,
515 &speaker_on_active.to_string()
516 ),
517 Err(KtoolError::InvalidArguments(_))
518 ));
519
520 drop(system);
521 fs::remove_file(database).unwrap();
522 }
523
524 #[test]
525 fn only_successor_names_dispatch() {
526 let database = path();
527 let system = open(&database).unwrap();
528
529 assert_eq!(KTOOLS.len(), 2);
530 assert_eq!(KTOOLS[0].name, IDENTIFY_KTOOL);
531 assert_eq!(KTOOLS[1].name, SET_SAMPLE_STATE_KTOOL);
532 assert!(matches!(
533 execute(&system, "kcode-speech-classifier/train", "{}"),
534 Err(KtoolError::UnknownTool)
535 ));
536
537 drop(system);
538 fs::remove_file(database).unwrap();
539 }
540}