lean_ctx/proxy/
prefix_replay.rs1use std::collections::HashMap;
16use std::sync::{Mutex, OnceLock};
17
18use serde_json::Value;
19
20const MAX_TRACKED: usize = 2048;
21
22#[derive(Clone)]
23struct ConversationPrefix {
24 forwarded_bytes: Vec<u8>,
25 original_hashes: Vec<u64>,
26 count: usize,
27}
28
29pub struct AppendDelta {
30 pub prefix_bytes: Vec<u8>,
31 pub delta_start: usize,
32}
33
34fn store() -> &'static Mutex<HashMap<u64, ConversationPrefix>> {
35 static STORE: OnceLock<Mutex<HashMap<u64, ConversationPrefix>>> = OnceLock::new();
36 STORE.get_or_init(|| Mutex::new(HashMap::new()))
37}
38
39fn message_hash(msg: &Value) -> u64 {
40 use std::hash::{Hash, Hasher};
41 let canonical = serde_json::to_string(msg).unwrap_or_default();
42 let mut hasher = std::collections::hash_map::DefaultHasher::new();
43 canonical.hash(&mut hasher);
44 hasher.finish()
45}
46
47pub fn conversation_id(system: Option<&Value>, messages: &[Value]) -> u64 {
49 use std::hash::{Hash, Hasher};
50 let mut hasher = std::collections::hash_map::DefaultHasher::new();
51 if let Some(sys) = system {
52 serde_json::to_string(sys)
53 .unwrap_or_default()
54 .hash(&mut hasher);
55 }
56 if let Some(first) = messages.first() {
57 serde_json::to_string(first)
58 .unwrap_or_default()
59 .hash(&mut hasher);
60 }
61 hasher.finish()
62}
63
64pub fn detect_append_only(conv_id: u64, messages: &[Value]) -> Option<AppendDelta> {
67 let guard = store().lock().ok()?;
68 let prev = guard.get(&conv_id)?;
69
70 if messages.len() <= prev.count {
71 return None;
72 }
73
74 let current_hashes: Vec<u64> = messages.iter().map(message_hash).collect();
75 for (i, prev_hash) in prev.original_hashes.iter().enumerate() {
76 if current_hashes.get(i) != Some(prev_hash) {
77 return None;
78 }
79 }
80
81 Some(AppendDelta {
82 prefix_bytes: prev.forwarded_bytes.clone(),
83 delta_start: prev.count,
84 })
85}
86
87pub fn record_forwarded(conv_id: u64, forwarded: Vec<u8>, originals: &[Value], msg_count: usize) {
89 let hashes: Vec<u64> = originals.iter().take(msg_count).map(message_hash).collect();
90 let entry = ConversationPrefix {
91 forwarded_bytes: forwarded,
92 original_hashes: hashes,
93 count: msg_count,
94 };
95
96 if let Ok(mut guard) = store().lock() {
97 if guard.len() >= MAX_TRACKED
98 && !guard.contains_key(&conv_id)
99 && let Some(&oldest) = guard.keys().next()
100 {
101 guard.remove(&oldest);
102 }
103 guard.insert(conv_id, entry);
104 }
105}
106
107pub fn overlay_prefix(prefix_bytes: &[u8], delta_messages: &[Value]) -> Option<Vec<u8>> {
110 if delta_messages.is_empty() {
111 return Some(prefix_bytes.to_vec());
112 }
113
114 let prefix_str = std::str::from_utf8(prefix_bytes).ok()?;
115 let trimmed = prefix_str.trim_end();
116 if !trimmed.ends_with(']') {
117 return None;
118 }
119 let without_bracket = &trimmed[..trimmed.len() - 1];
120
121 let mut result = without_bracket.as_bytes().to_vec();
122 for msg in delta_messages {
123 result.extend_from_slice(b",");
124 let serialised = serde_json::to_string(msg).ok()?;
125 result.extend_from_slice(serialised.as_bytes());
126 }
127 result.push(b']');
128 Some(result)
129}
130
131#[cfg(test)]
132pub fn clear() {
133 if let Ok(mut guard) = store().lock() {
134 guard.clear();
135 }
136}
137
138#[cfg(test)]
139mod tests {
140 use super::*;
141 use serde_json::json;
142
143 fn sample_messages() -> Vec<Value> {
144 vec![
145 json!({"role": "user", "content": "hello"}),
146 json!({"role": "assistant", "content": "hi there"}),
147 ]
148 }
149
150 #[test]
151 fn append_only_detection_works() {
152 clear();
153 let msgs = sample_messages();
154 let conv = conversation_id(None, &msgs);
155 let forwarded = serde_json::to_vec(&msgs).unwrap();
156
157 record_forwarded(conv, forwarded.clone(), &msgs, msgs.len());
158
159 let mut extended = msgs.clone();
160 extended.push(json!({"role": "user", "content": "what's up?"}));
161
162 let delta = detect_append_only(conv, &extended).expect("should detect append-only");
163 assert_eq!(delta.delta_start, 2);
164 assert_eq!(delta.prefix_bytes, forwarded);
165 }
166
167 #[test]
168 fn detection_fails_on_modified_prefix() {
169 clear();
170 let msgs = sample_messages();
171 let conv = conversation_id(None, &msgs);
172 let forwarded = serde_json::to_vec(&msgs).unwrap();
173 record_forwarded(conv, forwarded, &msgs, msgs.len());
174
175 let mut modified = msgs;
176 modified[0] = json!({"role": "user", "content": "different"});
177 modified.push(json!({"role": "user", "content": "extra"}));
178 assert!(detect_append_only(conv, &modified).is_none());
179 }
180
181 #[test]
182 fn prefix_replay_is_byte_identical_across_turns() {
183 clear();
184 let msgs = sample_messages();
185 let conv = conversation_id(None, &msgs);
186 let forwarded = serde_json::to_vec(&msgs).unwrap();
187 record_forwarded(conv, forwarded.clone(), &msgs, msgs.len());
188
189 let mut turn2 = msgs.clone();
190 turn2.push(json!({"role": "user", "content": "next"}));
191
192 let delta = detect_append_only(conv, &turn2).unwrap();
193 let result = overlay_prefix(&delta.prefix_bytes, &turn2[delta.delta_start..]).unwrap();
194 let result_str = String::from_utf8(result).unwrap();
195 let parsed: Vec<Value> = serde_json::from_str(&result_str).unwrap();
196 assert_eq!(parsed.len(), 3);
197
198 let prefix_portion = &result_str[..forwarded.len() - 1];
199 let original_prefix = std::str::from_utf8(&forwarded[..forwarded.len() - 1]).unwrap();
200 assert_eq!(
201 prefix_portion, original_prefix,
202 "prefix bytes must be identical"
203 );
204 }
205
206 #[test]
207 fn overlay_with_empty_delta_returns_prefix() {
208 let msgs = sample_messages();
209 let bytes = serde_json::to_vec(&msgs).unwrap();
210 let result = overlay_prefix(&bytes, &[]).unwrap();
211 assert_eq!(result, bytes);
212 }
213
214 #[test]
215 fn conversation_id_is_deterministic() {
216 let sys = json!("You are helpful");
217 let msgs = sample_messages();
218 assert_eq!(
219 conversation_id(Some(&sys), &msgs),
220 conversation_id(Some(&sys), &msgs)
221 );
222 }
223
224 #[test]
225 fn max_tracked_evicts_oldest() {
226 clear();
227 for i in 0..MAX_TRACKED + 10 {
228 let msgs = vec![json!({"role": "user", "content": format!("msg {i}")})];
229 let conv = conversation_id(None, &msgs);
230 record_forwarded(conv, serde_json::to_vec(&msgs).unwrap(), &msgs, 1);
231 }
232 assert!(store().lock().unwrap().len() <= MAX_TRACKED);
233 }
234}