use misanthropic::prompt::{
Prompt,
index::{BlockIndex, Index, IndexMut},
message::CacheControl,
};
use crate::reactor::inference::Quirks;
const MAX_CACHE_CONTROLS_PER_REQUEST: usize = 4;
pub const ROLL_WINDOW: usize = 2;
pub fn roll_breakpoints(quirks: &Quirks, prompt: &mut Prompt) {
roll_breakpoints_with(quirks, prompt, CacheControl::one_hour());
}
pub fn roll_breakpoints_with(
quirks: &Quirks,
prompt: &mut Prompt,
cache_control: CacheControl,
) {
if quirks.cache_markers_ignored {
return;
}
let Some(anchor) = prompt.messages.len().checked_sub(1) else {
return;
};
windowed(prompt, ROLL_WINDOW, anchor, cache_control);
}
fn windowed(
prompt: &mut Prompt,
n: usize,
anchor: usize,
cache_control: CacheControl,
) {
let server_tool_markers = prompt.tools.as_ref().map_or(0, |tools| {
tools
.iter()
.filter(|t| t.as_method().is_none() && t.is_cached())
.count()
});
let pinned = server_tool_markers
+ prompt
.indices()
.filter(|&i| {
!matches!(i, Index::Block(BlockIndex::Message(_)))
&& index_is_cached(prompt, i)
})
.count();
let n = n.min(MAX_CACHE_CONTROLS_PER_REQUEST.saturating_sub(pinned));
let mut tail: std::collections::HashSet<usize> =
std::collections::HashSet::with_capacity(n);
for k in 0..n {
let Some(m) = anchor.checked_sub(2 * k) else {
break;
};
if !prompt.messages[m].content.has_cache() {
prompt.messages[m].content.cache_with(cache_control.clone());
}
tail.insert(m);
}
let budget = MAX_CACHE_CONTROLS_PER_REQUEST
.saturating_sub(pinned)
.saturating_sub(tail.len());
let stragglers: Vec<Index> = prompt
.indices()
.filter(|&i| match i {
Index::Block(BlockIndex::Message((m, _))) => {
!tail.contains(&m) && index_is_cached(prompt, i)
}
_ => false,
})
.collect();
for &index in stragglers.iter().skip(budget) {
if let Some(IndexMut::Block(block)) = prompt.get_mut(index) {
block.uncache();
}
}
}
fn index_is_cached(prompt: &Prompt, index: Index) -> bool {
use misanthropic::prompt::index::IndexRef;
match prompt.get(index) {
Some(IndexRef::Method(method)) => method.is_cached(),
Some(IndexRef::Block(block)) => block.is_cached(),
None => false,
}
}
pub fn divergence(prev: &Prompt, next: &Prompt) -> Option<String> {
fn wire(value: &impl serde::Serialize) -> serde_json::Value {
fn strip(value: &mut serde_json::Value) {
match value {
serde_json::Value::Object(map) => {
map.remove("cache_control");
map.values_mut().for_each(strip);
}
serde_json::Value::Array(items) => {
items.iter_mut().for_each(strip)
}
_ => {}
}
}
let mut value =
serde_json::to_value(value).expect("a prompt always serializes");
strip(&mut value);
value
}
let heads = [
("tools", wire(&prev.tools), wire(&next.tools)),
("system", wire(&prev.system), wire(&next.system)),
("thinking", wire(&prev.thinking), wire(&next.thinking)),
(
"tool_choice",
wire(&prev.tool_choice),
wire(&next.tool_choice),
),
];
for (field, a, b) in heads {
if a != b {
return Some(format!("`{field}` changed: {a} -> {b}"));
}
}
for (i, a) in prev.messages.iter().enumerate() {
let Some(b) = next.messages.get(i) else {
return Some(format!(
"message {i} of {} dropped",
prev.messages.len()
));
};
let (a, b) = (wire(a), wire(b));
if a != b {
return Some(format!("message {i} changed: {a} -> {b}"));
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use misanthropic::prompt::message::Role;
fn quirks(f: impl FnOnce(&mut Quirks)) -> Quirks {
let mut q = Quirks::default();
f(&mut q);
q
}
fn session(pairs: usize) -> Prompt {
let mut prompt = Prompt::default()
.system("system text")
.add_message((Role::User, "intro"))
.unwrap()
.cache_1h();
prompt.system.as_mut().unwrap().cache_1h();
for i in 0..pairs {
prompt
.push_message((Role::Assistant, format!("asst {i}")))
.unwrap();
prompt
.push_message((Role::User, format!("results {i}")))
.unwrap();
}
prompt
}
fn marked(prompt: &Prompt) -> Vec<usize> {
prompt
.messages
.iter()
.enumerate()
.filter(|(_, m)| m.content.has_cache())
.map(|(i, _)| i)
.collect()
}
fn total_markers(prompt: &Prompt) -> usize {
serde_json::to_string(prompt)
.unwrap()
.matches(r#""cache_control":"#)
.count()
}
#[test]
fn canonical_rolls_onto_user_turns() {
let mut prompt = session(2);
roll_breakpoints(&Quirks::default(), &mut prompt);
assert_eq!(marked(&prompt), vec![0, 2, 4]);
assert_eq!(prompt.messages[2].role, Role::User);
assert_eq!(prompt.messages[4].role, Role::User);
}
#[test]
fn ollama_is_a_no_op() {
let mut prompt = session(2);
let before = total_markers(&prompt);
let q = quirks(|q| q.cache_markers_ignored = true);
roll_breakpoints(&q, &mut prompt);
assert_eq!(total_markers(&prompt), before);
}
#[test]
fn empty_prompt_is_harmless() {
let mut prompt = Prompt::default();
roll_breakpoints(&Quirks::default(), &mut prompt);
assert_eq!(total_markers(&prompt), 0);
}
#[test]
fn round_loop_never_exceeds_budget_or_jumps_role() {
let (q, role) = (Quirks::default(), Role::User);
let mut prompt = session(0);
for i in 0..10 {
prompt
.push_message((Role::Assistant, format!("asst {i}")))
.unwrap();
prompt
.push_message((Role::User, format!("results {i}")))
.unwrap();
roll_breakpoints(&q, &mut prompt);
assert!(
total_markers(&prompt) <= MAX_CACHE_CONTROLS_PER_REQUEST,
"round {i}: {} markers",
total_markers(&prompt)
);
assert!(
prompt.system.as_ref().unwrap().has_cache(),
"round {i}: system marker evicted"
);
assert!(
prompt.messages[0].content.has_cache(),
"round {i}: intro marker evicted"
);
for idx in marked(&prompt).into_iter().skip(1) {
assert_eq!(
prompt.messages[idx].role, role,
"round {i}: rolling marker jumped role at {idx}"
);
}
}
let json = serde_json::to_string(&prompt).unwrap();
assert!(
!json.contains(r#""cache_control":{"type":"ephemeral"}"#),
"5m marker present:\n{json}"
);
}
#[test]
fn re_rolling_without_new_messages_is_idempotent() {
let mut prompt = session(3);
roll_breakpoints(&Quirks::default(), &mut prompt);
let first = marked(&prompt);
roll_breakpoints(&Quirks::default(), &mut prompt);
assert_eq!(marked(&prompt), first);
assert!(total_markers(&prompt) <= MAX_CACHE_CONTROLS_PER_REQUEST);
}
#[test]
fn window_shrinks_before_evicting_prefix_markers() {
let mut prompt = session(3);
prompt.messages[1].content.cache_1h(); roll_breakpoints(&Quirks::default(), &mut prompt);
assert!(total_markers(&prompt) <= MAX_CACHE_CONTROLS_PER_REQUEST);
assert!(prompt.system.as_ref().unwrap().has_cache());
assert!(prompt.messages[0].content.has_cache());
}
#[cfg(feature = "client")]
#[tokio::test]
#[ignore = "hits the live Anthropic API (cents, not dollars)"]
async fn live_roll_breakpoints_hit_the_cache() {
let key = std::env::var("ANTHROPIC_API_KEY").unwrap_or_else(|_| {
let path = format!(
"{}/Projects/agora/secrets/anthropic_api_key",
std::env::var("HOME").expect("HOME")
);
std::fs::read_to_string(path)
.expect("no ANTHROPIC_API_KEY and no key file")
.trim()
.to_string()
});
let client = misanthropic::Client::new(key).expect("client");
let padding: String = (0..500)
.map(|i| {
format!(
"Fact {i}: the {i}th cache line holds a distinct \
sentence so the prefix is long and incompressible.\n"
)
})
.collect();
let mut prompt = Prompt::default()
.max_tokens(std::num::NonZeroU32::new(32).unwrap())
.system(format!(
"You are terse. Reply with a single word.\n\n{padding}"
))
.add_message((Role::User, "Say the word: one."))
.unwrap();
prompt.system.as_mut().unwrap().cache();
let quirks = Quirks::default();
roll_breakpoints_with(&quirks, &mut prompt, CacheControl::ephemeral());
let first = client.message(&prompt).await.expect("round 1");
let wrote = first.usage.cache_creation_input_tokens.unwrap_or(0)
+ first.usage.cache_read_input_tokens.unwrap_or(0);
assert!(
wrote > 0,
"round 1 neither wrote nor read cache (prefix under the \
minimum?): {:?}",
first.usage
);
prompt.push_message(first).unwrap();
prompt
.push_message((Role::User, "Say the word: two."))
.unwrap();
roll_breakpoints_with(&quirks, &mut prompt, CacheControl::ephemeral());
let second = client.message(&prompt).await.expect("round 2");
let read = second.usage.cache_read_input_tokens.unwrap_or(0);
println!(
"round 2 usage: read={read} create={:?} input={}",
second.usage.cache_creation_input_tokens, second.usage.input_tokens
);
assert!(read > 0, "round 2 read nothing: {:?}", second.usage);
}
#[test]
fn divergence_sees_only_what_a_prefix_cache_does() {
let prev = session(1);
let mut next = prev.clone();
next.push_message((Role::Assistant, "asst 1")).unwrap();
next.push_message((Role::User, "results 1")).unwrap();
roll_breakpoints(&Quirks::default(), &mut next);
assert_eq!(divergence(&prev, &next), None);
let mut grown = prev.clone();
grown.messages.last_mut().unwrap().extend(["a note"]);
let why = divergence(&prev, &grown).unwrap();
assert!(why.starts_with("message 2 changed"), "{why}");
let mut dropped = prev.clone();
dropped.messages.pop();
assert!(divergence(&prev, &dropped).unwrap().contains("dropped"));
let other = prev.clone().system("other system");
assert!(divergence(&prev, &other).unwrap().starts_with("`system`"));
}
}