use serde::{Deserialize, Serialize};
use crate::serve::api::engine::{GrammarKind, PromptCache, PromptCacheKey, ToolCallPolicy};
pub const PROMPT_CACHE_FORMAT_VERSION: u32 = 1;
pub const PROMPT_CACHE_PAYLOAD_KIND: &str = "prompt-cache";
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PromptCacheSnapshot {
pub format_version: u32,
pub tokens: Vec<u32>,
pub key: PromptCacheKeyPersist,
pub text: String,
pub reasoning_text: Option<String>,
pub completion_tokens: usize,
pub reasoning_tokens: Option<usize>,
pub finish_reason: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PromptCacheKeyPersist {
pub max_tokens: usize,
pub stop_strings: Vec<String>,
pub logit_bias_sorted: Vec<(u32, u32)>,
pub frequency_penalty_bits: u32,
pub presence_penalty_bits: u32,
pub min_p_bits: u32,
pub logprobs: bool,
pub top_logprobs: u32,
pub parallel_tool_calls: bool,
pub grammar_was_none: bool,
}
pub fn try_serialize(cache: &PromptCache) -> Option<Vec<u8>> {
if cache.tokens.is_empty() {
return None;
}
if cache.key.grammar.is_some() {
return None;
}
let snap = PromptCacheSnapshot {
format_version: PROMPT_CACHE_FORMAT_VERSION,
tokens: cache.tokens.clone(),
key: PromptCacheKeyPersist {
max_tokens: cache.key.max_tokens,
stop_strings: cache.key.stop_strings.clone(),
logit_bias_sorted: cache.key.logit_bias_sorted.clone(),
frequency_penalty_bits: cache.key.frequency_penalty_bits,
presence_penalty_bits: cache.key.presence_penalty_bits,
min_p_bits: cache.key.min_p_bits,
logprobs: cache.key.logprobs,
top_logprobs: cache.key.top_logprobs,
parallel_tool_calls: cache.key.parallel_tool_calls,
grammar_was_none: true,
},
text: cache.text.clone(),
reasoning_text: cache.reasoning_text.clone(),
completion_tokens: cache.completion_tokens,
reasoning_tokens: cache.reasoning_tokens,
finish_reason: cache.finish_reason.to_string(),
};
serde_json::to_vec(&snap).ok()
}
pub fn try_deserialize(bytes: &[u8]) -> Option<PromptCache> {
let snap: PromptCacheSnapshot = serde_json::from_slice(bytes).ok()?;
if snap.format_version != PROMPT_CACHE_FORMAT_VERSION {
return None;
}
let finish_reason: &'static str = match snap.finish_reason.as_str() {
"stop" => "stop",
"length" => "length",
"tool_calls" => "tool_calls",
_ => return None,
};
Some(PromptCache {
tokens: snap.tokens,
key: PromptCacheKey {
max_tokens: snap.key.max_tokens,
stop_strings: snap.key.stop_strings,
logit_bias_sorted: snap.key.logit_bias_sorted,
grammar: None,
grammar_kind: GrammarKind::default(),
frequency_penalty_bits: snap.key.frequency_penalty_bits,
presence_penalty_bits: snap.key.presence_penalty_bits,
min_p_bits: snap.key.min_p_bits,
tool_call_policy: ToolCallPolicy::default(),
logprobs: snap.key.logprobs,
top_logprobs: snap.key.top_logprobs,
parallel_tool_calls: snap.key.parallel_tool_calls,
},
text: snap.text,
reasoning_text: snap.reasoning_text,
completion_tokens: snap.completion_tokens,
reasoning_tokens: snap.reasoning_tokens,
finish_reason,
fragments: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn fresh_cache(tokens: Vec<u32>, text: &str) -> PromptCache {
let mut c = PromptCache::new();
c.tokens = tokens;
c.text = text.to_string();
c.completion_tokens = 4;
c.finish_reason = "length";
c
}
#[test]
fn empty_cache_yields_none() {
let c = PromptCache::new();
assert!(try_serialize(&c).is_none());
}
#[test]
fn round_trip_default_key_byte_recoverable() {
let original = fresh_cache(vec![1, 2, 3, 4], "hello world");
let bytes = try_serialize(&original).expect("serialize");
let restored = try_deserialize(&bytes).expect("deserialize");
assert_eq!(restored.tokens, original.tokens);
assert_eq!(restored.text, original.text);
assert_eq!(restored.completion_tokens, original.completion_tokens);
assert_eq!(restored.finish_reason, original.finish_reason);
assert_eq!(restored.key.max_tokens, original.key.max_tokens);
assert_eq!(
restored.key.parallel_tool_calls,
original.key.parallel_tool_calls
);
assert!(restored.key.grammar.is_none());
assert!(restored.fragments.is_none());
}
#[test]
fn unknown_finish_reason_yields_none_on_deserialize() {
let mut original = fresh_cache(vec![1], "x");
let leaked: &'static str = Box::leak(Box::new("bogus".to_string()));
original.finish_reason = leaked;
let bytes = try_serialize(&original).expect("serialize");
let restored = try_deserialize(&bytes);
assert!(
restored.is_none(),
"unknown finish_reason should yield None"
);
}
#[test]
fn version_mismatch_yields_none() {
let original = fresh_cache(vec![1], "x");
let bytes = try_serialize(&original).expect("serialize");
let mut s = String::from_utf8(bytes).expect("utf8");
s = s.replace(
&format!("\"format_version\":{}", PROMPT_CACHE_FORMAT_VERSION),
"\"format_version\":9999",
);
let restored = try_deserialize(s.as_bytes());
assert!(
restored.is_none(),
"future-version snapshot should yield None on this hf2q"
);
}
}