1//! A set of named questions, and the exact request bytes.
2
3use std::marker::PhantomData;
4use std::sync::atomic::{AtomicU64, Ordering};
5
6use serde::Serialize;
7
8use crate::answer::FromAnswer;
9use crate::{ChoiceAnswer, Choice, Json, ModelId, Noul, NoulAnswer, ProtocolError, Question, Score, ScoreAnswer};
10
11/// Questions to send together, in order. Each gets an id of the caller's
12/// choosing; the API says the id "is not sent to the underlying model".
13#[derive(Debug)]
14pub struct Questions {
15    /// Distinguishes this set's keys from any other's.
16    pub(crate) set: u64,
17    pub(crate) entries: Vec<(String, Question)>,
18}
19
20/// Reads the answer to one question, typed by its kind, so a Choice
21/// answer cannot be read as a Noul. Valid only for the `Questions` that
22/// issued it.
23pub struct Key<A> {
24    pub(crate) set: u64,
25    pub(crate) index: usize,
26    _answer: PhantomData<fn() -> A>,
27}
28
29impl<A> Clone for Key<A> {
30    fn clone(&self) -> Self {
31        *self
32    }
33}
34
35impl<A> Copy for Key<A> {}
36
37impl Default for Questions {
38    fn default() -> Self {
39        Self::new()
40    }
41}
42
43impl Questions {
44    pub fn new() -> Self {
45        static NEXT: AtomicU64 = AtomicU64::new(0);
46        Questions { set: NEXT.fetch_add(1, Ordering::Relaxed), entries: Vec::new() }
47    }
48
49    pub fn noul(&mut self, id: &str, question: Noul) -> Result<Key<NoulAnswer>, ProtocolError> {
50        self.push(id, Question::Noul(question))
51    }
52
53    pub fn choice(&mut self, id: &str, question: Choice) -> Result<Key<ChoiceAnswer>, ProtocolError> {
54        self.push(id, Question::Choice(question))
55    }
56
57    pub fn score(&mut self, id: &str, question: Score) -> Result<Key<ScoreAnswer>, ProtocolError> {
58        self.push(id, Question::Score(question))
59    }
60
61    pub fn len(&self) -> usize {
62        self.entries.len()
63    }
64
65    pub fn is_empty(&self) -> bool {
66        self.entries.is_empty()
67    }
68
69    /// Each question with its id, in the order they are sent.
70    pub fn iter(&self) -> impl ExactSizeIterator<Item = (&str, &Question)> {
71        self.entries.iter().map(|(id, q)| (id.as_str(), q))
72    }
73
74    fn push<A: FromAnswer>(&mut self, id: &str, question: Question) -> Result<Key<A>, ProtocolError> {
75        if id.is_empty() {
76            return Err(ProtocolError::Invalid("a question id is empty".into()));
77        }
78        if self.entries.iter().any(|(existing, _)| existing == id) {
79            return Err(ProtocolError::Invalid(format!("the question id {id:?} is used twice")));
80        }
81        self.entries.push((id.to_owned(), question));
82        Ok(Key { set: self.set, index: self.entries.len() - 1, _answer: PhantomData })
83    }
84}
85
86/// The body of `POST /v1/systemone`: `{"model","questions","state"}`,
87/// with `state` sent byte for byte as given.
88pub fn request_bytes(model: &ModelId, state: &Json, questions: &Questions) -> Result<Vec<u8>, ProtocolError> {
89    if questions.is_empty() {
90        return Err(ProtocolError::Invalid("a request needs at least one question".into()));
91    }
92    #[derive(Serialize)]
93    struct Body<'a> {
94        model: &'a str,
95        questions: QuestionMap<'a>,
96        state: &'a Json,
97    }
98    struct QuestionMap<'a>(&'a [(String, Question)]);
99    impl Serialize for QuestionMap<'_> {
100        fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
101            s.collect_map(self.0.iter().map(|(id, q)| (id, q)))
102        }
103    }
104    serde_json::to_vec(&Body { model: model.as_str(), questions: QuestionMap(&questions.entries), state })
105        .map_err(|e| ProtocolError::Invalid(e.to_string()))
106}
107
108/// The documented per-request input ceiling (digest §4), so a huge `state`
109/// cannot make a worst case run away.
110pub const MAX_REQUEST_TOKENS: f64 = 64_000.0;
111
112/// A conservative worst case, in tokens, for what one attempt could cost:
113/// for checking a budget BEFORE sending, when the true count (which only the
114/// response carries) is not yet known.
115///
116/// A UTF-8 character is never shorter than one token can be cheaper than, so
117/// counting characters can only overstate the token count - which is exactly
118/// what a worst case should do. The documented 64k-token ceiling per request
119/// caps it either way.
120pub fn worst_case_tokens(model: &ModelId, state: &Json, questions: &Questions) -> f64 {
121    let chars = request_bytes(model, state, questions)
122        .map(|bytes| String::from_utf8_lossy(&bytes).chars().count() as f64)
123        .unwrap_or(MAX_REQUEST_TOKENS);
124    chars.min(MAX_REQUEST_TOKENS)
125}
126
127/// [`worst_case_tokens`] in dollars, with every attempt the retry policy may
128/// make (`1 + retry::MAX_RETRIES`) billed in full.
129pub fn worst_case_dollars(model: &ModelId, state: &Json, questions: &Questions) -> f64 {
130    let attempts = f64::from(1 + crate::retry::MAX_RETRIES);
131    worst_case_tokens(model, state, questions) * attempts / 1e6 * crate::DOLLARS_PER_MTOK
132}
133
134/// One question object exactly as it appears in the request body,
135/// without its id: the bytes a cache key hashes.
136pub fn question_bytes(question: &Question) -> Result<Vec<u8>, ProtocolError> {
137    serde_json::to_vec(question).map_err(|e| ProtocolError::Invalid(e.to_string()))
138}
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143
144    #[test]
145    fn exact_bytes() {
146        let mut questions = Questions::new();
147        questions.noul("q", Noul::new(Json::text("Urgent?"))).unwrap();
148        let state = Json::canonical(r#"{"id":1,"body":"hi"}"#).unwrap();
149        let bytes = request_bytes(&ModelId::pinned("jev-1.13.0").unwrap(), &state, &questions).unwrap();
150        assert_eq!(
151            std::str::from_utf8(&bytes).unwrap(),
152            r#"{"model":"jev-1.13.0","questions":{"q":{"type":"noul","instructions":"Urgent?"}},"state":{"body":"hi","id":1}}"#
153        );
154    }
155
156    #[test]
157    fn question_bytes_are_the_request_bytes() {
158        let mut questions = Questions::new();
159        questions.noul("q", Noul::new(Json::text("Urgent?"))).unwrap();
160        let (_, question) = questions.iter().next().unwrap();
161        let alone = question_bytes(question).unwrap();
162        let body = request_bytes(&ModelId::pinned("jev-1.13.0").unwrap(), &Json::text("s"), &questions).unwrap();
163        let body = std::str::from_utf8(&body).unwrap();
164        assert!(body.contains(&format!(r#""q":{}"#, std::str::from_utf8(&alone).unwrap())));
165    }
166
167    #[test]
168    fn ids_are_unique_and_nonempty() {
169        let mut questions = Questions::new();
170        questions.noul("a", Noul::new(Json::text("?"))).unwrap();
171        assert!(questions.noul("a", Noul::new(Json::text("?"))).is_err());
172        assert!(questions.noul("", Noul::new(Json::text("?"))).is_err());
173    }
174
175    #[test]
176    fn no_questions_no_request() {
177        let model = ModelId::pinned("jev-1.13.0").unwrap();
178        assert!(request_bytes(&model, &Json::text("s"), &Questions::new()).is_err());
179    }
180
181    #[test]
182    fn a_huge_state_is_capped_at_the_documented_ceiling_not_left_unbounded() {
183        let model = ModelId::pinned("jev-1.13.0").unwrap();
184        let huge = Json::text(&"x".repeat(1_000_000));
185        let mut questions = Questions::new();
186        questions.noul("q", Noul::new(Json::text("long?"))).unwrap();
187        assert_eq!(worst_case_tokens(&model, &huge, &questions), MAX_REQUEST_TOKENS);
188        let small = Json::text("x");
189        assert!(worst_case_dollars(&model, &small, &questions) < worst_case_dollars(&model, &huge, &questions));
190    }
191}