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());
}
}