From b15e149b9d79d71594efc0f3efd4906ee0d33a1b Mon Sep 17 00:00:00 2001 From: leookun Date: Fri, 28 Aug 2026 10:33:00 +0800 Subject: [PATCH] =?UTF-8?q?fix=EF=BC=9Aadd=20anthropic=20cache=20block?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- server/src/provider/anthropic.rs | 53 ++++++++++++++++++++++---------- server/tests/provider_stream.rs | 29 ++++++++++++----- 2 files changed, 59 insertions(+), 23 deletions(-) diff --git a/server/src/provider/anthropic.rs b/server/src/provider/anthropic.rs index dcce5a6..1a50d22 100644 --- a/server/src/provider/anthropic.rs +++ b/server/src/provider/anthropic.rs @@ -51,16 +51,32 @@ impl Provider for AnthropicProvider { let ModelInvocation { call_id, request, .. } = invocation; let mut messages = anthropic_messages(&request.history)?; 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) .unwrap_or(DEFAULT_MAX_OUTPUT_TOKENS); 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 }); if !request.prompt.tools.is_empty() { - body["tools"] = json!(request.prompt.tools.iter().map(|tool| json!({ - "name": tool.name, "description": tool.description, "input_schema": tool.parameters - })).collect::>()); + 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 + }); + if index + 1 == tool_count { + value["cache_control"] = json!({"type": "ephemeral"}); + } + value + }).collect::>()); } apply_model(&mut body, &request.model)?; merge_extra_params(&mut body, &request.model.extra_params)?; @@ -376,19 +392,24 @@ fn anthropic_messages(messages: &[ProjectedMessage]) -> Result> { } fn mark_cache_breakpoint(messages: &mut [Value]) { - for message in messages { - let Some(content) = message.get_mut("content").and_then(Value::as_array_mut) else { + let Some(message) = messages + .iter_mut() + .rev() + .find(|message| message.get("role").and_then(Value::as_str) == Some("user")) + else { + return; + }; + 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); + if !matches!(kind, Some("text" | "image" | "tool_result")) { continue; - }; - for block in content { - let kind = block.get("type").and_then(Value::as_str); - if matches!(kind, Some("thinking" | "redacted_thinking")) { - continue; - } - if let Some(block) = block.as_object_mut() { - block.insert("cache_control".into(), json!({"type": "ephemeral"})); - return; - } + } + if let Some(block) = block.as_object_mut() { + block.insert("cache_control".into(), json!({"type": "ephemeral"})); + return; } } } diff --git a/server/tests/provider_stream.rs b/server/tests/provider_stream.rs index 1eff2ba..991cfa8 100644 --- a/server/tests/provider_stream.rs +++ b/server/tests/provider_stream.rs @@ -3,7 +3,7 @@ use cursor_server::{ config::{ProviderConfig, ProviderKind}, model::{ ContentPart, ModelInvocation, ModelLatency, ModelRequest, ModelSpec, ProjectedContent, - ProjectedMessage, PromptSpec, Role, Usage, + ProjectedMessage, PromptSpec, Role, ToolDefinition, Usage, }, provider::{ FinishReason, ModelEvent, OpenAiChatProvider, OpenAiResponsesProvider, Provider, @@ -12,7 +12,7 @@ use cursor_server::{ run::{consume_model_cycle, RunFailure}, }; use futures_util::{stream, StreamExt}; -use serde_json::Value; +use serde_json::{json, Value}; use std::{sync::Arc, time::Duration}; use tokio_util::sync::CancellationToken; @@ -658,6 +658,18 @@ async fn anthropic_raw_stream_uses_explicit_and_default_token_limits() { ); let mut request = invocation(); 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 body = requests.recv().await.unwrap(); 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!(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!( - body["messages"][0]["content"][0]["cache_control"]["type"], + body["messages"][0]["content"][1]["cache_control"]["type"], "ephemeral" ); assert_eq!(default_body["max_tokens"], 65_000); assert!(default_body.get("cache_control").is_none()); - assert_eq!( - default_body["messages"][0]["content"][0]["cache_control"]["type"], - "ephemeral" - ); + assert_eq!(default_body["messages"][0]["content"][0]["type"], "text"); assert!(body.get("service_tier").is_none()); assert_eq!(body["messages"][0]["content"][1]["type"], "image"); assert_eq!(