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 {
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),
}
}
}
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"})
);
}
}