Skip to main content

lean_ctx/proxy/
prefix_replay.rs

1//! Append-only delta detection and forwarded-byte replay for provider prefix
2//! cache stability.
3//!
4//! When the proxy compresses and forwards a request, the serialised prefix
5//! bytes become the provider's cache key. Re-serialising an identical `Value`
6//! can produce subtly different bytes (JSON key order, float precision, Unicode
7//! escaping), causing a cache miss even though the content has not changed.
8//!
9//! This module caches the **exact forwarded bytes** and replays them verbatim
10//! on subsequent turns when the new request is an append-only extension.
11//! Only the delta (new messages) is freshly serialised.
12//!
13//! Active only in `ProxyMode::Cache`.
14
15use 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
47/// Conversation identity from system prompt + first user message.
48pub 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
64/// Detect whether current messages are an append-only extension of the previous
65/// turn. Returns cached prefix bytes and delta start index if so.
66pub 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
87/// Record the forwarded prefix bytes after a successful upstream send.
88pub 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
107/// Overlay cached prefix bytes with fresh delta bytes, producing a valid JSON
108/// array where the prefix portion is byte-identical to the previous forward.
109pub 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}