1//! The durable cache's key: the hash of what was actually sent, so an
2//! incomplete key cannot be written (contract *Cache*).
3//!
4//! `sha256(model ‖ prompt version ‖ namespace ‖ layout ‖ state ‖ question)`,
5//! where state and question are the exact bytes of the request body and
6//! the question object carries no id: the API says the id "is not sent to
7//! the underlying model", so two ids for one question are one judgment.
8//! Every field is length-prefixed, so no two different inputs concatenate
9//! to the same bytes.
10
11use std::fmt;
12
13use jev_protocol::{ModelId, ProtocolError, Question, Questions, question_bytes};
14use sha2::{Digest, Sha256};
15
16use crate::batch::Placement;
17
18/// Changes when this extension changes how it turns SQL into questions,
19/// so answers to the old wording are never served for the new one.
20#[derive(Clone, Copy, Debug, PartialEq, Eq)]
21pub struct PromptVersion(pub u32);
22
23pub const PROMPT_VERSION: PromptVersion = PromptVersion(1);
24
25/// How a row reaches the model (contract *Execution*). Answers are not
26/// assumed equal across layouts, so the layout is part of the key.
27#[derive(Clone, Copy, Debug, PartialEq, Eq)]
28pub enum Layout {
29    /// The row is the state; each question is a branch.
30    RowAsState,
31    /// The state is shared context; the row is in the question's instructions.
32    RowInInstructions,
33}
34
35impl Layout {
36    /// The name hashed into the key and stored in `jev_cache.layout`:
37    /// stable across reorderings of the enum.
38    pub fn tag(self) -> &'static str {
39        match self {
40            Layout::RowAsState => "row-as-state",
41            Layout::RowInInstructions => "row-in-instructions",
42        }
43    }
44}
45
46/// Everything in a key besides the row's placement and the question:
47/// fixed for a scan.
48#[derive(Clone, Debug)]
49pub struct KeyScope {
50    /// `jev.model`. Lookup uses the pin because the answering model is not
51    /// known before the call; the response's model is checked against it.
52    pub model: ModelId,
53    pub prompt: PromptVersion,
54    /// `jev.cache_namespace`. Changing it re-judges every input.
55    pub namespace: String,
56}
57
58#[derive(Clone, Copy, PartialEq, Eq, Hash)]
59pub struct CacheKey([u8; 32]);
60
61impl CacheKey {
62    /// The key for one question asked about a row placed as the batch
63    /// planner placed it: its layout and state.
64    pub fn new(scope: &KeyScope, placed: &Placement, question: &Question) -> Result<Self, ProtocolError> {
65        Ok(Self::from_bytes(scope, placed, &question_bytes(question)?))
66    }
67
68    /// One key per question, in the order the questions are sent.
69    pub fn for_each(scope: &KeyScope, placed: &Placement, questions: &Questions) -> Result<Vec<Self>, ProtocolError> {
70        questions.iter().map(|(_, q)| Self::new(scope, placed, q)).collect()
71    }
72
73    fn from_bytes(scope: &KeyScope, placed: &Placement, question: &[u8]) -> Self {
74        let mut h = Sha256::new();
75        let mut field = |bytes: &[u8]| {
76            h.update((bytes.len() as u64).to_be_bytes());
77            h.update(bytes);
78        };
79        field(b"postjevsql cache key v1");
80        field(scope.model.as_str().as_bytes());
81        field(&scope.prompt.0.to_be_bytes());
82        field(scope.namespace.as_bytes());
83        field(placed.layout().tag().as_bytes());
84        field(placed.state().as_str().as_bytes());
85        field(question);
86        CacheKey(h.finalize().into())
87    }
88
89    pub fn as_bytes(&self) -> &[u8; 32] {
90        &self.0
91    }
92}
93
94impl fmt::Display for CacheKey {
95    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96        self.0.iter().try_for_each(|b| write!(f, "{b:02x}"))
97    }
98}
99
100impl fmt::Debug for CacheKey {
101    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
102        write!(f, "CacheKey({self})")
103    }
104}
105
106#[cfg(test)]
107mod tests {
108    use jev_protocol::{Choice, Json, Noul};
109
110    use super::*;
111
112    fn scope() -> KeyScope {
113        KeyScope {
114            model: ModelId::pinned("jev-1.13.0").unwrap(),
115            prompt: PROMPT_VERSION,
116            namespace: String::new(),
117        }
118    }
119
120    fn state() -> Placement {
121        placed(r#"{"id":1,"body":"hi"}"#)
122    }
123
124    fn placed(row: &str) -> Placement {
125        Placement::new(Layout::RowAsState, Json::canonical(row).unwrap())
126    }
127
128    fn urgent() -> Question {
129        Question::Noul(Noul::new(Json::text("Urgent?")))
130    }
131
132    fn key(scope: &KeyScope) -> CacheKey {
133        CacheKey::new(scope, &state(), &urgent()).unwrap()
134    }
135
136    #[test]
137    fn deterministic() {
138        assert_eq!(key(&scope()), key(&scope()));
139        assert_eq!(key(&scope()).to_string().len(), 64);
140    }
141
142    #[test]
143    fn changes_with_layout() {
144        // Built directly: the planner does not place a row in
145        // instructions until the layout-parity gate passes.
146        let other = Placement::new(Layout::RowInInstructions, state().state().clone());
147        assert_ne!(key(&scope()), CacheKey::new(&scope(), &other, &urgent()).unwrap());
148    }
149
150    #[test]
151    fn changes_with_namespace() {
152        let other = KeyScope { namespace: "rejudge".into(), ..scope() };
153        assert_ne!(key(&scope()), key(&other));
154    }
155
156    #[test]
157    fn changes_with_model() {
158        let other = KeyScope { model: ModelId::pinned("jev-1.12.0").unwrap(), ..scope() };
159        assert_ne!(key(&scope()), key(&other));
160    }
161
162    #[test]
163    fn changes_with_prompt_version() {
164        let other = KeyScope { prompt: PromptVersion(PROMPT_VERSION.0 + 1), ..scope() };
165        assert_ne!(key(&scope()), key(&other));
166    }
167
168    #[test]
169    fn changes_with_state_and_question() {
170        let s = scope();
171        let other_state = placed(r#"{"id":2,"body":"hi"}"#);
172        assert_ne!(key(&s), CacheKey::new(&s, &other_state, &urgent()).unwrap());
173        let other_q = Question::Noul(Noul::new(Json::text("Spam?")));
174        assert_ne!(key(&s), CacheKey::new(&s, &state(), &other_q).unwrap());
175    }
176
177    #[test]
178    fn choice_option_order_is_part_of_the_key() {
179        let choice = |labels: [&str; 2]| {
180            Question::Choice(
181                Choice::new(Json::text("Which?"), labels.map(|l| (l.to_owned(), None))).unwrap(),
182            )
183        };
184        let s = scope();
185        assert_ne!(
186            CacheKey::new(&s, &state(), &choice(["a", "b"])).unwrap(),
187            CacheKey::new(&s, &state(), &choice(["b", "a"])).unwrap(),
188        );
189    }
190
191    #[test]
192    fn not_with_question_id() {
193        let ask = |id: &str| {
194            let mut questions = Questions::new();
195            questions.noul(id, Noul::new(Json::text("Urgent?"))).unwrap();
196            CacheKey::for_each(&scope(), &state(), &questions).unwrap()
197        };
198        assert_eq!(ask("q0"), ask("row-17/urgent"));
199        assert_eq!(ask("q0"), vec![key(&scope())]);
200    }
201}