use std::borrow::Cow;
use std::collections::HashMap;
use schemars::JsonSchema;
use serde::de::{Error as _, IgnoredAny, MapAccess, Visitor};
use serde::{Deserialize, Deserializer, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum Answer {
Noul(NoulAnswer),
Choice(ChoiceAnswer),
Score(ScoreAnswer),
}
impl<'de> Deserialize<'de> for Answer {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct AnswerVisitor;
impl<'de> Visitor<'de> for AnswerVisitor {
type Value = Answer;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a tagged answer object")
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut tag: Option<Cow<'de, str>> = None;
let mut noul: Option<f64> = None;
let mut choice: Option<String> = None;
let mut probabilities: Option<HashMap<String, f64>> = None;
let mut confidence: Option<f64> = None;
let mut score: Option<f64> = None;
let mut legend: Option<HashMap<String, String>> = None;
let mut early: Vec<(&'static str, serde_json::Value)> = Vec::new();
while let Some(key) = map.next_key::<Cow<'de, str>>()? {
let field: &'static str = match key.as_ref() {
"type" => {
if tag.is_some() {
return Err(A::Error::duplicate_field("type"));
}
let value = map.next_value::<Cow<'de, str>>()?;
tag = Some(match value.as_ref() {
"noul" | "choice" | "score" => value,
other => {
return Err(A::Error::unknown_variant(
other,
&["noul", "choice", "score"],
))
}
});
continue;
}
"noul" => "noul",
"choice" => "choice",
"probabilities" => "probabilities",
"confidence" => "confidence",
"score" => "score",
"legend" => "legend",
_ => {
let _ = map.next_value::<IgnoredAny>()?;
continue;
}
};
let Some(answer_tag) = tag.as_deref() else {
early.push((field, map.next_value()?));
continue;
};
if !field_belongs_to_answer(answer_tag, field) {
let _ = map.next_value::<IgnoredAny>()?;
continue;
}
let already = |early: &[(&'static str, serde_json::Value)]| {
early.iter().any(|(name, _)| *name == field)
};
let duplicate = match field {
"noul" => noul.is_some(),
"choice" => choice.is_some(),
"probabilities" => probabilities.is_some(),
"confidence" => confidence.is_some(),
"score" => score.is_some(),
"legend" => legend.is_some(),
_ => false,
} || already(&early);
if duplicate {
return Err(A::Error::duplicate_field(field));
}
match field {
"noul" => noul = Some(map.next_value()?),
"choice" => choice = Some(map.next_value()?),
"probabilities" => probabilities = Some(map.next_value()?),
"confidence" => confidence = Some(map.next_value()?),
"score" => score = Some(map.next_value()?),
"legend" => legend = Some(map.next_value()?),
_ => unreachable!("field matched the list above"),
}
}
let tag = tag.ok_or_else(|| A::Error::missing_field("type"))?;
let mut replayed_fields = Vec::new();
for (field, value) in early {
if !field_belongs_to_answer(&tag, field) {
continue;
}
if replayed_fields.contains(&field) {
return Err(A::Error::duplicate_field(field));
}
replayed_fields.push(field);
let converted: serde_json::Value = value;
match field {
"noul" => {
noul =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
"choice" => {
choice =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
"probabilities" => {
probabilities =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
"confidence" => {
confidence =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
"score" => {
score =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
"legend" => {
legend =
Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
}
_ => unreachable!("field matched the list above"),
}
}
match tag.as_ref() {
"noul" => Ok(Answer::Noul(NoulAnswer {
noul: noul.ok_or_else(|| A::Error::missing_field("noul"))?,
})),
"choice" => Ok(Answer::Choice(ChoiceAnswer {
choice: choice.ok_or_else(|| A::Error::missing_field("choice"))?,
probabilities: probabilities
.ok_or_else(|| A::Error::missing_field("probabilities"))?,
confidence: confidence
.ok_or_else(|| A::Error::missing_field("confidence"))?,
})),
"score" => Ok(Answer::Score(ScoreAnswer {
score: score.ok_or_else(|| A::Error::missing_field("score"))?,
legend: legend.ok_or_else(|| A::Error::missing_field("legend"))?,
probabilities: probabilities
.ok_or_else(|| A::Error::missing_field("probabilities"))?,
confidence: confidence
.ok_or_else(|| A::Error::missing_field("confidence"))?,
})),
_ => unreachable!("tag validated against the variant list above"),
}
}
}
deserializer.deserialize_map(AnswerVisitor)
}
}
fn field_belongs_to_answer(tag: &str, field: &str) -> bool {
match tag {
"noul" => field == "noul",
"choice" => matches!(field, "choice" | "probabilities" | "confidence"),
"score" => matches!(field, "score" | "legend" | "probabilities" | "confidence"),
_ => false,
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
pub struct NoulAnswer {
pub noul: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
pub struct ChoiceAnswer {
pub choice: String,
pub probabilities: HashMap<String, f64>,
pub confidence: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
pub struct ScoreAnswer {
pub score: f64,
pub legend: HashMap<String, String>,
pub probabilities: HashMap<String, f64>,
pub confidence: f64,
}