1use std::borrow::Cow;
9use std::collections::HashMap;
10
11use schemars::JsonSchema;
12use serde::de::{Error as _, IgnoredAny, MapAccess, Visitor};
13use serde::{Deserialize, Deserializer, Serialize};
14
15#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
17#[serde(tag = "type", rename_all = "lowercase")]
18pub enum Answer {
19 Noul(NoulAnswer),
21 Choice(ChoiceAnswer),
23 Score(ScoreAnswer),
25}
26
27impl<'de> Deserialize<'de> for Answer {
34 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
35 where
36 D: Deserializer<'de>,
37 {
38 struct AnswerVisitor;
39
40 impl<'de> Visitor<'de> for AnswerVisitor {
41 type Value = Answer;
42
43 fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
44 formatter.write_str("a tagged answer object")
45 }
46
47 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
48 where
49 A: MapAccess<'de>,
50 {
51 let mut tag: Option<Cow<'de, str>> = None;
52 let mut noul: Option<f64> = None;
53 let mut choice: Option<String> = None;
54 let mut probabilities: Option<HashMap<String, f64>> = None;
55 let mut confidence: Option<f64> = None;
56 let mut score: Option<f64> = None;
57 let mut legend: Option<HashMap<String, String>> = None;
58 let mut early: Vec<(&'static str, serde_json::Value)> = Vec::new();
60
61 while let Some(key) = map.next_key::<Cow<'de, str>>()? {
62 let field: &'static str = match key.as_ref() {
63 "type" => {
64 if tag.is_some() {
65 return Err(A::Error::duplicate_field("type"));
66 }
67 let value = map.next_value::<Cow<'de, str>>()?;
68 tag = Some(match value.as_ref() {
69 "noul" | "choice" | "score" => value,
70 other => {
71 return Err(A::Error::unknown_variant(
72 other,
73 &["noul", "choice", "score"],
74 ))
75 }
76 });
77 continue;
78 }
79 "noul" => "noul",
80 "choice" => "choice",
81 "probabilities" => "probabilities",
82 "confidence" => "confidence",
83 "score" => "score",
84 "legend" => "legend",
85 _ => {
86 let _ = map.next_value::<IgnoredAny>()?;
87 continue;
88 }
89 };
90 let Some(answer_tag) = tag.as_deref() else {
91 early.push((field, map.next_value()?));
92 continue;
93 };
94 if !field_belongs_to_answer(answer_tag, field) {
97 let _ = map.next_value::<IgnoredAny>()?;
98 continue;
99 }
100 let already = |early: &[(&'static str, serde_json::Value)]| {
101 early.iter().any(|(name, _)| *name == field)
102 };
103 let duplicate = match field {
104 "noul" => noul.is_some(),
105 "choice" => choice.is_some(),
106 "probabilities" => probabilities.is_some(),
107 "confidence" => confidence.is_some(),
108 "score" => score.is_some(),
109 "legend" => legend.is_some(),
110 _ => false,
111 } || already(&early);
112 if duplicate {
113 return Err(A::Error::duplicate_field(field));
114 }
115 match field {
116 "noul" => noul = Some(map.next_value()?),
117 "choice" => choice = Some(map.next_value()?),
118 "probabilities" => probabilities = Some(map.next_value()?),
119 "confidence" => confidence = Some(map.next_value()?),
120 "score" => score = Some(map.next_value()?),
121 "legend" => legend = Some(map.next_value()?),
122 _ => unreachable!("field matched the list above"),
123 }
124 }
125
126 let tag = tag.ok_or_else(|| A::Error::missing_field("type"))?;
128 let mut replayed_fields = Vec::new();
129 for (field, value) in early {
130 if !field_belongs_to_answer(&tag, field) {
131 continue;
132 }
133 if replayed_fields.contains(&field) {
134 return Err(A::Error::duplicate_field(field));
135 }
136 replayed_fields.push(field);
137 let converted: serde_json::Value = value;
138 match field {
139 "noul" => {
140 noul =
141 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
142 }
143 "choice" => {
144 choice =
145 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
146 }
147 "probabilities" => {
148 probabilities =
149 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
150 }
151 "confidence" => {
152 confidence =
153 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
154 }
155 "score" => {
156 score =
157 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
158 }
159 "legend" => {
160 legend =
161 Some(serde_json::from_value(converted).map_err(A::Error::custom)?)
162 }
163 _ => unreachable!("field matched the list above"),
164 }
165 }
166
167 match tag.as_ref() {
168 "noul" => Ok(Answer::Noul(NoulAnswer {
169 noul: noul.ok_or_else(|| A::Error::missing_field("noul"))?,
170 })),
171 "choice" => Ok(Answer::Choice(ChoiceAnswer {
172 choice: choice.ok_or_else(|| A::Error::missing_field("choice"))?,
173 probabilities: probabilities
174 .ok_or_else(|| A::Error::missing_field("probabilities"))?,
175 confidence: confidence
176 .ok_or_else(|| A::Error::missing_field("confidence"))?,
177 })),
178 "score" => Ok(Answer::Score(ScoreAnswer {
179 score: score.ok_or_else(|| A::Error::missing_field("score"))?,
180 legend: legend.ok_or_else(|| A::Error::missing_field("legend"))?,
181 probabilities: probabilities
182 .ok_or_else(|| A::Error::missing_field("probabilities"))?,
183 confidence: confidence
184 .ok_or_else(|| A::Error::missing_field("confidence"))?,
185 })),
186 _ => unreachable!("tag validated against the variant list above"),
187 }
188 }
189 }
190
191 deserializer.deserialize_map(AnswerVisitor)
192 }
193}
194
195fn field_belongs_to_answer(tag: &str, field: &str) -> bool {
196 match tag {
197 "noul" => field == "noul",
198 "choice" => matches!(field, "choice" | "probabilities" | "confidence"),
199 "score" => matches!(field, "score" | "legend" | "probabilities" | "confidence"),
200 _ => false,
201 }
202}
203
204#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
210pub struct NoulAnswer {
211 pub noul: f64,
213}
214
215#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
218pub struct ChoiceAnswer {
219 pub choice: String,
221 pub probabilities: HashMap<String, f64>,
223 pub confidence: f64,
225}
226
227#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
231pub struct ScoreAnswer {
232 pub score: f64,
234 pub legend: HashMap<String, String>,
236 pub probabilities: HashMap<String, f64>,
238 pub confidence: f64,
240}