mirror of
https://github.com/block/goose.git
synced 2026-07-17 12:56:20 +02:00
feat: add caching support for Anthropic models via Databricks (#6320)
This commit is contained in:
@@ -407,6 +407,92 @@ struct StreamingChunk {
|
||||
model: String,
|
||||
}
|
||||
|
||||
/// Check if the model name indicates a Claude/Anthropic model that supports cache control.
|
||||
fn is_claude_model(model_name: &str) -> bool {
|
||||
model_name.contains("claude")
|
||||
}
|
||||
|
||||
/// Add Anthropic-style cache_control fields to the request payload for Claude models.
|
||||
/// This enables prompt caching to reduce costs when using Claude via Databricks.
|
||||
///
|
||||
/// Cache control is added to:
|
||||
/// - The system message
|
||||
/// - The last two user messages (for incremental caching across turns)
|
||||
/// - The last tool definition (so all tools are cached as a single prefix)
|
||||
pub fn apply_cache_control_for_claude(payload: &mut Value) {
|
||||
if let Some(messages_spec) = payload
|
||||
.as_object_mut()
|
||||
.and_then(|obj| obj.get_mut("messages"))
|
||||
.and_then(|messages| messages.as_array_mut())
|
||||
{
|
||||
// Add cache_control to the last two user messages for incremental caching.
|
||||
// The last message gets cached so future turns can read from it.
|
||||
// The second-to-last user message is also cached to read from the previous cache.
|
||||
let mut user_count = 0;
|
||||
for message in messages_spec.iter_mut().rev() {
|
||||
if message.get("role") == Some(&json!("user")) {
|
||||
if let Some(content) = message.get_mut("content") {
|
||||
if let Some(content_str) = content.as_str() {
|
||||
*content = json!([{
|
||||
"type": "text",
|
||||
"text": content_str,
|
||||
"cache_control": { "type": "ephemeral" }
|
||||
}]);
|
||||
} else if let Some(content_array) = content.as_array_mut() {
|
||||
// Content is already an array, add cache_control to the last element
|
||||
if let Some(last_content) = content_array.last_mut() {
|
||||
if let Some(obj) = last_content.as_object_mut() {
|
||||
obj.insert(
|
||||
"cache_control".to_string(),
|
||||
json!({ "type": "ephemeral" }),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
user_count += 1;
|
||||
if user_count >= 2 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Add cache_control to the system message
|
||||
if let Some(system_message) = messages_spec
|
||||
.iter_mut()
|
||||
.find(|msg| msg.get("role") == Some(&json!("system")))
|
||||
{
|
||||
if let Some(content) = system_message.get_mut("content") {
|
||||
if let Some(content_str) = content.as_str() {
|
||||
*system_message = json!({
|
||||
"role": "system",
|
||||
"content": [{
|
||||
"type": "text",
|
||||
"text": content_str,
|
||||
"cache_control": { "type": "ephemeral" }
|
||||
}]
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Add cache_control to the last tool definition
|
||||
if let Some(tools_spec) = payload
|
||||
.as_object_mut()
|
||||
.and_then(|obj| obj.get_mut("tools"))
|
||||
.and_then(|tools| tools.as_array_mut())
|
||||
{
|
||||
if let Some(last_tool) = tools_spec.last_mut() {
|
||||
if let Some(function) = last_tool.get_mut("function") {
|
||||
if let Some(obj) = function.as_object_mut() {
|
||||
obj.insert("cache_control".to_string(), json!({ "type": "ephemeral" }));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Validates and fixes tool schemas to ensure they have proper parameter structure.
|
||||
/// If parameters exist, ensures they have properties and required fields, or removes parameters entirely.
|
||||
pub fn validate_tool_schemas(tools: &mut [Value]) {
|
||||
@@ -589,6 +675,11 @@ pub fn create_request(
|
||||
}
|
||||
}
|
||||
|
||||
// Apply cache control for Claude models to enable prompt caching
|
||||
if is_claude_model(&model_config.model_name) {
|
||||
apply_cache_control_for_claude(&mut payload);
|
||||
}
|
||||
|
||||
Ok(payload)
|
||||
}
|
||||
|
||||
@@ -1149,4 +1240,241 @@ mod tests {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_claude_model() {
|
||||
assert!(is_claude_model("databricks-claude-sonnet-4"));
|
||||
assert!(is_claude_model("databricks-claude-3-7-sonnet"));
|
||||
assert!(is_claude_model("claude-sonnet-4"));
|
||||
assert!(is_claude_model("goose-claude-sonnet"));
|
||||
assert!(!is_claude_model("gpt-4o"));
|
||||
assert!(!is_claude_model("gemini-2-5-flash"));
|
||||
assert!(!is_claude_model("databricks-meta-llama-3-3-70b"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_apply_cache_control_for_claude_system_message() -> anyhow::Result<()> {
|
||||
let mut payload = json!({
|
||||
"model": "databricks-claude-sonnet-4",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a helpful assistant."
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello"
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
apply_cache_control_for_claude(&mut payload);
|
||||
|
||||
let messages = payload["messages"].as_array().unwrap();
|
||||
let system_msg = &messages[0];
|
||||
|
||||
// System message content should be converted to array with cache_control
|
||||
assert!(system_msg["content"].is_array());
|
||||
let content = system_msg["content"].as_array().unwrap();
|
||||
assert_eq!(content.len(), 1);
|
||||
assert_eq!(content[0]["type"], "text");
|
||||
assert_eq!(content[0]["text"], "You are a helpful assistant.");
|
||||
assert_eq!(content[0]["cache_control"]["type"], "ephemeral");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_apply_cache_control_for_claude_user_messages() -> anyhow::Result<()> {
|
||||
let mut payload = json!({
|
||||
"model": "databricks-claude-sonnet-4",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are helpful"
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "First question"
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "First answer"
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Second question"
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Second answer"
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Third question"
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
apply_cache_control_for_claude(&mut payload);
|
||||
|
||||
let messages = payload["messages"].as_array().unwrap();
|
||||
|
||||
// First user message should NOT have cache_control (only last 2)
|
||||
let first_user = &messages[1];
|
||||
assert_eq!(first_user["content"], "First question");
|
||||
|
||||
// Second-to-last user message should have cache_control
|
||||
let second_user = &messages[3];
|
||||
assert!(second_user["content"].is_array());
|
||||
assert_eq!(
|
||||
second_user["content"][0]["cache_control"]["type"],
|
||||
"ephemeral"
|
||||
);
|
||||
|
||||
// Last user message should have cache_control
|
||||
let last_user = &messages[5];
|
||||
assert!(last_user["content"].is_array());
|
||||
assert_eq!(
|
||||
last_user["content"][0]["cache_control"]["type"],
|
||||
"ephemeral"
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_apply_cache_control_for_claude_tools() -> anyhow::Result<()> {
|
||||
let mut payload = json!({
|
||||
"model": "databricks-claude-sonnet-4",
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are helpful"
|
||||
}
|
||||
],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "tool1",
|
||||
"description": "First tool"
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "tool2",
|
||||
"description": "Second tool"
|
||||
}
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
apply_cache_control_for_claude(&mut payload);
|
||||
|
||||
let tools = payload["tools"].as_array().unwrap();
|
||||
|
||||
// First tool should NOT have cache_control
|
||||
assert!(tools[0]["function"].get("cache_control").is_none());
|
||||
|
||||
// Last tool should have cache_control
|
||||
assert_eq!(tools[1]["function"]["cache_control"]["type"], "ephemeral");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_request_claude_has_cache_control() -> anyhow::Result<()> {
|
||||
let model_config = ModelConfig {
|
||||
model_name: "databricks-claude-sonnet-4".to_string(),
|
||||
context_limit: Some(200000),
|
||||
temperature: None,
|
||||
max_tokens: Some(8192),
|
||||
toolshim: false,
|
||||
toolshim_model: None,
|
||||
fast_model: None,
|
||||
};
|
||||
|
||||
let messages = vec![
|
||||
Message::user().with_text("Hello"),
|
||||
Message::assistant().with_text("Hi there!"),
|
||||
Message::user().with_text("How are you?"),
|
||||
];
|
||||
|
||||
let tool = Tool::new(
|
||||
"test_tool",
|
||||
"A test tool",
|
||||
object!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}),
|
||||
);
|
||||
|
||||
let request = create_request(
|
||||
&model_config,
|
||||
"You are helpful",
|
||||
&messages,
|
||||
&[tool],
|
||||
&ImageFormat::OpenAi,
|
||||
)?;
|
||||
|
||||
// Verify system message has cache_control
|
||||
let messages_arr = request["messages"].as_array().unwrap();
|
||||
let system_msg = &messages_arr[0];
|
||||
assert!(system_msg["content"].is_array());
|
||||
assert_eq!(
|
||||
system_msg["content"][0]["cache_control"]["type"],
|
||||
"ephemeral"
|
||||
);
|
||||
|
||||
// Verify last tool has cache_control
|
||||
let tools = request["tools"].as_array().unwrap();
|
||||
assert_eq!(tools[0]["function"]["cache_control"]["type"], "ephemeral");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_request_non_claude_no_cache_control() -> anyhow::Result<()> {
|
||||
let model_config = ModelConfig {
|
||||
model_name: "gpt-4o".to_string(),
|
||||
context_limit: Some(128000),
|
||||
temperature: None,
|
||||
max_tokens: Some(4096),
|
||||
toolshim: false,
|
||||
toolshim_model: None,
|
||||
fast_model: None,
|
||||
};
|
||||
|
||||
let messages = vec![Message::user().with_text("Hello")];
|
||||
|
||||
let tool = Tool::new(
|
||||
"test_tool",
|
||||
"A test tool",
|
||||
object!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}),
|
||||
);
|
||||
|
||||
let request = create_request(
|
||||
&model_config,
|
||||
"You are helpful",
|
||||
&messages,
|
||||
&[tool],
|
||||
&ImageFormat::OpenAi,
|
||||
)?;
|
||||
|
||||
// Verify system message does NOT have cache_control (it's a plain string)
|
||||
let messages_arr = request["messages"].as_array().unwrap();
|
||||
let system_msg = &messages_arr[0];
|
||||
assert!(system_msg["content"].is_string());
|
||||
|
||||
// Verify tool does NOT have cache_control
|
||||
let tools = request["tools"].as_array().unwrap();
|
||||
assert!(tools[0]["function"].get("cache_control").is_none());
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user