use super::canonical_json;
use crate::{
hex::lower_hex,
output::redact_sensitive_text,
persistence::atomic_write,
providers::{ChatMessage, ProviderConversationItem},
};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::{
fs,
io::Read,
path::{Path, PathBuf},
};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct FileFingerprint {
pub path: PathBuf,
pub len: u64,
pub modified_nanos: u128,
pub content_hash: String,
}
impl FileFingerprint {
pub fn from_path(path: impl AsRef<Path>) -> anyhow::Result<Self> {
let path = path.as_ref();
let metadata = fs::metadata(path)?;
let modified_nanos = metadata
.modified()?
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos();
let mut file = fs::File::open(path)?;
let mut hasher = Sha256::new();
let mut buffer = [0_u8; 64 * 1024];
loop {
let bytes_read = file.read(&mut buffer)?;
if bytes_read == 0 {
break;
}
hasher.update(&buffer[..bytes_read]);
}
Ok(Self {
path: path.to_path_buf(),
len: metadata.len(),
modified_nanos,
content_hash: lower_hex(hasher.finalize()),
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ContextCacheEntry {
pub key: String,
pub token_estimate: usize,
#[serde(default)]
pub messages: Vec<ChatMessage>,
#[serde(default)]
pub input_material: String,
}
impl ContextCacheEntry {
fn redacted_for_persistence(&self) -> Self {
Self {
key: self.key.clone(),
token_estimate: self.token_estimate,
messages: self
.messages
.iter()
.map(|message| ChatMessage {
role: message.role.clone(),
content: redact_sensitive_text(&message.content),
})
.collect(),
input_material: redact_sensitive_text(&self.input_material),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ContextCache {
root: PathBuf,
}
impl ContextCache {
pub fn new(root: PathBuf) -> Self {
Self { root }
}
pub fn key(
provider: &str,
model: &str,
system_prompt: &str,
messages: &[ChatMessage],
files: &[FileFingerprint],
) -> String {
let items = messages
.iter()
.cloned()
.map(ProviderConversationItem::Message)
.collect::<Vec<_>>();
Self::key_for_conversation(provider, model, system_prompt, &items, files)
}
pub fn key_for_conversation(
provider: &str,
model: &str,
system_prompt: &str,
conversation_items: &[ProviderConversationItem],
files: &[FileFingerprint],
) -> String {
let input_material =
conversation_cache_material(provider, model, system_prompt, conversation_items);
Self::key_for_material(&input_material, files)
}
pub(crate) fn key_for_material(input_material: &str, files: &[FileFingerprint]) -> String {
if files.is_empty() {
return digest_material(input_material);
}
let mut data = String::from(input_material);
let mut sorted_files = files.iter().collect::<Vec<_>>();
sorted_files.sort_by(|left, right| left.path.cmp(&right.path));
for file in sorted_files {
push_material_field(&mut data, "file.path", &file.path.display().to_string());
push_material_field(&mut data, "file.len", &file.len.to_string());
push_material_field(
&mut data,
"file.modified_nanos",
&file.modified_nanos.to_string(),
);
push_material_field(&mut data, "file.content_hash", &file.content_hash);
}
digest_material(&data)
}
pub fn read(&self, key: &str) -> anyhow::Result<Option<ContextCacheEntry>> {
let path = self.path_for_key(key)?;
if !path.exists() {
return Ok(None);
}
Ok(Some(serde_json::from_str(&fs::read_to_string(path)?)?))
}
pub fn write(&self, entry: &ContextCacheEntry) -> anyhow::Result<()> {
fs::create_dir_all(&self.root)?;
let path = self.path_for_key(&entry.key)?;
let sanitized = entry.redacted_for_persistence();
atomic_write(&path, serde_json::to_string_pretty(&sanitized)?.as_bytes())?;
Ok(())
}
fn path_for_key(&self, key: &str) -> anyhow::Result<PathBuf> {
validate_cache_key(key)?;
let path = self.root.join(format!("{key}.json"));
let normalized_root = lexical_normalize(&self.root);
let normalized_path = lexical_normalize(&path);
if !normalized_path.starts_with(&normalized_root) {
anyhow::bail!("context cache key escapes cache root");
}
Ok(path)
}
}
fn digest_material(data: &str) -> String {
let digest = Sha256::digest(data.as_bytes());
lower_hex(digest)
}
fn validate_cache_key(key: &str) -> anyhow::Result<()> {
if key.len() != 64
|| !key
.chars()
.all(|ch| ch.is_ascii_digit() || ('a'..='f').contains(&ch))
{
anyhow::bail!("invalid context cache key");
}
Ok(())
}
fn lexical_normalize(path: &Path) -> PathBuf {
let mut normalized = PathBuf::new();
for component in path.components() {
match component {
std::path::Component::CurDir => {}
std::path::Component::ParentDir => {
normalized.pop();
}
other => normalized.push(other.as_os_str()),
}
}
normalized
}
pub fn conversation_cache_material(
provider: &str,
model: &str,
system_prompt: &str,
conversation_items: &[ProviderConversationItem],
) -> String {
let mut data = String::new();
push_material_field(&mut data, "provider", provider);
push_material_field(&mut data, "model", model);
push_material_field(&mut data, "system_prompt", system_prompt);
for item in conversation_items {
match item {
ProviderConversationItem::Message(message) => {
push_material_field(&mut data, "item", "message");
push_material_field(&mut data, "role", message.role.as_api_str());
push_material_field(&mut data, "content", &message.content);
}
ProviderConversationItem::ResponseItem(value) => {
push_material_field(&mut data, "item", "response_item");
push_material_field(&mut data, "value", &canonical_json(value));
}
ProviderConversationItem::ToolResult(result) => {
push_material_field(&mut data, "item", "tool_result");
push_material_field(&mut data, "call_id", &result.call_id);
push_material_field(&mut data, "tool_name", &result.tool_name);
push_material_field(&mut data, "success", &result.success.to_string());
push_material_field(&mut data, "output", &result.output);
}
ProviderConversationItem::LegacyReplayNote {
event_type,
content,
} => {
push_material_field(&mut data, "item", "legacy_replay_note");
push_material_field(&mut data, "event_type", event_type);
push_material_field(&mut data, "content", content);
}
}
}
data
}
fn push_material_field(data: &mut String, key: &str, value: &str) {
data.push_str(key);
data.push('=');
data.push_str(&value.len().to_string());
data.push(':');
data.push_str(value);
data.push('\n');
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn file_fingerprint_hashes_file_content() {
let temp_dir = tempfile::tempdir().expect("temp dir");
let path = temp_dir.path().join("input.txt");
let content = b"context cache fingerprint content";
fs::write(&path, content).expect("write file");
let fingerprint = FileFingerprint::from_path(&path).expect("fingerprint");
assert_eq!(fingerprint.len, content.len() as u64);
assert_eq!(fingerprint.content_hash, lower_hex(Sha256::digest(content)));
}
#[test]
fn key_for_material_is_independent_of_file_order() {
let first = FileFingerprint {
path: PathBuf::from("b.txt"),
len: 10,
modified_nanos: 1,
content_hash: "b".repeat(64),
};
let second = FileFingerprint {
path: PathBuf::from("a.txt"),
len: 20,
modified_nanos: 2,
content_hash: "a".repeat(64),
};
let forward = ContextCache::key_for_material("material", &[first.clone(), second.clone()]);
let reverse = ContextCache::key_for_material("material", &[second, first]);
assert_eq!(forward, reverse);
}
#[test]
fn write_redacts_persisted_messages_and_input_material_without_changing_key_metadata() {
let temp_dir = tempfile::tempdir().expect("temp dir");
let cache = ContextCache::new(temp_dir.path().to_path_buf());
let message_bearer = format!("Bearer message{}", "x".repeat(24));
let provider_token = format!("sk-{}", "x".repeat(24));
let input_bearer = format!("Bearer input{}", "x".repeat(24));
let input_api_key = "input-secret";
let message_api_key = "message-secret";
let entry = ContextCacheEntry {
key: "a".repeat(64),
token_estimate: 42,
messages: vec![
ChatMessage::user(format!("api_key='{message_api_key}' {message_bearer}")),
ChatMessage::assistant(format!("provider token {provider_token}")),
],
input_material: format!(
"input api_key={input_api_key} {input_bearer} {provider_token}"
),
};
cache.write(&entry).expect("write cache entry");
let persisted_path = temp_dir.path().join(format!("{}.json", entry.key));
let persisted = fs::read_to_string(persisted_path).expect("persisted cache json");
assert!(
persisted.contains(&format!("\"key\": \"{}\"", entry.key)),
"{persisted}"
);
assert!(persisted.contains("\"token_estimate\": 42"), "{persisted}");
for secret in [
message_api_key,
message_bearer.as_str(),
provider_token.as_str(),
input_api_key,
input_bearer.as_str(),
] {
assert!(!persisted.contains(secret), "{persisted}");
}
assert!(
persisted.contains("<redacted>") || persisted.contains("[REDACTED]"),
"{persisted}"
);
}
}