jevcrates.git / jev-protocol / src / response.rs
1//! A 200 response, verified against the questions that were asked
2//! before any answer is readable.
3
4use std::collections::HashMap;
5
6use serde::Deserialize;
7
8use crate::answer::{Answer, FromAnswer};
9use crate::{ChoiceAnswer, Key, ModelId, NoulAnswer, ProtocolError, Question, Questions, ScoreAnswer, Usage};
10
11/// Answers aligned with the `Questions` they answer.
12#[derive(Debug)]
13pub struct Response {
14    set: u64,
15    model: String,
16    answers: Vec<Answer>,
17    usage: Usage,
18}
19
20#[derive(Deserialize)]
21struct Wire {
22    model: String,
23    answers: HashMap<String, WireAnswer>,
24    usage: Usage,
25}
26
27#[derive(Deserialize)]
28#[serde(tag = "type", rename_all = "lowercase")]
29enum WireAnswer {
30    Noul { noul: f64 },
31    Choice { choice: String, confidence: f64, probabilities: HashMap<String, f64> },
32    Score { score: f64, confidence: f64, legend: HashMap<String, String>, probabilities: HashMap<String, f64> },
33}
34
35impl Response {
36    /// Parses a 200 body and checks it answers exactly `questions`, from
37    /// the pinned model. A response from another model is refused, so a
38    /// moved model is never recorded under the pinned version.
39    pub fn parse(pinned: &ModelId, questions: &Questions, body: &[u8]) -> Result<Self, ProtocolError> {
40        let bad = |why: String| ProtocolError::Response(why);
41        let mut wire: Wire = serde_json::from_slice(body).map_err(|e| bad(e.to_string()))?;
42        if wire.model != pinned.as_str() {
43            return Err(ProtocolError::ModelMismatch { pinned: pinned.to_string(), answered: wire.model });
44        }
45        let mut answers = Vec::with_capacity(questions.entries.len());
46        for (id, question) in &questions.entries {
47            let answer = wire.answers.remove(id).ok_or_else(|| bad(format!("no answer to {id:?}")))?;
48            answers.push(verify(id, question, answer).map_err(bad)?);
49        }
50        if let Some(extra) = wire.answers.keys().next() {
51            return Err(bad(format!("an answer to {extra:?}, which was not asked")));
52        }
53        Ok(Response { set: questions.set, model: wire.model, answers, usage: wire.usage })
54    }
55
56    /// The answer `key` names. Panics on a key from other `Questions`,
57    /// which is a bug in the caller, not a property of the response.
58    pub fn get<A: FromAnswer>(&self, key: Key<A>) -> &A {
59        assert_eq!(key.set, self.set, "a Key from other Questions");
60        A::from_answer(&self.answers[key.index]).expect("verified to match its question")
61    }
62
63    pub fn model(&self) -> &str {
64        &self.model
65    }
66
67    pub fn usage(&self) -> Usage {
68        self.usage
69    }
70}
71
72fn verify(id: &str, question: &Question, answer: WireAnswer) -> Result<Answer, String> {
73    let unit = |what: &str, p: f64| {
74        if (0.0..=1.0).contains(&p) { Ok(p) } else { Err(format!("{id:?}: {what} {p} is outside [0, 1]")) }
75    };
76    match (question, answer) {
77        (Question::Noul(_), WireAnswer::Noul { noul }) => {
78            Ok(Answer::Noul(NoulAnswer { noul: unit("noul", noul)? }))
79        }
80        (Question::Choice(q), WireAnswer::Choice { choice, confidence, mut probabilities }) => {
81            if !q.labels().any(|l| l == choice) {
82                return Err(format!("{id:?}: chose {choice:?}, which was not an option"));
83            }
84            let ordered = q
85                .labels()
86                .map(|label| {
87                    let p = probabilities.remove(label).ok_or_else(|| format!("{id:?}: no probability for {label:?}"))?;
88                    Ok((label.to_owned(), unit("probability", p)?))
89                })
90                .collect::<Result<Vec<_>, String>>()?;
91            if let Some(extra) = probabilities.keys().next() {
92                return Err(format!("{id:?}: a probability for {extra:?}, which was not an option"));
93            }
94            Ok(Answer::Choice(ChoiceAnswer { choice, confidence: unit("confidence", confidence)?, probabilities: ordered }))
95        }
96        (Question::Score(q), WireAnswer::Score { score, confidence, mut legend, mut probabilities }) => {
97            let n = q.level_count();
98            let mut ps = Vec::with_capacity(n);
99            let mut descriptions = Vec::with_capacity(n);
100            for level in 0..n {
101                let key = level.to_string();
102                ps.push(unit("probability", probabilities.remove(&key).ok_or_else(|| format!("{id:?}: no probability for level {level}"))?)?);
103                descriptions.push(legend.remove(&key).ok_or_else(|| format!("{id:?}: no legend for level {level}"))?);
104            }
105            if !probabilities.is_empty() || !legend.is_empty() {
106                return Err(format!("{id:?}: levels beyond the {n} asked"));
107            }
108            if !(0.0..=(n - 1) as f64).contains(&score) {
109                return Err(format!("{id:?}: score {score} is outside 0..={}", n - 1));
110            }
111            Ok(Answer::Score(ScoreAnswer { score, confidence: unit("confidence", confidence)?, probabilities: ps, legend: descriptions }))
112        }
113        (_, answer) => Err(format!("{id:?}: answered as {}, asked as another type", kind(&answer))),
114    }
115}
116
117fn kind(answer: &WireAnswer) -> &'static str {
118    match answer {
119        WireAnswer::Noul { .. } => "noul",
120        WireAnswer::Choice { .. } => "choice",
121        WireAnswer::Score { .. } => "score",
122    }
123}
124
125#[cfg(test)]
126mod tests {
127    use super::*;
128    use crate::{Choice, Json, Noul, Score};
129
130    fn pinned() -> ModelId {
131        ModelId::pinned("jev-1.13.0").unwrap()
132    }
133
134    fn body(answers: &str) -> Vec<u8> {
135        format!(r#"{{"model":"jev-1.13.0","answers":{answers},"usage":{{"input_tokens":1,"output_tokens":2}}}}"#).into_bytes()
136    }
137
138    #[test]
139    fn typed_answers() {
140        let mut qs = Questions::new();
141        let urgent = qs.noul("u", Noul::new(Json::text("?"))).unwrap();
142        let team = qs
143            .choice("t", Choice::new(Json::text("?"), [("sales".into(), None), ("billing".into(), None)]).unwrap())
144            .unwrap();
145        let mood = qs.score("m", Score::new(Json::text("?"), [Json::text("calm"), Json::text("angry")]).unwrap()).unwrap();
146        let r = Response::parse(
147            &pinned(),
148            &qs,
149            &body(
150                r#"{"u":{"type":"noul","noul":0.95},
151                    "t":{"type":"choice","choice":"billing","confidence":0.8,"probabilities":{"billing":0.9,"sales":0.1}},
152                    "m":{"type":"score","score":0.2,"confidence":0.7,"legend":{"0":"calm","1":"angry"},"probabilities":{"0":0.8,"1":0.2}}}"#,
153            ),
154        )
155        .unwrap();
156        assert_eq!(r.get(urgent).noul, 0.95);
157        assert_eq!(r.get(team).probabilities, vec![("sales".into(), 0.1), ("billing".into(), 0.9)]);
158        assert_eq!(r.get(mood).probabilities, vec![0.8, 0.2]);
159        assert_eq!(r.get(mood).normalized(), 0.2);
160        assert_eq!(r.usage(), Usage { input_tokens: 1, output_tokens: 2 });
161    }
162
163    #[test]
164    fn refuses_what_was_not_asked() {
165        let mut qs = Questions::new();
166        qs.choice("t", Choice::new(Json::text("?"), [("a".into(), None), ("b".into(), None)]).unwrap()).unwrap();
167        for answers in [
168            r#"{}"#,
169            r#"{"t":{"type":"noul","noul":0.5}}"#,
170            r#"{"t":{"type":"choice","choice":"c","confidence":1,"probabilities":{"a":0,"b":1}}}"#,
171            r#"{"t":{"type":"choice","choice":"a","confidence":1,"probabilities":{"a":1}}}"#,
172            r#"{"t":{"type":"choice","choice":"a","confidence":1,"probabilities":{"a":1,"b":0,"c":0}}}"#,
173            r#"{"t":{"type":"choice","choice":"a","confidence":1,"probabilities":{"a":1.5,"b":0}}}"#,
174            r#"{"t":{"type":"choice","choice":"a","confidence":1,"probabilities":{"a":1,"b":0}},"x":{"type":"noul","noul":0}}"#,
175        ] {
176            assert!(Response::parse(&pinned(), &qs, &body(answers)).is_err(), "{answers}");
177        }
178    }
179
180    #[test]
181    fn refuses_another_model() {
182        let mut qs = Questions::new();
183        qs.noul("u", Noul::new(Json::text("?"))).unwrap();
184        let b = br#"{"model":"jev-1.14.0","answers":{"u":{"type":"noul","noul":0.5}},"usage":{"input_tokens":1,"output_tokens":1}}"#;
185        assert!(matches!(Response::parse(&pinned(), &qs, b), Err(ProtocolError::ModelMismatch { .. })));
186    }
187
188    #[test]
189    #[should_panic(expected = "a Key from other Questions")]
190    fn foreign_key_panics() {
191        let mut a = Questions::new();
192        a.noul("u", Noul::new(Json::text("?"))).unwrap();
193        let mut b = Questions::new();
194        let foreign = b.noul("u", Noul::new(Json::text("?"))).unwrap();
195        let r = Response::parse(&pinned(), &a, &body(r#"{"u":{"type":"noul","noul":0.5}}"#)).unwrap();
196        r.get(foreign);
197    }
198}