Files
agcore/src/llm/provider/anthropic.rs
T
徐涛 5e475e1303 docs(llm): 清理 build_request_builder 重复 doc comment
- anthropic.rs:删除 build_request_builder doc comment 中重复的 4 行(Round 1 安全性修复时追加内容未清理原段落)
- openai_response.rs:在 build_request_builder doc comment 补充「构造 HTTP POST 请求 builder(含认证头与额外请求头)」描述句,与另两 provider 对齐
2026-07-20 15:09:21 +08:00

1529 lines
55 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Anthropic Messages API Provider —— Phase 1 引入。
//!
//! Anthropic SSE 事件序列:`message_start` → `content_block_start` → `content_block_delta`
//! → `content_block_stop` → `message_delta` → `message_stop`。与 OpenAI 不同,
//! Anthropic 提供显式 block 边界事件,状态机相对简单。
use std::collections::HashMap;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
use async_trait::async_trait;
use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use tracing::{debug, error, info, warn};
use crate::llm::error::LlmError;
use crate::llm::types::message::{ContentBlock, ContentBlockType, Message};
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{
MessageResponse, PartialMessageResponse, PartialUsage, StopReason, StreamEvent,
};
use crate::llm::types::usage::Usage;
use crate::llm::{LlmProvider, ProviderCapabilities, ProviderFeatures};
/// Anthropic Provider 默认 `max_tokens` 兜底值。
///
/// Anthropic Messages API 要求 `max_tokens` 为必填字段,`MessageRequest.max_tokens`
/// 为 `Option<u32>`。未设置时使用此默认值。
const DEFAULT_MAX_TOKENS: u32 = 4096;
pub struct AnthropicProvider {
http_client: Client,
base_url: String,
#[allow(dead_code)]
api_key: String,
model: String,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
/// Provider 级别固定请求头(如平台标识头),所有请求自动携带。
extra_headers: Vec<(String, String)>,
}
impl AnthropicProvider {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let key_header =
HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
let version_header = HeaderValue::from_static("2023-06-01");
let http_client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.expect("创建 HTTP 客户端失败");
Self {
http_client,
base_url: if base_url.is_empty() {
"https://api.anthropic.com".to_string()
} else {
base_url
},
api_key,
model,
timeout_secs,
extra_headers: Vec::new(),
}
}
/// ⚠️ 替换 HTTP Client**丢弃** `new()` 中设置的默认 headers`x-api-key` / `anthropic-version`)。
///
/// 调用此方法后,所有 Anthropic API 请求将以**无认证头**发送出去,预期会 401/403 失败。
/// 推荐改用 [`Self::with_timeout`],它会重建 client 并保留默认 headers。
///
/// 此方法仍保留以兼容调用方自定义 client 但不需要默认 headers 的极端场景。
#[deprecated(
since = "0.2.0",
note = "此方法会丢弃默认 headersx-api-key / anthropic-version),改为使用 `with_timeout` 或带 headers 的 `Client::builder()`"
)]
pub fn with_client(mut self, client: Client) -> Self {
self.http_client = client;
self
}
/// 替换 HTTP Client 的超时配置(重建底层 client,保留默认 headers)。
///
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后通过 `with_client` 注入的 Client
/// 将被替换。headers 构造逻辑与 `new()` 中的保持一致(`x-api-key` / `anthropic-version`)。
///
/// ponytail: 同值调用短路。当 `secs == self.timeout_secs` 时跳过 client 重建,
/// 避免 `create_provider` 路径 `new(timeout).with_timeout(timeout)` 的双重构造。
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
if secs == self.timeout_secs {
return Ok(self);
}
// ponytail: 重建 http_client 时保留已有默认 headersx-api-key / anthropic-version)。
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
let key_header = HeaderValue::from_str(&self.api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
self.http_client = Client::builder()
.timeout(Duration::from_secs(secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
self.timeout_secs = secs;
Ok(self)
}
/// 一次性构造 —— `create_provider` 路径专用,避免 `new(...)` + `with_timeout(...)` 的双重 client 构造。
///
/// 调用方负责预先构造好符合 Anthropic 协议要求的 `http_client`(带正确的 `x-api-key` /
/// `anthropic-version` 默认 headers + 指定 timeout)。
pub(crate) fn from_parts(
base_url: String,
api_key: String,
model: String,
http_client: Client,
timeout_secs: u64,
extra_headers: Vec<(String, String)>,
) -> Self {
Self {
http_client,
base_url: if base_url.is_empty() {
"https://api.anthropic.com".to_string()
} else {
base_url
},
api_key,
model,
timeout_secs,
extra_headers,
}
}
/// 设置 Provider 级别固定头,替换已有的 extra_headers(如有)。
/// 返回 self 以支持链式调用。如需追加语义,在外部自行 `extend`。
pub fn with_extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
self.extra_headers = headers;
self
}
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
}
/// `MessageRequest` → Anthropic Messages API 请求体。
fn build_request_body(
&self,
request: MessageRequest,
) -> Result<AnthropicRequestBody, LlmError> {
let mut system_prompts: Vec<String> = Vec::new();
let mut api_messages: Vec<AnthropicMessage> = Vec::new();
for msg in request.messages.iter() {
match msg {
Message::System { content } => {
// 提取所有 Text block 拼为单个 system prompt
let text: String = content
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(text.as_str()),
_ => None,
})
.collect::<Vec<&str>>()
.join("\n");
system_prompts.push(text);
}
Message::User { content } => {
api_messages.push(AnthropicMessage::user(content));
}
Message::UserImage {
data,
mime_type,
detail,
} => {
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
let source = if data.starts_with("http://") || data.starts_with("https://") {
AnthropicImageSource::Url { url: data.clone() }
} else {
AnthropicImageSource::Base64 {
media_type: mime_type.clone(),
data: data.clone(),
}
};
let _ = detail; // 当前忽略 ImageDetailAnthropic 不支持
api_messages.push(AnthropicMessage::User {
content: vec![AnthropicContentPart::Image { source }],
});
}
Message::Assistant { content } => {
api_messages.push(AnthropicMessage::assistant(content));
}
Message::ToolResult {
tool_call_id,
content,
is_error,
} => {
api_messages.push(AnthropicMessage::User {
content: vec![AnthropicContentPart::ToolResult {
tool_use_id: tool_call_id.clone(),
content: serialize_tool_result_content(content),
is_error: *is_error,
}],
});
}
}
}
let max_tokens = self.resolve_max_tokens(&request);
// ponytail: 提前抽取 custom_headers,避免后续 into_iter 消耗 request.tools 后借用失败。
let custom_headers: HashMap<String, String> =
request.get_extra_opt("custom_headers").unwrap_or_default();
let tools = if request.tools.is_empty() {
None
} else {
Some(
request
.tools
.into_iter()
.map(|t| AnthropicToolDefinition {
name: t.name.clone(),
description: t.description.clone().unwrap_or_default(),
input_schema: t.parameters.clone(),
})
.collect(),
)
};
let thinking = request.thinking.clone().map(|tc| AnthropicThinking {
ty: "enabled",
budget_tokens: tc.budget_tokens,
});
Ok(AnthropicRequestBody {
model: if request.model.is_empty() {
self.model.clone()
} else {
request.model
},
max_tokens,
system: if system_prompts.is_empty() {
None
} else {
Some(system_prompts.join("\n"))
},
messages: api_messages,
tools,
thinking,
stream: if request.stream { Some(true) } else { None },
custom_headers,
})
}
/// 构造 HTTP POST 请求 builder(含认证头 + 自定义头)。
/// 认证头(x-api-key / anthropic-version)已由 Client 的 default_headers 提供。
///
/// 头融合顺序:认证头(default_headers)→ Provider 级 extra_headers → 请求级 custom_headers
/// 后者覆盖前者。非法 header 名/值(如控制字符)静默跳过 + warn,避免 reqwest panic。
fn build_request_builder(
&self,
body: &AnthropicRequestBody,
) -> Result<reqwest::RequestBuilder, LlmError> {
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
let mut builder = self.http_client.post(&url).json(body);
for (k, v) in &self.extra_headers {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(k.as_bytes()),
HeaderValue::from_str(v),
) {
builder = builder.header(name, value);
} else {
warn!(header = %k, "skipping invalid extra_header (key or value contains illegal characters)");
}
}
for (key, value) in &body.custom_headers {
if let (Ok(name), Ok(value)) = (
HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(value),
) {
builder = builder.header(name, value);
} else {
warn!(header = %key, "skipping invalid custom_header (key or value contains illegal characters)");
}
}
Ok(builder)
}
async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
let body = self.build_request_body(request)?;
info!(model = %body.model, "Anthropic: 发送非流式请求");
let response = self
.build_request_builder(&body)?
.send()
.await
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let body_text = response.text().await.unwrap_or_default();
debug!(body = %body_text, "Anthropic: 收到响应体");
let resp: AnthropicResponseBody = serde_json::from_str(&body_text).map_err(|e| {
error!(error = %e, "Anthropic 响应解析失败");
LlmError::Other(format!("响应解析失败: {e}"))
})?;
Ok(Self::convert_response(resp))
}
async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let mut body = self.build_request_body(request)?;
body.stream = Some(true);
info!(model = %body.model, "Anthropic: 发送流式请求");
let response = self
.build_request_builder(&body)?
.send()
.await
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let byte_stream = response
.bytes_stream()
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
Box::pin(byte_stream);
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
}
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(self.timeout_secs),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {e}"))
} else {
LlmError::Other(format!("请求失败: {e}"))
}
}
/// Anthropic 错误映射:529overloaded)→ RateLimit401/403 → Authentication。
async fn handle_error_response(response: reqwest::Response) -> LlmError {
let status = response.status().as_u16();
// ponytail: 尝试解析 Anthropic 风格 error body`{error: {type, message}}`
// 但即使解析失败也用 HTTP status + 原始 body 兜底。
let retry_after = response
.headers()
.get("retry-after")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok())
.map(Duration::from_secs);
let body = response.text().await.unwrap_or_default();
match status {
401 | 403 => LlmError::Authentication(body),
429 => LlmError::RateLimit { retry_after },
529 => LlmError::RateLimit { retry_after: None },
_ => LlmError::Request { status, body },
}
}
/// Anthropic 非流式响应 → `MessageResponse`。
fn convert_response(resp: AnthropicResponseBody) -> MessageResponse {
let mut blocks: Vec<ContentBlock> = Vec::new();
for part in resp.content {
match part {
AnthropicContentBlockResp::Text { text } => {
blocks.push(ContentBlock::Text { text });
}
AnthropicContentBlockResp::ToolUse { id, name, input } => {
blocks.push(ContentBlock::ToolUse { id, name, input });
}
AnthropicContentBlockResp::Thinking {
thinking,
signature,
} => {
blocks.push(ContentBlock::Thinking {
text: thinking,
signature,
});
}
}
}
let stop_reason = match resp.stop_reason.as_deref() {
Some("end_turn") => StopReason::Stop,
Some("max_tokens") => StopReason::MaxTokens,
Some("tool_use") => StopReason::ToolUse,
Some("stop_sequence") => StopReason::StopSequence,
_ => StopReason::Stop,
};
let usage = Usage::from_input_output(resp.usage.input_tokens, resp.usage.output_tokens);
MessageResponse {
id: resp.id,
model: resp.model,
message: Message::Assistant { content: blocks },
usage,
stop_reason,
extra: Default::default(),
}
}
}
#[async_trait]
impl LlmProvider for AnthropicProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.chat_blocking(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.chat_stream_inner(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities {
provider_name: "anthropic",
supported_models: Some(vec![
"claude-sonnet-4-20250514".into(),
"claude-3-5-sonnet-20241022".into(),
]),
features: ProviderFeatures {
streaming: true,
thinking: true,
vision: true,
audio_input: false,
tool_use: true,
parallel_tool_calls: false,
system_prompt_in_messages: false,
max_context_window: 200_000,
},
}
}
}
// =============================================================================
// Anthropic 协议类型(私有)
// =============================================================================
#[derive(Debug, Serialize)]
struct AnthropicRequestBody {
model: String,
max_tokens: u32,
#[serde(skip_serializing_if = "Option::is_none")]
system: Option<String>,
messages: Vec<AnthropicMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
tools: Option<Vec<AnthropicToolDefinition>>,
#[serde(skip_serializing_if = "Option::is_none")]
thinking: Option<AnthropicThinking>,
#[serde(skip_serializing_if = "Option::is_none")]
stream: Option<bool>,
/// 请求级别自定义 HTTP 头。运行时注入,不进入 JSON 请求体。
#[serde(skip)]
custom_headers: HashMap<String, String>,
}
#[derive(Debug, Serialize)]
struct AnthropicThinking {
ty: &'static str, // "enabled"
budget_tokens: u32,
}
#[derive(Debug, Serialize)]
struct AnthropicToolDefinition {
name: String,
description: String,
input_schema: Value,
}
#[derive(Debug, Serialize)]
#[serde(tag = "role", rename_all = "lowercase")]
enum AnthropicMessage {
User { content: Vec<AnthropicContentPart> },
Assistant { content: Vec<AnthropicContentPart> },
}
impl AnthropicMessage {
fn user(content: &[ContentBlock]) -> Self {
AnthropicMessage::User {
content: content_to_parts(content),
}
}
fn assistant(content: &[ContentBlock]) -> Self {
AnthropicMessage::Assistant {
content: content_to_parts(content),
}
}
}
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentPart {
Text {
text: String,
},
Image {
source: AnthropicImageSource,
},
ToolUse {
id: String,
name: String,
input: Value,
},
ToolResult {
tool_use_id: String,
content: Value,
is_error: bool,
},
}
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicImageSource {
Base64 { media_type: String, data: String },
Url { url: String },
}
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(AnthropicContentPart::Text { text: text.clone() }),
ContentBlock::Image { source } => {
let api_source = if source.is_url {
AnthropicImageSource::Url {
url: source.data.clone(),
}
} else {
AnthropicImageSource::Base64 {
media_type: source.mime_type.clone(),
data: source.data.clone(),
}
};
Some(AnthropicContentPart::Image { source: api_source })
}
ContentBlock::ToolUse { id, name, input } => Some(AnthropicContentPart::ToolUse {
id: id.clone(),
name: name.clone(),
input: input.clone(),
}),
ContentBlock::Thinking { text, signature: _ } => {
// ponytail: 暂未映射 Thinking 到 Anthropic wireAnthropic 自动产生)。
// Phase 2 后通过专门的 thinking 支持扩展。
if text.is_empty() {
None
} else {
Some(AnthropicContentPart::Text { text: text.clone() })
}
}
// ponytail: Audio / File / ToolResult / Extension 在 Anthropic 上
// 由 ToolResult 单独通过 AnthropicMessage 路径处理,不走 content_to_parts。
// 这里保持非穷举的兜底匹配。
_ => None,
})
.collect()
}
fn serialize_tool_result_content(blocks: &[ContentBlock]) -> Value {
// Anthropic tool_result.content 接受 string 或 array of content blocks。
// 简化:单文本 → string;多块或非文本 → array。
if let [ContentBlock::Text { text }] = blocks {
return Value::String(text.clone());
}
let arr: Vec<Value> = blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(json!({"type": "text", "text": text})),
_ => None,
})
.collect();
Value::Array(arr)
}
#[derive(Debug, Deserialize)]
struct AnthropicResponseBody {
id: String,
#[allow(dead_code)]
#[serde(rename = "type")]
ty: String,
model: String,
content: Vec<AnthropicContentBlockResp>,
stop_reason: Option<String>,
usage: AnthropicUsage,
}
#[derive(Debug, Deserialize)]
struct AnthropicUsage {
input_tokens: u32,
output_tokens: u32,
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentBlockResp {
Text {
text: String,
},
ToolUse {
id: String,
name: String,
input: Value,
},
Thinking {
thinking: String,
signature: Option<String>,
},
}
// =============================================================================
// SSE 流 → StreamEvent 状态机
// =============================================================================
/// Anthropic SSE 事件。
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicSseEvent {
MessageStart {
message: AnthropicMessageStart,
},
ContentBlockStart {
index: u32,
content_block: AnthropicContentBlockStart,
},
ContentBlockDelta {
index: u32,
delta: AnthropicDelta,
},
ContentBlockStop {
index: u32,
},
MessageDelta {
delta: AnthropicMessageDeltaInner,
usage: Option<AnthropicMessageDeltaUsage>,
},
MessageStop,
Ping,
#[serde(other)]
Unknown,
}
#[derive(Debug, Deserialize)]
struct AnthropicMessageStart {
id: String,
model: String,
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentBlockStart {
Text { text: String },
ToolUse { id: String, name: String },
Thinking { thinking: String },
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicDelta {
TextDelta { text: String },
InputJson { partial_json: String },
ThinkingDelta { thinking: String },
SignatureDelta { signature: String },
}
#[derive(Debug, Deserialize)]
struct AnthropicMessageDeltaInner {
stop_reason: Option<String>,
#[serde(default)]
#[allow(dead_code)]
stop_sequence: Option<String>,
}
#[derive(Debug, Deserialize)]
struct AnthropicMessageDeltaUsage {
output_tokens: u32,
input_tokens: Option<u32>,
}
pub struct AnthropicSseStream {
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
buffer: String,
partial: PartialMessageResponse,
#[allow(dead_code)]
next_block_index: u32,
saw_terminal: bool,
}
impl AnthropicSseStream {
fn new(chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>) -> Self {
Self {
chunks,
buffer: String::new(),
partial: PartialMessageResponse::new(),
next_block_index: 0,
saw_terminal: false,
}
}
fn next_event_line(&mut self) -> Option<String> {
// Anthropic SSE 多行:event: ... / data: ... 之间有空行(清空字段)。
// 简化:每次找出 `\n` 切分后找下一个 `data:` 行。
while let Some(pos) = self.buffer.find('\n') {
let line: String = self.buffer.drain(..=pos).collect::<String>();
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
if let Some(data) = trimmed.strip_prefix("data: ") {
return Some(data.to_string());
}
// 跳 event: 行 / 空注释
}
None
}
fn handle_event_json(&mut self, data: &str) -> Vec<StreamEvent> {
let event: AnthropicSseEvent = match serde_json::from_str(data) {
Ok(e) => e,
Err(err) => {
warn!(error = %err, raw = %data, "Anthropic SSE event parse failed");
return vec![StreamEvent::Error {
message: format!("SSE 解析失败: {err}"),
}];
}
};
let mut events = Vec::new();
match event {
AnthropicSseEvent::MessageStart { message } => {
events.push(StreamEvent::MessageStart {
id: message.id.clone(),
model: message.model.clone(),
});
self.partial.id = Some(message.id);
self.partial.model = Some(message.model);
}
AnthropicSseEvent::ContentBlockStart {
index,
content_block,
} => {
// 先把所有字段提前,避免 match 中 part-move
let block_type = match &content_block {
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
},
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
};
events.push(StreamEvent::ContentBlockStart { index, block_type });
let builder = match content_block {
AnthropicContentBlockStart::Text { text } => {
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
}
AnthropicContentBlockStart::ToolUse { id, name } => {
crate::llm::types::response_v2::ContentBlockBuilder::ToolUse {
id,
name,
arguments: String::new(),
}
}
AnthropicContentBlockStart::Thinking { thinking } => {
crate::llm::types::response_v2::ContentBlockBuilder::Thinking {
buffer: thinking,
signature: None,
}
}
};
self.partial.blocks.entry(index).or_insert(builder);
self.partial.last_open_index = Some(index);
}
AnthropicSseEvent::ContentBlockDelta { index, delta } => match delta {
AnthropicDelta::TextDelta { text } => {
events.push(StreamEvent::TextDelta { text });
// ponytail: builder 通过末尾 `partial.apply_to(&ev)` 自动更新,
// 不需要此处手动 push —— 之前写错导致双重拼接。
self.partial.last_open_index = Some(index);
}
AnthropicDelta::InputJson { partial_json } => {
events.push(StreamEvent::ToolCallArgumentsDelta {
index,
arguments: partial_json,
});
}
AnthropicDelta::ThinkingDelta { thinking } => {
events.push(StreamEvent::ThinkingDelta { text: thinking });
}
AnthropicDelta::SignatureDelta { signature } => {
self.partial.set_thinking_signature(signature);
}
},
AnthropicSseEvent::ContentBlockStop { index } => {
// 不区分 text/refusal 与 tool_use,统一用 ContentBlockEnd
// 消费方在 ContentBlockStart 已知类型。
events.push(StreamEvent::ContentBlockEnd { index });
self.partial.block_completion.insert(index);
if self.partial.last_open_index == Some(index) {
self.partial.last_open_index = None;
}
}
AnthropicSseEvent::MessageDelta { delta, usage } => {
if let Some(reason) = delta.stop_reason.as_deref() {
self.partial.stop_reason = Some(match reason {
"end_turn" => StopReason::Stop,
"max_tokens" => StopReason::MaxTokens,
"tool_use" => StopReason::ToolUse,
"stop_sequence" => StopReason::StopSequence,
_ => StopReason::Other,
});
}
if let Some(u) = usage {
let partial_usage = PartialUsage {
prompt_tokens: u.input_tokens,
completion_tokens: Some(u.output_tokens),
total_tokens: None,
completion_tokens_details: None,
prompt_tokens_details: None,
};
events.push(StreamEvent::CostUpdate {
usage: partial_usage,
});
}
}
AnthropicSseEvent::MessageStop => {
// 收尾:partial.finalize() → MessageComplete
if self.saw_terminal {
return events;
}
self.saw_terminal = true;
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete {
full_response: full,
});
}
Err(e) => {
events.push(StreamEvent::Error {
message: e.to_string(),
});
}
}
}
AnthropicSseEvent::Ping | AnthropicSseEvent::Unknown => {
// 忽略
}
}
// 应用 side effects
for ev in &events {
self.partial.apply_to(ev);
}
events
}
}
fn _unused_marker() {}
impl Stream for AnthropicSseStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(data) = self.next_event_line() {
let mut events = self.handle_event_json(&data);
if !events.is_empty() {
return Poll::Ready(Some(Ok(events.remove(0))));
}
continue;
}
match Pin::new(&mut self.chunks).poll_next(cx) {
Poll::Ready(Some(Ok(bytes))) => {
if let Ok(s) = std::str::from_utf8(&bytes) {
self.buffer.push_str(s);
}
}
Poll::Ready(Some(Err(e))) => return Poll::Ready(Some(Err(e))),
Poll::Ready(None) => {
// 流关闭:如未收到 message_stop,主动 finalize
if !self.saw_terminal {
self.saw_terminal = true;
match self.partial.clone().finalize() {
Ok(full) => {
return Poll::Ready(Some(Ok(StreamEvent::MessageComplete {
full_response: full,
})));
}
Err(e) => {
return Poll::Ready(Some(Ok(StreamEvent::Error {
message: e.to_string(),
})));
}
}
}
return Poll::Ready(None);
}
Poll::Pending => return Poll::Pending,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::types::request_v2::MessageRequest;
use serde_json::json;
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn make_provider(base_url: String) -> AnthropicProvider {
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
AnthropicProvider::new(
base_url,
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
30,
)
}
#[tokio::test]
async fn anthropic_chat_basic_text_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_01test",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "Hello from Claude!"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "Hello from Claude!");
assert_eq!(response.stop_reason, StopReason::Stop);
assert_eq!(response.usage.prompt_tokens, 10);
}
#[tokio::test]
async fn anthropic_401_maps_to_authentication() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(401).set_body_string("invalid api key"))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
assert!(matches!(err, LlmError::Authentication(_)));
}
#[tokio::test]
async fn anthropic_529_overloaded_maps_to_rate_limit() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(529).set_body_string("overloaded"))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
assert!(matches!(err, LlmError::RateLimit { .. }));
}
#[tokio::test]
async fn anthropic_stream_emits_message_complete() {
let server = MockServer::start().await;
let sse = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_s1\",\"model\":\"claude-sonnet-4-20250514\"}}\n\n\
event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\
event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hi\"}}\n\n\
event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\" there\"}}\n\n\
event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n\
event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":3,\"input_tokens\":7}}\n\n\
event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("content-type", "text/event-stream")
.set_body_string(sse),
)
.mount(&server)
.await;
let provider = make_provider(server.uri());
let mut stream = provider
.chat_stream_inner(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hello")],
stream: true,
..Default::default()
})
.await
.unwrap();
use futures_util::StreamExt;
let mut collected: Vec<StreamEvent> = Vec::new();
while let Some(ev) = stream.next().await {
collected.push(ev.unwrap());
}
let complete = collected
.iter()
.find_map(|e| match e {
StreamEvent::MessageComplete { full_response } => Some(full_response.clone()),
_ => None,
})
.expect("expected MessageComplete");
assert_eq!(complete.text(), "Hi there");
assert_eq!(complete.stop_reason, StopReason::Stop);
assert_eq!(complete.usage.prompt_tokens, 7);
// 应当看到 MessageStart + ContentBlockStart + 2 个 TextDelta + ContentBlockEnd + CostUpdate + MessageComplete
let starts = collected
.iter()
.filter(|e| matches!(e, StreamEvent::MessageStart { .. }))
.count();
assert_eq!(starts, 1);
let cost = collected.iter().find_map(|e| match e {
StreamEvent::CostUpdate { usage } => Some(usage.clone()),
_ => None,
});
assert_eq!(cost.and_then(|u| u.completion_tokens), Some(3));
}
#[test]
fn anthropic_capabilities_reports_thinking_and_tool_use() {
let caps = AnthropicProvider::new(
"http://x".into(),
"k".into(),
"claude-sonnet-4-20250514".into(),
30,
)
.capabilities();
assert_eq!(caps.provider_name, "anthropic");
assert!(caps.features.thinking);
assert!(caps.features.tool_use);
assert_eq!(caps.features.max_context_window, 200_000);
}
#[tokio::test]
async fn anthropic_default_max_tokens_when_missing() {
// 不通过 wire —— 仅测试 convert_request 默认 max_tokens
let provider = make_provider("http://x".into());
let req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
};
let body = provider.build_request_body(req).unwrap();
assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS);
assert_eq!(body.model, "claude-sonnet-4-20250514");
}
// ===== Phase 11 Step 11.2 wiremock roundtrip 测试 =====
#[tokio::test]
async fn anthropic_401_structured_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!({
"type": "error",
"error": {
"type": "authentication_error",
"message": "Invalid API key provided: sk-ant-test"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Authentication(msg) => assert!(msg.contains("Invalid API key")),
other => panic!("expected Authentication, got {other:?}"),
}
}
#[tokio::test]
async fn anthropic_tool_use_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_tool",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [
{"type": "text", "text": "Let me check."},
{"type": "tool_use", "id": "toolu_abc", "name": "lookup", "input": {"q": "rust"}}
],
"stop_reason": "tool_use",
"usage": {"input_tokens": 8, "output_tokens": 12}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Look up rust")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.stop_reason, StopReason::ToolUse);
let tool_use = match &response.message {
Message::Assistant { content } => content.iter().find_map(|b| match b {
ContentBlock::ToolUse { id, name, .. } => Some((id.clone(), name.clone())),
_ => None,
}),
_ => None,
};
assert_eq!(tool_use, Some(("toolu_abc".into(), "lookup".into())));
}
#[tokio::test]
async fn anthropic_version_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.and(header("anthropic-version", "2023-06-01"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_v",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "OK"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "OK");
}
#[tokio::test]
async fn anthropic_529_overloaded_structured() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(529).set_body_json(json!({
"type": "error",
"error": {
"type": "overloaded_error",
"message": "Overloaded: Anthropic API is temporarily overloaded"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::RateLimit { retry_after } => assert!(retry_after.is_none()),
other => panic!("expected RateLimit, got {other:?}"),
}
}
// ===== custom_headers (Phase 8 Step 8.7) =====
fn mock_messages_body() -> serde_json::Value {
json!({
"id": "msg_test",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "OK"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})
}
fn make_provider_with_extra_headers(
base_url: String,
extra_headers: Vec<(String, String)>,
) -> AnthropicProvider {
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
AnthropicProvider::from_parts(
base_url,
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
client,
30,
extra_headers,
)
}
#[test]
fn anthropic_custom_headers_from_extra() {
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra("custom_headers", json!({"X-Custom": "v1", "X-Other": "v2"}));
let body = provider.build_request_body(req).unwrap();
assert_eq!(body.custom_headers.get("X-Custom").unwrap(), "v1");
assert_eq!(body.custom_headers.get("X-Other").unwrap(), "v2");
}
#[test]
fn anthropic_custom_headers_skipped_in_json_body() {
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
let body = provider.build_request_body(req).unwrap();
let value = serde_json::to_value(&body).unwrap();
assert!(
value.get("custom_headers").is_none(),
"custom_headers 不应进入 JSON body"
);
}
#[test]
fn anthropic_custom_headers_invalid_type_fallback() {
let provider = make_provider_with_extra_headers("http://x".into(), Vec::new());
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra("custom_headers", json!("not_an_object"));
let body = provider.build_request_body(req).unwrap();
assert!(body.custom_headers.is_empty());
}
#[tokio::test]
async fn anthropic_custom_headers_are_sent() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.and(header("X-Custom", "v1"))
.and(header("anthropic-version", "2023-06-01"))
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra("custom_headers", json!({"X-Custom": "v1"}));
let resp = provider.chat_blocking(req).await.unwrap();
assert_eq!(resp.text(), "OK");
}
#[tokio::test]
async fn anthropic_provider_level_headers_are_sent() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.and(header("X-Platform", "anthropic-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
.mount(&server)
.await;
let provider = make_provider_with_extra_headers(
server.uri(),
vec![("X-Platform".into(), "anthropic-test".into())],
);
let req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
let resp = provider.chat_blocking(req).await.unwrap();
assert_eq!(resp.text(), "OK");
}
#[tokio::test]
async fn anthropic_custom_headers_override_provider_headers() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
.mount(&server)
.await;
let provider = make_provider_with_extra_headers(
server.uri(),
vec![("X-Platform".into(), "provider-level".into())],
);
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra("custom_headers", json!({"X-Platform": "request-wins"}));
let resp = provider.chat_blocking(req).await.unwrap();
assert_eq!(resp.text(), "OK");
let received = server.received_requests().await.unwrap();
assert_eq!(received.len(), 1);
let platforms: Vec<&str> = received[0]
.headers
.get_all("X-Platform")
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
assert!(platforms.contains(&"request-wins"));
}
#[tokio::test]
async fn anthropic_custom_headers_can_override_auth_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(mock_messages_body()))
.mount(&server)
.await;
let provider = make_provider_with_extra_headers(server.uri(), Vec::new());
let mut req = MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
};
req.set_extra(
"custom_headers",
json!({"x-api-key": "from-custom-headers"}),
);
let resp = provider.chat_blocking(req).await.unwrap();
assert_eq!(resp.text(), "OK");
let received = server.received_requests().await.unwrap();
assert_eq!(received.len(), 1);
let keys: Vec<&str> = received[0]
.headers
.get_all("x-api-key")
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
assert!(
keys.contains(&"from-custom-headers"),
"custom_headers 应能覆盖 x-api-key 头,实际收到: {keys:?}"
);
}
#[test]
fn anthropic_extra_headers_from_constructor_unit() {
// ponytail: 用 RequestBuilder::build() 直检 headers,无需 wiremock。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = AnthropicProvider::from_parts(
"http://x".into(),
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
client,
30,
vec![("X-Platform".into(), "anthropic-test".into())],
);
let body = AnthropicRequestBody {
model: "claude-sonnet-4-20250514".into(),
max_tokens: 4096,
system: None,
messages: Vec::new(),
tools: None,
thinking: None,
stream: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Platform").unwrap(), "anthropic-test");
}
#[test]
fn anthropic_invalid_header_name_is_skipped() {
// ponytail: extra_headers 含非法 key 应被跳过,不应让 reqwest panic。
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("create http client");
let provider = AnthropicProvider::from_parts(
"http://x".into(),
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
client,
30,
Vec::new(),
)
.with_extra_headers(vec![
("X-Valid".into(), "v1".into()),
("bad\nname".into(), "v2".into()),
]);
let body = AnthropicRequestBody {
model: "claude-sonnet-4-20250514".into(),
max_tokens: 4096,
system: None,
messages: Vec::new(),
tools: None,
thinking: None,
stream: None,
custom_headers: HashMap::new(),
};
let req = provider
.build_request_builder(&body)
.unwrap()
.build()
.unwrap();
assert_eq!(req.headers().get("X-Valid").unwrap(), "v1");
assert!(
req.headers().get("bad\nname").is_none(),
"非法 header 名应被静默跳过"
);
}
}