use misanthropic::prompt::{
Prompt,
index::{BlockIndex, Index, IndexMut},
message::{CacheControl, Role},
};
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 anchor = if quirks.breakpoint_after_assistant {
prompt
.messages
.iter()
.rposition(|m| m.role == Role::Assistant)
} else {
prompt.messages.len().checked_sub(1)
};
let Some(anchor) = anchor 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,
}
}
#[cfg(test)]
mod tests {
use super::*;
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 blallama_rolls_onto_assistant_turns() {
let mut prompt = session(2);
let q = quirks(|q| q.breakpoint_after_assistant = true);
roll_breakpoints(&q, &mut prompt);
assert_eq!(marked(&prompt), vec![0, 1, 3]);
assert_eq!(prompt.messages[1].role, Role::Assistant);
assert_eq!(prompt.messages[3].role, Role::Assistant);
}
#[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;
q.breakpoint_after_assistant = true;
});
roll_breakpoints(&q, &mut prompt);
assert_eq!(total_markers(&prompt), before);
}
#[test]
fn blallama_skips_until_an_assistant_turn_exists() {
let mut prompt = session(0);
let q = quirks(|q| q.breakpoint_after_assistant = true);
roll_breakpoints(&q, &mut prompt);
assert_eq!(marked(&prompt), vec![0], "prefix marker only");
}
#[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() {
for (q, role) in [
(Quirks::default(), Role::User),
(
quirks(|q| q.breakpoint_after_assistant = true),
Role::Assistant,
),
] {
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);
}
}