Skip to main content

kcode_speaker_system/
ktool.rs

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}