- 修复测试编译回归:补全 session.rs/cycle.rs 测试模块导入; convert.rs 2 处 irrefutable if let 改为 let - composer.rs 迁移至 IR:OpenaiChatMessage → Message, ContentField/OpenaiContentPart → ContentBlock;删除 set_message_name 和 build_request;developer 消息映射为 Message::System - knowledge.rs 锁修复:std::sync::Mutex → tokio::sync::Mutex; search() 优化锁粒度(锁内仅 clone IDs,避免锁内异步 IO) - 标记 ChatResponse / ToolDefinition 为废弃(#[deprecated(since = "0.1.0")]), 内部使用点加 #[allow(deprecated)] 抑制警告 - clippy 清零:合并冗余 if、手动 strip_prefix 改 strip_prefix、 多处 dead_code 抑制、测试代码清理
1004 lines
35 KiB
Rust
1004 lines
35 KiB
Rust
//! 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::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::header::{HeaderMap, HeaderValue};
|
||
use reqwest::Client;
|
||
use serde::{Deserialize, Serialize};
|
||
use serde_json::{json, Value};
|
||
use tracing::{debug, error, info, warn};
|
||
|
||
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||
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;
|
||
|
||
/// 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,
|
||
}
|
||
|
||
impl AnthropicProvider {
|
||
pub fn new(base_url: String, api_key: String, model: String) -> 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(120))
|
||
.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,
|
||
}
|
||
}
|
||
|
||
pub fn with_client(mut self, client: Client) -> Self {
|
||
self.http_client = client;
|
||
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; // 当前忽略 ImageDetail,Anthropic 不支持
|
||
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);
|
||
|
||
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 },
|
||
})
|
||
}
|
||
|
||
async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||
let body = self.build_request_body(request)?;
|
||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||
|
||
info!(model = %body.model, "Anthropic: 发送非流式请求");
|
||
|
||
let response = self
|
||
.http_client
|
||
.post(&url)
|
||
.json(&body)
|
||
.send()
|
||
.await
|
||
.map_err(Self::map_reqwest_error)?;
|
||
|
||
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);
|
||
|
||
let url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
|
||
|
||
info!(model = %body.model, "Anthropic: 发送流式请求");
|
||
|
||
let response = self
|
||
.http_client
|
||
.post(&url)
|
||
.json(&body)
|
||
.send()
|
||
.await
|
||
.map_err(Self::map_reqwest_error)?;
|
||
|
||
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(e: reqwest::Error) -> LlmError {
|
||
if e.is_timeout() {
|
||
LlmError::Timeout {
|
||
duration: Duration::from_secs(120),
|
||
}
|
||
} else if e.is_connect() {
|
||
LlmError::Other(format!("连接失败: {e}"))
|
||
} else {
|
||
LlmError::Other(format!("请求失败: {e}"))
|
||
}
|
||
}
|
||
|
||
/// Anthropic 错误映射:529(overloaded)→ RateLimit;401/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>,
|
||
}
|
||
|
||
#[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 wire(Anthropic 自动产生)。
|
||
// 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::{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())
|
||
}
|
||
|
||
#[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(),
|
||
)
|
||
.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");
|
||
}
|
||
}
|