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}