Skip to main content

openkind_core/
answer.rs

1//! Answer types — one per Question type.
2//!
3//! Spec: <https://docs.typesafe.ai/api#answer-types>
4//! > Every answer carries a `type` matching its question. Choice and Score
5//! > answers also carry a `confidence` between 0 to 1, derived from the
6//! > answer's probability distribution.
7
8use 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/// Tagged union of answer models matching the evaluated question types.
16#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
17#[serde(tag = "type", rename_all = "lowercase")]
18pub enum Answer {
19    /// Boolean probability answer containing a single probability value.
20    Noul(NoulAnswer),
21    /// Categorical choice answer containing selected label, probability distribution, and confidence.
22    Choice(ChoiceAnswer),
23    /// Ordinal rating answer containing the evaluated score, rubric legend, probabilities, and confidence.
24    Score(ScoreAnswer),
25}
26
27// Streaming deserialization for the same reason as `Question`: serde's
28// internally-tagged derive buffers every answer into `Content` before
29// picking a variant, which showed up as a measurable share of client-side
30// response decoding. The visitor reads `type` and deserializes the rest of
31// the fields straight into the selected variant; fields arriving before the
32// tag are buffered as JSON values and converted once the tag is known.
33impl<'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                // Fields seen before the tag, keyed in arrival order.
59                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                    // Known fields from other variants are still unknown
95                    // fields for this variant, so their values are ignored.
96                    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                // Replay any fields that arrived before the tag.
127                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/// Spec: `noul` is a number 0..1. **No `confidence` field** — Noul answers
205/// don't get one per the spec.
206///
207/// `f64` so wire-format precision matches the API spec exactly (a 0.92
208/// going through f32 round-trips to 0.9200000166893005 — wrong).
209#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
210pub struct NoulAnswer {
211    /// Probability value in `[0.0, 1.0]` representing the likelihood that the answer is true/yes.
212    pub noul: f64,
213}
214
215/// `probabilities` is a full distribution (sums to 1) over the criteria keys.
216/// `confidence` is required (per spec). `f64` for wire-format precision.
217#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
218pub struct ChoiceAnswer {
219    /// Selected option identifier matching one of the keys in question criteria.
220    pub choice: String,
221    /// Normalized probability distribution over all criteria options summing to 1.0.
222    pub probabilities: HashMap<String, f64>,
223    /// Model confidence score in `[0.0, 1.0]` derived from the probability distribution.
224    pub confidence: f64,
225}
226
227/// `legend` is level-index (string) → level description.
228/// `probabilities` is keyed by the same index strings. `confidence` required.
229/// `f64` for wire-format precision.
230#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
231pub struct ScoreAnswer {
232    /// Inferred expected score value along the ordinal rubric.
233    pub score: f64,
234    /// Mapping of numeric level index strings (`"0"`, `"1"`, ...) to rubric descriptions.
235    pub legend: HashMap<String, String>,
236    /// Normalized probability distribution over the level indices summing to 1.0.
237    pub probabilities: HashMap<String, f64>,
238    /// Model confidence score in `[0.0, 1.0]` derived from the probability distribution.
239    pub confidence: f64,
240}