Skip to main content

magi_code/context/
cache.rs

1use super::canonical_json;
2use crate::{
3    hex::lower_hex,
4    output::redact_sensitive_text,
5    persistence::atomic_write,
6    providers::{ChatMessage, ProviderConversationItem},
7};
8use serde::{Deserialize, Serialize};
9use sha2::{Digest, Sha256};
10use std::{
11    fs,
12    io::Read,
13    path::{Path, PathBuf},
14};
15
16#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
17pub struct FileFingerprint {
18    pub path: PathBuf,
19    pub len: u64,
20    pub modified_nanos: u128,
21    pub content_hash: String,
22}
23
24impl FileFingerprint {
25    pub fn from_path(path: impl AsRef<Path>) -> anyhow::Result<Self> {
26        let path = path.as_ref();
27        let metadata = fs::metadata(path)?;
28        let modified_nanos = metadata
29            .modified()?
30            .duration_since(std::time::UNIX_EPOCH)
31            .unwrap_or_default()
32            .as_nanos();
33        let mut file = fs::File::open(path)?;
34        let mut hasher = Sha256::new();
35        let mut buffer = [0_u8; 64 * 1024];
36        loop {
37            let bytes_read = file.read(&mut buffer)?;
38            if bytes_read == 0 {
39                break;
40            }
41            hasher.update(&buffer[..bytes_read]);
42        }
43        Ok(Self {
44            path: path.to_path_buf(),
45            len: metadata.len(),
46            modified_nanos,
47            content_hash: lower_hex(hasher.finalize()),
48        })
49    }
50}
51
52#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
53pub struct ContextCacheEntry {
54    pub key: String,
55    pub token_estimate: usize,
56    #[serde(default)]
57    pub messages: Vec<ChatMessage>,
58    #[serde(default)]
59    pub input_material: String,
60}
61
62impl ContextCacheEntry {
63    fn redacted_for_persistence(&self) -> Self {
64        Self {
65            key: self.key.clone(),
66            token_estimate: self.token_estimate,
67            messages: self
68                .messages
69                .iter()
70                .map(|message| ChatMessage {
71                    role: message.role.clone(),
72                    content: redact_sensitive_text(&message.content),
73                })
74                .collect(),
75            input_material: redact_sensitive_text(&self.input_material),
76        }
77    }
78}
79
80#[derive(Debug, Clone, PartialEq, Eq)]
81pub struct ContextCache {
82    root: PathBuf,
83}
84
85impl ContextCache {
86    pub fn new(root: PathBuf) -> Self {
87        Self { root }
88    }
89
90    pub fn key(
91        provider: &str,
92        model: &str,
93        system_prompt: &str,
94        messages: &[ChatMessage],
95        files: &[FileFingerprint],
96    ) -> String {
97        let items = messages
98            .iter()
99            .cloned()
100            .map(ProviderConversationItem::Message)
101            .collect::<Vec<_>>();
102        Self::key_for_conversation(provider, model, system_prompt, &items, files)
103    }
104
105    pub fn key_for_conversation(
106        provider: &str,
107        model: &str,
108        system_prompt: &str,
109        conversation_items: &[ProviderConversationItem],
110        files: &[FileFingerprint],
111    ) -> String {
112        let input_material =
113            conversation_cache_material(provider, model, system_prompt, conversation_items);
114        Self::key_for_material(&input_material, files)
115    }
116
117    pub(crate) fn key_for_material(input_material: &str, files: &[FileFingerprint]) -> String {
118        if files.is_empty() {
119            return digest_material(input_material);
120        }
121        let mut data = String::from(input_material);
122        let mut sorted_files = files.iter().collect::<Vec<_>>();
123        sorted_files.sort_by(|left, right| left.path.cmp(&right.path));
124        for file in sorted_files {
125            push_material_field(&mut data, "file.path", &file.path.display().to_string());
126            push_material_field(&mut data, "file.len", &file.len.to_string());
127            push_material_field(
128                &mut data,
129                "file.modified_nanos",
130                &file.modified_nanos.to_string(),
131            );
132            push_material_field(&mut data, "file.content_hash", &file.content_hash);
133        }
134        digest_material(&data)
135    }
136
137    pub fn read(&self, key: &str) -> anyhow::Result<Option<ContextCacheEntry>> {
138        let path = self.path_for_key(key)?;
139        if !path.exists() {
140            return Ok(None);
141        }
142        Ok(Some(serde_json::from_str(&fs::read_to_string(path)?)?))
143    }
144
145    pub fn write(&self, entry: &ContextCacheEntry) -> anyhow::Result<()> {
146        fs::create_dir_all(&self.root)?;
147        let path = self.path_for_key(&entry.key)?;
148        let sanitized = entry.redacted_for_persistence();
149        atomic_write(&path, serde_json::to_string_pretty(&sanitized)?.as_bytes())?;
150        Ok(())
151    }
152
153    fn path_for_key(&self, key: &str) -> anyhow::Result<PathBuf> {
154        validate_cache_key(key)?;
155        let path = self.root.join(format!("{key}.json"));
156        let normalized_root = lexical_normalize(&self.root);
157        let normalized_path = lexical_normalize(&path);
158        if !normalized_path.starts_with(&normalized_root) {
159            anyhow::bail!("context cache key escapes cache root");
160        }
161        Ok(path)
162    }
163}
164
165fn digest_material(data: &str) -> String {
166    let digest = Sha256::digest(data.as_bytes());
167    lower_hex(digest)
168}
169
170fn validate_cache_key(key: &str) -> anyhow::Result<()> {
171    if key.len() != 64
172        || !key
173            .chars()
174            .all(|ch| ch.is_ascii_digit() || ('a'..='f').contains(&ch))
175    {
176        anyhow::bail!("invalid context cache key");
177    }
178    Ok(())
179}
180
181fn lexical_normalize(path: &Path) -> PathBuf {
182    let mut normalized = PathBuf::new();
183    for component in path.components() {
184        match component {
185            std::path::Component::CurDir => {}
186            std::path::Component::ParentDir => {
187                normalized.pop();
188            }
189            other => normalized.push(other.as_os_str()),
190        }
191    }
192    normalized
193}
194
195pub fn conversation_cache_material(
196    provider: &str,
197    model: &str,
198    system_prompt: &str,
199    conversation_items: &[ProviderConversationItem],
200) -> String {
201    let mut data = String::new();
202    push_material_field(&mut data, "provider", provider);
203    push_material_field(&mut data, "model", model);
204    push_material_field(&mut data, "system_prompt", system_prompt);
205    for item in conversation_items {
206        match item {
207            ProviderConversationItem::Message(message) => {
208                push_material_field(&mut data, "item", "message");
209                push_material_field(&mut data, "role", message.role.as_api_str());
210                push_material_field(&mut data, "content", &message.content);
211            }
212            ProviderConversationItem::ResponseItem(value) => {
213                push_material_field(&mut data, "item", "response_item");
214                push_material_field(&mut data, "value", &canonical_json(value));
215            }
216            ProviderConversationItem::ToolResult(result) => {
217                push_material_field(&mut data, "item", "tool_result");
218                push_material_field(&mut data, "call_id", &result.call_id);
219                push_material_field(&mut data, "tool_name", &result.tool_name);
220                push_material_field(&mut data, "success", &result.success.to_string());
221                push_material_field(&mut data, "output", &result.output);
222            }
223            ProviderConversationItem::LegacyReplayNote {
224                event_type,
225                content,
226            } => {
227                push_material_field(&mut data, "item", "legacy_replay_note");
228                push_material_field(&mut data, "event_type", event_type);
229                push_material_field(&mut data, "content", content);
230            }
231        }
232    }
233    data
234}
235
236fn push_material_field(data: &mut String, key: &str, value: &str) {
237    data.push_str(key);
238    data.push('=');
239    data.push_str(&value.len().to_string());
240    data.push(':');
241    data.push_str(value);
242    data.push('\n');
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248
249    #[test]
250    fn file_fingerprint_hashes_file_content() {
251        let temp_dir = tempfile::tempdir().expect("temp dir");
252        let path = temp_dir.path().join("input.txt");
253        let content = b"context cache fingerprint content";
254        fs::write(&path, content).expect("write file");
255
256        let fingerprint = FileFingerprint::from_path(&path).expect("fingerprint");
257
258        assert_eq!(fingerprint.len, content.len() as u64);
259        assert_eq!(fingerprint.content_hash, lower_hex(Sha256::digest(content)));
260    }
261
262    #[test]
263    fn key_for_material_is_independent_of_file_order() {
264        let first = FileFingerprint {
265            path: PathBuf::from("b.txt"),
266            len: 10,
267            modified_nanos: 1,
268            content_hash: "b".repeat(64),
269        };
270        let second = FileFingerprint {
271            path: PathBuf::from("a.txt"),
272            len: 20,
273            modified_nanos: 2,
274            content_hash: "a".repeat(64),
275        };
276
277        let forward = ContextCache::key_for_material("material", &[first.clone(), second.clone()]);
278        let reverse = ContextCache::key_for_material("material", &[second, first]);
279
280        assert_eq!(forward, reverse);
281    }
282
283    #[test]
284    fn write_redacts_persisted_messages_and_input_material_without_changing_key_metadata() {
285        let temp_dir = tempfile::tempdir().expect("temp dir");
286        let cache = ContextCache::new(temp_dir.path().to_path_buf());
287        let message_bearer = format!("Bearer message{}", "x".repeat(24));
288        let provider_token = format!("sk-{}", "x".repeat(24));
289        let input_bearer = format!("Bearer input{}", "x".repeat(24));
290        let input_api_key = "input-secret";
291        let message_api_key = "message-secret";
292        let entry = ContextCacheEntry {
293            key: "a".repeat(64),
294            token_estimate: 42,
295            messages: vec![
296                ChatMessage::user(format!("api_key='{message_api_key}' {message_bearer}")),
297                ChatMessage::assistant(format!("provider token {provider_token}")),
298            ],
299            input_material: format!(
300                "input api_key={input_api_key} {input_bearer} {provider_token}"
301            ),
302        };
303
304        cache.write(&entry).expect("write cache entry");
305
306        let persisted_path = temp_dir.path().join(format!("{}.json", entry.key));
307        let persisted = fs::read_to_string(persisted_path).expect("persisted cache json");
308        assert!(
309            persisted.contains(&format!("\"key\": \"{}\"", entry.key)),
310            "{persisted}"
311        );
312        assert!(persisted.contains("\"token_estimate\": 42"), "{persisted}");
313        for secret in [
314            message_api_key,
315            message_bearer.as_str(),
316            provider_token.as_str(),
317            input_api_key,
318            input_bearer.as_str(),
319        ] {
320            assert!(!persisted.contains(secret), "{persisted}");
321        }
322        assert!(
323            persisted.contains("<redacted>") || persisted.contains("[REDACTED]"),
324            "{persisted}"
325        );
326    }
327}