kcode-speech-classification 0.1.1

Typed open-set speaker classification with SQLite-backed training and an NDJSON Unix-socket service.
Documentation
use crate::model::{
    Cohort, DeleteOutcome, FeatureRow, IdentifyOutcome, ObservationKey, TrainOutcome,
};
use serde::{Deserialize, Serialize};

#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
#[serde(tag = "operation", rename_all = "snake_case")]
pub enum ProtocolRequest {
    Identify {
        key: ObservationKey,
        cohort: Cohort,
        row: FeatureRow,
        threshold: f64,
    },
    Train {
        key: ObservationKey,
        cohort: Cohort,
        row: FeatureRow,
        speaker_id: String,
    },
    Delete {
        key: ObservationKey,
    },
}

#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
#[serde(tag = "operation", content = "outcome", rename_all = "snake_case")]
pub enum ProtocolResult {
    Identify(IdentifyOutcome),
    Train(TrainOutcome),
    Delete(DeleteOutcome),
}

#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ProtocolError {
    pub code: String,
    pub message: String,
}

#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
#[serde(tag = "status", rename_all = "snake_case")]
pub enum ProtocolResponse {
    Success { result: ProtocolResult },
    Error { error: ProtocolError },
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::model::tests::{sample_cohort, sample_row};
    use crate::model::{CandidateEvidence, IdentifyEvidence};

    fn key() -> ObservationKey {
        ObservationKey {
            object_id: "object-17".to_owned(),
            piece_index: 3,
        }
    }

    #[test]
    fn every_request_round_trips() {
        let requests = [
            ProtocolRequest::Identify {
                key: key(),
                cohort: sample_cohort(),
                row: sample_row(),
                threshold: 2.5,
            },
            ProtocolRequest::Train {
                key: key(),
                cohort: sample_cohort(),
                row: sample_row(),
                speaker_id: "speaker-a".to_owned(),
            },
            ProtocolRequest::Delete { key: key() },
        ];

        for request in requests {
            let json = serde_json::to_string(&request).unwrap();
            let decoded: ProtocolRequest = serde_json::from_str(&json).unwrap();
            assert_eq!(decoded, request);
        }
    }

    #[test]
    fn success_and_error_responses_round_trip() {
        let evidence = IdentifyEvidence {
            best: CandidateEvidence {
                speaker_id: "speaker-a".to_owned(),
                cost: 4.0,
            },
            runner_up: Some(CandidateEvidence {
                speaker_id: "speaker-b".to_owned(),
                cost: 7.0,
            }),
            background_population_cost: 8.0,
            absolute_gap: 4.0,
            runner_up_gap: Some(3.0),
            confidence_score: 3.0,
        };
        let responses = [
            ProtocolResponse::Success {
                result: ProtocolResult::Identify(IdentifyOutcome {
                    speaker_id: Some("speaker-a".to_owned()),
                    evidence: Some(evidence),
                }),
            },
            ProtocolResponse::Success {
                result: ProtocolResult::Train(TrainOutcome::Corrected),
            },
            ProtocolResponse::Success {
                result: ProtocolResult::Delete(DeleteOutcome::NotFound),
            },
            ProtocolResponse::Error {
                error: ProtocolError {
                    code: "validation".to_owned(),
                    message: "bad input".to_owned(),
                },
            },
        ];

        for response in responses {
            let json = serde_json::to_string(&response).unwrap();
            let decoded: ProtocolResponse = serde_json::from_str(&json).unwrap();
            assert_eq!(decoded, response);
        }
    }

    #[test]
    fn request_rejects_an_unknown_operation() {
        let json = r#"{"operation":"export"}"#;
        assert!(serde_json::from_str::<ProtocolRequest>(json).is_err());
    }
}