kcode-speaker-system 0.2.0

Small in-process speaker classifier with identify, train, and delete operations
Documentation
//! Stable Kennedy Ktool contract for the in-process speaker classifier.

use std::fmt;

use serde::Deserialize;
use serde_json::{Value, json};

use crate::{
    CandidateEvidence, Cohort, DeleteOutcome, Error, FeatureRow, IdentifyEvidence, IdentifyOutcome,
    ObservationKey, SpeechClassifier, TrainOutcome,
};

pub const IDENTIFY_TOOL: &str = "kcode-speaker-system/identify";
pub const TRAIN_TOOL: &str = "kcode-speaker-system/train";
pub const DELETE_TOOL: &str = "kcode-speaker-system/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 speaker-system 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: FeatureRow,
    threshold: f64,
}

#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct TrainArguments {
    key: ObservationKeyInput,
    cohort: CohortInput,
    row: FeatureRow,
    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,
        }
    }
}

impl SpeechClassifier {
    /// Executes and renders one already decoded classifier operation.
    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),
        }
    }
}

/// Decodes one exact Ktool name and argument object without touching classifier state.
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,
        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,
        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("classifier outcomes contain finite JSON")
}

#[cfg(test)]
mod tests {
    use super::*;

    fn cohort() -> Value {
        json!({
            "provider":"google",
            "model":"gemini-example",
            "promptVersion":"speaker-24-1",
            "schemaVersion":"speaker-24-1",
            "primaryLanguage":"eng"
        })
    }

    fn row(seed: u8) -> Value {
        Value::Array(
            (0..crate::FEATURE_COUNT)
                .map(|index| json!((usize::from(seed) + index) % 100))
                .collect(),
        )
    }

    fn key() -> Value {
        json!({"objectId":"AAECAwQF", "pieceIndex":3})
    }

    #[test]
    fn strict_contract_decodes_all_three_operations() {
        assert!(
            decode_ktool(
                IDENTIFY_TOOL,
                &json!({"key":key(), "cohort":cohort(), "row":row(1), "threshold":2.5})
            )
            .is_ok()
        );
        assert!(
            decode_ktool(
                TRAIN_TOOL,
                &json!({
                    "key":key(),
                    "cohort":cohort(),
                    "row":row(1),
                    "speakerId":"Full Name"
                })
            )
            .is_ok()
        );
        assert!(decode_ktool(DELETE_TOOL, &json!({"key":key()})).is_ok());
        assert!(decode_ktool("kcode-speaker-system/set-sample-state", &json!({})).is_err());
    }

    #[test]
    fn strict_contract_rejects_unknown_fields_and_bad_vectors() {
        assert!(
            decode_ktool(DELETE_TOOL, &json!({"key":key(), "speakerId":"unexpected"})).is_err()
        );
        assert!(
            decode_ktool(
                TRAIN_TOOL,
                &json!({
                    "key":key(),
                    "cohort":cohort(),
                    "row":[1, 2],
                    "speakerId":"Full Name"
                })
            )
            .is_err()
        );
    }

    #[test]
    fn outcomes_remain_stable_camel_case_json() {
        let rendered: Value = serde_json::from_str(&render_identify_outcome(IdentifyOutcome {
            speaker_id: Some("Full Name".to_owned()),
            evidence: Some(IdentifyEvidence {
                best: CandidateEvidence {
                    speaker_id: "Full Name".to_owned(),
                    cost: -4.0,
                },
                runner_up: None,
                background_population_cost: 0.0,
                absolute_gap: 4.0,
                runner_up_gap: None,
                confidence_score: 4.0,
            }),
        }))
        .unwrap();
        assert_eq!(rendered["speakerId"], "Full Name");
        assert_eq!(rendered["retained"], true);
        assert_eq!(rendered["evidence"]["confidenceScore"], 4.0);
        assert_eq!(
            serde_json::from_str::<Value>(&render_train_outcome(TrainOutcome::Corrected)).unwrap(),
            json!({"operation":"train", "outcome":"corrected"})
        );
        assert_eq!(
            serde_json::from_str::<Value>(&render_delete_outcome(DeleteOutcome::NotFound)).unwrap(),
            json!({"operation":"delete", "outcome":"not_found"})
        );
    }
}