fix:add anthropic cache block

This commit is contained in:
leookun
2026-08-28 10:33:00 +08:00
parent 15f2280e07
commit b15e149b9d
2 changed files with 59 additions and 23 deletions
+30 -9
View File
@@ -51,16 +51,32 @@ impl Provider for AnthropicProvider {
let ModelInvocation { call_id, request, .. } = invocation; let ModelInvocation { call_id, request, .. } = invocation;
let mut messages = anthropic_messages(&request.history)?; let mut messages = anthropic_messages(&request.history)?;
mark_cache_breakpoint(&mut messages); mark_cache_breakpoint(&mut messages);
let system = if request.prompt.instructions.is_empty() {
Value::String(String::new())
} else {
json!([{
"type": "text",
"text": request.prompt.instructions,
"cache_control": {"type": "ephemeral"}
}])
};
let max_tokens = request.model.max_output_tokens.or(config.max_output_tokens) let max_tokens = request.model.max_output_tokens.or(config.max_output_tokens)
.unwrap_or(DEFAULT_MAX_OUTPUT_TOKENS); .unwrap_or(DEFAULT_MAX_OUTPUT_TOKENS);
let mut body = json!({ let mut body = json!({
"model": request.model.model_id, "system": request.prompt.instructions, "messages": messages, "model": request.model.model_id, "system": system, "messages": messages,
"max_tokens": max_tokens, "stream": true "max_tokens": max_tokens, "stream": true
}); });
if !request.prompt.tools.is_empty() { if !request.prompt.tools.is_empty() {
body["tools"] = json!(request.prompt.tools.iter().map(|tool| json!({ let tool_count = request.prompt.tools.len();
body["tools"] = json!(request.prompt.tools.iter().enumerate().map(|(index, tool)| {
let mut value = json!({
"name": tool.name, "description": tool.description, "input_schema": tool.parameters "name": tool.name, "description": tool.description, "input_schema": tool.parameters
})).collect::<Vec<_>>()); });
if index + 1 == tool_count {
value["cache_control"] = json!({"type": "ephemeral"});
}
value
}).collect::<Vec<_>>());
} }
apply_model(&mut body, &request.model)?; apply_model(&mut body, &request.model)?;
merge_extra_params(&mut body, &request.model.extra_params)?; merge_extra_params(&mut body, &request.model.extra_params)?;
@@ -376,13 +392,19 @@ fn anthropic_messages(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
} }
fn mark_cache_breakpoint(messages: &mut [Value]) { fn mark_cache_breakpoint(messages: &mut [Value]) {
for message in messages { let Some(message) = messages
let Some(content) = message.get_mut("content").and_then(Value::as_array_mut) else { .iter_mut()
continue; .rev()
.find(|message| message.get("role").and_then(Value::as_str) == Some("user"))
else {
return;
}; };
for block in content { let Some(content) = message.get_mut("content").and_then(Value::as_array_mut) else {
return;
};
for block in content.iter_mut().rev() {
let kind = block.get("type").and_then(Value::as_str); let kind = block.get("type").and_then(Value::as_str);
if matches!(kind, Some("thinking" | "redacted_thinking")) { if !matches!(kind, Some("text" | "image" | "tool_result")) {
continue; continue;
} }
if let Some(block) = block.as_object_mut() { if let Some(block) = block.as_object_mut() {
@@ -391,7 +413,6 @@ fn mark_cache_breakpoint(messages: &mut [Value]) {
} }
} }
} }
}
fn anthropic_parts(role: &Role, parts: &[ContentPart]) -> Result<Vec<Value>> { fn anthropic_parts(role: &Role, parts: &[ContentPart]) -> Result<Vec<Value>> {
parts parts
+22 -7
View File
@@ -3,7 +3,7 @@ use cursor_server::{
config::{ProviderConfig, ProviderKind}, config::{ProviderConfig, ProviderKind},
model::{ model::{
ContentPart, ModelInvocation, ModelLatency, ModelRequest, ModelSpec, ProjectedContent, ContentPart, ModelInvocation, ModelLatency, ModelRequest, ModelSpec, ProjectedContent,
ProjectedMessage, PromptSpec, Role, Usage, ProjectedMessage, PromptSpec, Role, ToolDefinition, Usage,
}, },
provider::{ provider::{
FinishReason, ModelEvent, OpenAiChatProvider, OpenAiResponsesProvider, Provider, FinishReason, ModelEvent, OpenAiChatProvider, OpenAiResponsesProvider, Provider,
@@ -12,7 +12,7 @@ use cursor_server::{
run::{consume_model_cycle, RunFailure}, run::{consume_model_cycle, RunFailure},
}; };
use futures_util::{stream, StreamExt}; use futures_util::{stream, StreamExt};
use serde_json::Value; use serde_json::{json, Value};
use std::{sync::Arc, time::Duration}; use std::{sync::Arc, time::Duration};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
@@ -658,6 +658,18 @@ async fn anthropic_raw_stream_uses_explicit_and_default_token_limits() {
); );
let mut request = invocation(); let mut request = invocation();
request.request.model.latency = ModelLatency::Fast; request.request.model.latency = ModelLatency::Fast;
request.request.prompt.tools = vec![
ToolDefinition {
name: "first".into(),
description: "first tool".into(),
parameters: json!({"type":"object"}),
},
ToolDefinition {
name: "last".into(),
description: "last tool".into(),
parameters: json!({"type":"object"}),
},
];
let events = collect(provider.stream(request, CancellationToken::new())).await; let events = collect(provider.stream(request, CancellationToken::new())).await;
let body = requests.recv().await.unwrap(); let body = requests.recv().await.unwrap();
let continued = continued_invocation(&events); let continued = continued_invocation(&events);
@@ -673,16 +685,19 @@ async fn anthropic_raw_stream_uses_explicit_and_default_token_limits() {
assert_eq!(body["max_tokens"], 1234); assert_eq!(body["max_tokens"], 1234);
assert!(body.get("cache_control").is_none()); assert!(body.get("cache_control").is_none());
assert_eq!(body["system"][0]["cache_control"]["type"], "ephemeral");
assert!(body["tools"][0].get("cache_control").is_none());
assert_eq!(body["tools"][1]["cache_control"]["type"], "ephemeral");
assert!(body["messages"][0]["content"][0]
.get("cache_control")
.is_none());
assert_eq!( assert_eq!(
body["messages"][0]["content"][0]["cache_control"]["type"], body["messages"][0]["content"][1]["cache_control"]["type"],
"ephemeral" "ephemeral"
); );
assert_eq!(default_body["max_tokens"], 65_000); assert_eq!(default_body["max_tokens"], 65_000);
assert!(default_body.get("cache_control").is_none()); assert!(default_body.get("cache_control").is_none());
assert_eq!( assert_eq!(default_body["messages"][0]["content"][0]["type"], "text");
default_body["messages"][0]["content"][0]["cache_control"]["type"],
"ephemeral"
);
assert!(body.get("service_tier").is_none()); assert!(body.get("service_tier").is_none());
assert_eq!(body["messages"][0]["content"][1]["type"], "image"); assert_eq!(body["messages"][0]["content"][1]["type"], "image");
assert_eq!( assert_eq!(