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/// One question object exactly as it appears in the request body, 109/// without its id: the bytes a cache key hashes. 110pub fn question_bytes(question: &Question) -> Result<Vec<u8>, ProtocolError> { 111 serde_json::to_vec(question).map_err(|e| ProtocolError::Invalid(e.to_string())) 112} 113 114#[cfg(test)] 115mod tests { 116 use super::*; 117 118 #[test] 119 fn exact_bytes() { 120 let mut questions = Questions::new(); 121 questions.noul("q", Noul::new(Json::text("Urgent?"))).unwrap(); 122 let state = Json::canonical(r#"{"id":1,"body":"hi"}"#).unwrap(); 123 let bytes = request_bytes(&ModelId::pinned("jev-1.13.0").unwrap(), &state, &questions).unwrap(); 124 assert_eq!( 125 std::str::from_utf8(&bytes).unwrap(), 126 r#"{"model":"jev-1.13.0","questions":{"q":{"type":"noul","instructions":"Urgent?"}},"state":{"body":"hi","id":1}}"# 127 ); 128 } 129 130 #[test] 131 fn question_bytes_are_the_request_bytes() { 132 let mut questions = Questions::new(); 133 questions.noul("q", Noul::new(Json::text("Urgent?"))).unwrap(); 134 let (_, question) = questions.iter().next().unwrap(); 135 let alone = question_bytes(question).unwrap(); 136 let body = request_bytes(&ModelId::pinned("jev-1.13.0").unwrap(), &Json::text("s"), &questions).unwrap(); 137 let body = std::str::from_utf8(&body).unwrap(); 138 assert!(body.contains(&format!(r#""q":{}"#, std::str::from_utf8(&alone).unwrap()))); 139 } 140 141 #[test] 142 fn ids_are_unique_and_nonempty() { 143 let mut questions = Questions::new(); 144 questions.noul("a", Noul::new(Json::text("?"))).unwrap(); 145 assert!(questions.noul("a", Noul::new(Json::text("?"))).is_err()); 146 assert!(questions.noul("", Noul::new(Json::text("?"))).is_err()); 147 } 148 149 #[test] 150 fn no_questions_no_request() { 151 let model = ModelId::pinned("jev-1.13.0").unwrap(); 152 assert!(request_bytes(&model, &Json::text("s"), &Questions::new()).is_err()); 153 } 154}