feat(llm): 重写 Provider 适配层,支持 OpenAI Chat / Anthropic / DeepSeek / Qwen
重写 OpenaiChatProvider:移除 Phase 0 临时桥接,实现真实流式状态机 (ContentBlockStart / *Delta / ContentBlockEnd / ToolCallEnd), MessageComplete.full_response 由 PartialMessageResponse::finalize 产出。 新增 AnthropicProvider:实现 Messages API + SSE 事件序列 (message_start → content_block_start → content_block_delta → content_block_stop → message_delta → message_stop), 529(overloaded)映射为 RateLimit;thinking signature 由 partial.set_thinking_signature 内部写入。 新增 DeepSeekProvider / QwenProvider:OpenAI-compatible 协议的 newtype 包装,共享 GenericOpenaiProvider 的 HTTP / SSE / 转换逻辑; Qwen 通过 extra_headers 注入 X-DashScope-SSE: enable。 新增 convert.rs 公共转换模块:从 Phase 0 cycle.rs / Phase 0 OpenaiProvider 桥接层提取 from_openai / to_openai / content_to_blocks / blocks_to_content,避免跨 Provider 重复逻辑。 新增 wiremock dev-dependency + 14 个集成测试: - OpenaiChatProvider:基础文本 / 401 / 500 / 流式 MessageComplete - AnthropicProvider:基础 / 401 / 529 / 流式 SSE 序列 / max_tokens 默认值 - DeepSeek / Qwen:基础文本响应 171 个测试全部通过(之前 157,新增 14)。
This commit is contained in:
@@ -23,3 +23,4 @@ time = { version = "0.3", features = ["serde"] }
|
||||
|
||||
[dev-dependencies]
|
||||
dotenvy = "0.15.7"
|
||||
wiremock = "0.6"
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
//! LLM 调用周期 —— 大模型基础调用周期控制。
|
||||
|
||||
pub mod compact;
|
||||
pub mod convert;
|
||||
pub mod cycle;
|
||||
pub mod error;
|
||||
pub mod hooks;
|
||||
|
||||
@@ -0,0 +1,455 @@
|
||||
//! 跨 Provider 类型转换 —— `Message` ↔ OpenAI `OpenaiChatMessage`。
|
||||
//!
|
||||
//! Phase 1 引入:避免转换逻辑在 `LlmCycle`、各 Provider 中重复。
|
||||
//! 这些函数期望作为**纯函数**被调用 —— 无内部状态,方便跨 Provider 复用。
|
||||
//!
|
||||
//! 转换范围:仅处理 `OpenaiChatMessage` ↔ `Message`、`ContentField` ↔ `Vec<ContentBlock>`。
|
||||
//! 流式 chunk 转换见各 Provider 内部(OpenAI Chat / Anthropic)。
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::openai_message::{
|
||||
ContentField, OpenaiChatMessage, OpenaiContentPart,
|
||||
};
|
||||
use crate::llm::types::OpenaiToolCall;
|
||||
|
||||
/// `OpenaiChatMessage` → IR `Message`。
|
||||
///
|
||||
/// 转换规则:
|
||||
/// - `Developer` / `System` → `Message::System`
|
||||
/// - `User` → `Message::User`(含图片/音频等多模态 part → `ContentBlock`)
|
||||
/// - `Assistant` → `Message::Assistant`(`tool_calls` 转为 `ContentBlock::ToolUse`)
|
||||
/// - `Tool` → `Message::ToolResult`(`is_error` 暂为 `false`,OpenAI 不携带此标记)
|
||||
/// - `Function`(已废弃)→ `Message::ToolResult`(`name` 作为 `tool_call_id` 兜底)
|
||||
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
|
||||
match msg {
|
||||
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => {
|
||||
Message::System {
|
||||
content: content_to_blocks(content),
|
||||
}
|
||||
}
|
||||
OpenaiChatMessage::User { content, .. } => Message::User {
|
||||
content: content_to_blocks(content),
|
||||
},
|
||||
OpenaiChatMessage::Assistant {
|
||||
content,
|
||||
tool_calls,
|
||||
..
|
||||
} => {
|
||||
let mut blocks = content_to_blocks(content);
|
||||
if let Some(calls) = tool_calls {
|
||||
for call in calls {
|
||||
if let OpenaiToolCall::Function { id, function } = call {
|
||||
let input: Value =
|
||||
serde_json::from_str(&function.arguments).unwrap_or(Value::Null);
|
||||
blocks.push(ContentBlock::ToolUse {
|
||||
id: id.clone(),
|
||||
name: function.name.clone(),
|
||||
input,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
Message::Assistant { content: blocks }
|
||||
}
|
||||
OpenaiChatMessage::Tool {
|
||||
content,
|
||||
tool_call_id,
|
||||
} => Message::ToolResult {
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
content: content_to_blocks(content),
|
||||
is_error: false,
|
||||
},
|
||||
// ponytail: `function` 是 OpenAI 旧版 `function_call` API 残留变体;
|
||||
// 当前主流为 `tool_calls`,此处仅保留兼容路径。
|
||||
OpenaiChatMessage::Function { content, name } => Message::ToolResult {
|
||||
tool_call_id: name.clone(),
|
||||
content: content_to_blocks(content),
|
||||
is_error: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/// IR `Message` → `OpenaiChatMessage`。
|
||||
///
|
||||
/// ponytail 简化:当前 `Assistant.content` 仅序列化 Text → `content: String`,
|
||||
/// `ToolUse` 拆出到 `tool_calls` 字段;thinking / image / audio / file / extension
|
||||
/// 视为非 OpenAI 原生,**回传时丢弃**。Phase 2 切换为 `Vec<Message>` 后该丢弃
|
||||
/// 自然消失(不再需要回传旧 wire 格式)。
|
||||
pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
|
||||
match msg {
|
||||
Message::System { content } => OpenaiChatMessage::System {
|
||||
content: blocks_to_content(content),
|
||||
name: None,
|
||||
},
|
||||
Message::User { content } => OpenaiChatMessage::User {
|
||||
content: blocks_to_content(content),
|
||||
name: None,
|
||||
},
|
||||
Message::UserImage { data, mime_type, detail } => {
|
||||
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
|
||||
let mime = mime_type.clone();
|
||||
let is_url = data.starts_with("http://") || data.starts_with("https://");
|
||||
let image_url = if is_url {
|
||||
crate::llm::types::openai_message::ImageURL {
|
||||
url: data.clone(),
|
||||
detail: Some(*detail),
|
||||
}
|
||||
} else {
|
||||
crate::llm::types::openai_message::ImageURL {
|
||||
url: format!("data:{mime};base64,{data}"),
|
||||
detail: Some(*detail),
|
||||
}
|
||||
};
|
||||
OpenaiChatMessage::User {
|
||||
content: ContentField::Array(vec![OpenaiContentPart::Image {
|
||||
image_url,
|
||||
detail: Some(*detail),
|
||||
}]),
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
Message::Assistant { content } => {
|
||||
let mut text_blocks: Vec<String> = Vec::new();
|
||||
let mut tool_call_blocks: Vec<OpenaiToolCall> = Vec::new();
|
||||
for block in content {
|
||||
match block {
|
||||
ContentBlock::Text { text } => text_blocks.push(text.clone()),
|
||||
ContentBlock::ToolUse { id, name, input } => {
|
||||
tool_call_blocks.push(OpenaiToolCall::Function {
|
||||
id: id.clone(),
|
||||
function: crate::llm::types::tool::FunctionCall {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input)
|
||||
.unwrap_or_else(|_| "null".to_string()),
|
||||
},
|
||||
});
|
||||
}
|
||||
// ponytail: thinking/refusal/image/audio/file/extension 在回传时被丢弃。
|
||||
// Phase 2 切换后此函数整体移除,自然修复。
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
let text: String = text_blocks.into_iter().collect();
|
||||
OpenaiChatMessage::Assistant {
|
||||
content: if text.is_empty() {
|
||||
ContentField::Array(vec![])
|
||||
} else {
|
||||
ContentField::String(text)
|
||||
},
|
||||
refusal: None,
|
||||
name: None,
|
||||
tool_calls: if tool_call_blocks.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(tool_call_blocks)
|
||||
},
|
||||
}
|
||||
}
|
||||
Message::ToolResult {
|
||||
tool_call_id,
|
||||
content,
|
||||
is_error: _,
|
||||
} => OpenaiChatMessage::Tool {
|
||||
content: blocks_to_content(content),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/// `ContentField` → `Vec<ContentBlock>`。
|
||||
///
|
||||
/// 仅产出 OpenAI 有 wire 对应的 `Text` / `Image` / `Audio` /
|
||||
/// `Refusal`(合并为 `Text`)block;`File` 被截断。
|
||||
pub fn content_to_blocks(field: &ContentField) -> Vec<ContentBlock> {
|
||||
match field {
|
||||
ContentField::String(s) => vec![ContentBlock::Text { text: s.clone() }],
|
||||
ContentField::Array(parts) => parts
|
||||
.iter()
|
||||
.filter_map(|p| match p {
|
||||
OpenaiContentPart::Text { text } => {
|
||||
Some(ContentBlock::Text { text: text.clone() })
|
||||
}
|
||||
OpenaiContentPart::Refusal { refusal } => {
|
||||
Some(ContentBlock::Text { text: refusal.clone() })
|
||||
}
|
||||
OpenaiContentPart::Image { image_url, .. } => {
|
||||
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
|
||||
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
|
||||
let url = &image_url.url;
|
||||
if let Some(rest) = url.strip_prefix("data:") {
|
||||
if let Some((mime, b64)) = rest.split_once(";base64,") {
|
||||
return Some(ContentBlock::Image {
|
||||
source: crate::llm::types::message::ImageSource {
|
||||
data: b64.to_string(),
|
||||
mime_type: mime.to_string(),
|
||||
is_url: false,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
Some(ContentBlock::Image {
|
||||
source: crate::llm::types::message::ImageSource {
|
||||
data: url.clone(),
|
||||
mime_type: "image/url".to_string(),
|
||||
is_url: true,
|
||||
},
|
||||
})
|
||||
}
|
||||
OpenaiContentPart::InputAudio { input_audio } => Some(ContentBlock::Audio {
|
||||
source: crate::llm::types::message::AudioSource {
|
||||
data: input_audio.data.clone(),
|
||||
format: input_audio.format,
|
||||
},
|
||||
}),
|
||||
// File 暂不映射(OpenAI File API 与 IR 不对齐)
|
||||
OpenaiContentPart::File { .. } => None,
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
/// `Vec<ContentBlock>` → `ContentField`。
|
||||
///
|
||||
/// 单文本块 → `ContentField::String`;多块或非文本主导 → `ContentField::Array`。
|
||||
pub fn blocks_to_content(blocks: &[ContentBlock]) -> ContentField {
|
||||
if let [ContentBlock::Text { text }] = blocks {
|
||||
return ContentField::String(text.clone());
|
||||
}
|
||||
let parts: Vec<OpenaiContentPart> = blocks
|
||||
.iter()
|
||||
.filter_map(|b| match b {
|
||||
ContentBlock::Text { text } => Some(OpenaiContentPart::Text { text: text.clone() }),
|
||||
ContentBlock::Image { source } => Some(OpenaiContentPart::Image {
|
||||
image_url: if source.is_url {
|
||||
crate::llm::types::openai_message::ImageURL {
|
||||
url: source.data.clone(),
|
||||
detail: None,
|
||||
}
|
||||
} else {
|
||||
crate::llm::types::openai_message::ImageURL {
|
||||
url: format!("data:{};base64,{}", source.mime_type, source.data),
|
||||
detail: None,
|
||||
}
|
||||
},
|
||||
detail: None,
|
||||
}),
|
||||
ContentBlock::Audio { source } => Some(OpenaiContentPart::InputAudio {
|
||||
input_audio: crate::llm::types::openai_message::InputAudio {
|
||||
data: source.data.clone(),
|
||||
format: source.format,
|
||||
},
|
||||
}),
|
||||
// ponytail: 其他 block 类型在 OpenAI wire 上无对应,回传时被截断。
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
if parts.is_empty() {
|
||||
ContentField::Array(vec![])
|
||||
} else {
|
||||
ContentField::Array(parts)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::llm::types::message::ImageSource;
|
||||
use crate::llm::types::shared::{AudioFormat, ImageDetail};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn from_openai_system_maps_to_message_system() {
|
||||
let m = OpenaiChatMessage::system_text("you are helpful");
|
||||
let ir = from_openai(&m);
|
||||
match ir {
|
||||
Message::System { content } => {
|
||||
assert_eq!(content.len(), 1);
|
||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful"));
|
||||
}
|
||||
_ => panic!("expected System variant"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_openai_assistant_tool_call_to_tool_use_block() {
|
||||
let m = OpenaiChatMessage::Assistant {
|
||||
content: ContentField::String("ok".into()),
|
||||
refusal: None,
|
||||
name: None,
|
||||
tool_calls: Some(vec![OpenaiToolCall::Function {
|
||||
id: "call_1".into(),
|
||||
function: crate::llm::types::tool::FunctionCall {
|
||||
name: "search".into(),
|
||||
arguments: r#"{"q":"rust"}"#.into(),
|
||||
},
|
||||
}]),
|
||||
};
|
||||
let ir = from_openai(&m);
|
||||
match ir {
|
||||
Message::Assistant { content } => {
|
||||
assert_eq!(content.len(), 2);
|
||||
match &content[1] {
|
||||
ContentBlock::ToolUse { id, name, input } => {
|
||||
assert_eq!(id, "call_1");
|
||||
assert_eq!(name, "search");
|
||||
assert_eq!(input, &json!({"q": "rust"}));
|
||||
}
|
||||
_ => panic!("expected ToolUse"),
|
||||
}
|
||||
}
|
||||
_ => panic!("expected Assistant"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_openai_tool_to_tool_result() {
|
||||
let m = OpenaiChatMessage::Tool {
|
||||
content: ContentField::String("ok".into()),
|
||||
tool_call_id: "call_1".into(),
|
||||
};
|
||||
let ir = from_openai(&m);
|
||||
match ir {
|
||||
Message::ToolResult {
|
||||
tool_call_id,
|
||||
content,
|
||||
is_error,
|
||||
} => {
|
||||
assert_eq!(tool_call_id, "call_1");
|
||||
assert!(!is_error);
|
||||
assert_eq!(content.len(), 1);
|
||||
}
|
||||
_ => panic!("expected ToolResult"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn to_openai_assistant_text_singular_string_content() {
|
||||
let m = Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: "hi".into() }],
|
||||
};
|
||||
let wire = to_openai(&m);
|
||||
match wire {
|
||||
OpenaiChatMessage::Assistant {
|
||||
content,
|
||||
tool_calls,
|
||||
..
|
||||
} => {
|
||||
assert!(matches!(content, ContentField::String(s) if s == "hi"));
|
||||
assert!(tool_calls.is_none());
|
||||
}
|
||||
_ => panic!("expected Assistant"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn to_openai_assistant_with_tool_use_produces_tool_calls() {
|
||||
let m = Message::Assistant {
|
||||
content: vec![
|
||||
ContentBlock::Text { text: "ok".into() },
|
||||
ContentBlock::ToolUse {
|
||||
id: "call_1".into(),
|
||||
name: "search".into(),
|
||||
input: json!({"q": "rust"}),
|
||||
},
|
||||
],
|
||||
};
|
||||
let wire = to_openai(&m);
|
||||
match wire {
|
||||
OpenaiChatMessage::Assistant {
|
||||
content,
|
||||
tool_calls,
|
||||
..
|
||||
} => {
|
||||
assert!(matches!(content, ContentField::String(_)));
|
||||
let calls = tool_calls.expect("tool_calls");
|
||||
assert_eq!(calls.len(), 1);
|
||||
if let OpenaiToolCall::Function { id, function } = &calls[0] {
|
||||
assert_eq!(id, "call_1");
|
||||
assert_eq!(function.name, "search");
|
||||
assert!(function.arguments.contains("rust"));
|
||||
} else {
|
||||
panic!("expected Function variant");
|
||||
}
|
||||
}
|
||||
_ => panic!("expected Assistant"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn to_openai_user_image_builds_data_uri() {
|
||||
let m = Message::UserImage {
|
||||
data: "BASE64DATA".into(),
|
||||
mime_type: "image/png".into(),
|
||||
detail: ImageDetail::High,
|
||||
};
|
||||
let wire = to_openai(&m);
|
||||
match wire {
|
||||
OpenaiChatMessage::User { content, .. } => match content {
|
||||
ContentField::Array(parts) => {
|
||||
assert_eq!(parts.len(), 1);
|
||||
match &parts[0] {
|
||||
OpenaiContentPart::Image { image_url, .. } => {
|
||||
assert_eq!(
|
||||
image_url.url,
|
||||
"data:image/png;base64,BASE64DATA"
|
||||
);
|
||||
}
|
||||
_ => panic!("expected Image part"),
|
||||
}
|
||||
}
|
||||
_ => panic!("expected Array content"),
|
||||
},
|
||||
_ => panic!("expected User"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn blocks_to_content_singular_text_uses_string() {
|
||||
let blocks = vec![ContentBlock::Text { text: "hi".into() }];
|
||||
let f = blocks_to_content(&blocks);
|
||||
assert!(matches!(f, ContentField::String(s) if s == "hi"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn blocks_to_content_multiple_blocks_uses_array() {
|
||||
let blocks = vec![
|
||||
ContentBlock::Text { text: "a".into() },
|
||||
ContentBlock::Text { text: "b".into() },
|
||||
];
|
||||
let f = blocks_to_content(&blocks);
|
||||
assert!(matches!(f, ContentField::Array(parts) if parts.len() == 2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn content_to_blocks_refusal_merges_to_text() {
|
||||
let field = ContentField::Array(vec![OpenaiContentPart::Refusal {
|
||||
refusal: "policy violation".into(),
|
||||
}]);
|
||||
let blocks = content_to_blocks(&field);
|
||||
assert_eq!(blocks.len(), 1);
|
||||
assert!(matches!(&blocks[0], ContentBlock::Text { text } if text == "policy violation"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn roundtrip_image_via_data_uri() {
|
||||
let field = ContentField::Array(vec![OpenaiContentPart::Image {
|
||||
image_url: crate::llm::types::openai_message::ImageURL {
|
||||
url: "data:image/png;base64,AAA".into(),
|
||||
detail: Some(ImageDetail::Auto),
|
||||
},
|
||||
detail: None,
|
||||
}]);
|
||||
let blocks = content_to_blocks(&field);
|
||||
assert_eq!(blocks.len(), 1);
|
||||
match &blocks[0] {
|
||||
ContentBlock::Image { source } => {
|
||||
assert_eq!(source.data, "AAA");
|
||||
assert_eq!(source.mime_type, "image/png");
|
||||
assert!(!source.is_url);
|
||||
}
|
||||
_ => panic!("expected Image"),
|
||||
}
|
||||
}
|
||||
}
|
||||
+21
-13
@@ -1,4 +1,6 @@
|
||||
pub mod anthropic;
|
||||
pub mod openai;
|
||||
pub mod openai_compat;
|
||||
pub mod registry;
|
||||
|
||||
use std::pin::Pin;
|
||||
@@ -57,23 +59,29 @@ pub fn create_provider(
|
||||
config: ProviderConfig,
|
||||
) -> Result<Box<dyn LlmProvider>, LlmError> {
|
||||
match provider_type {
|
||||
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiProvider::new(
|
||||
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::OpenaiResponse => Err(LlmError::Other(
|
||||
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
|
||||
)),
|
||||
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new(
|
||||
config.base_url,
|
||||
config.api_key,
|
||||
config.model,
|
||||
))),
|
||||
ProviderType::OpenaiResponse => {
|
||||
unimplemented!("OpenAI Response Provider 将在 Phase 1 引入")
|
||||
}
|
||||
ProviderType::Anthropic => {
|
||||
unimplemented!("Anthropic Provider 将在 Phase 1 引入")
|
||||
}
|
||||
ProviderType::DeepSeek => {
|
||||
unimplemented!("DeepSeek Provider 将在 Phase 1 引入")
|
||||
}
|
||||
ProviderType::Qwen => {
|
||||
unimplemented!("Qwen Provider 将在 Phase 1 引入")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+814
-362
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,234 @@
|
||||
//! DeepSeek / Qwen Provider —— OpenAI-compatible 协议的 newtype 包装。
|
||||
//!
|
||||
//! DeepSeek 与 Qwen 都是 OpenAI-compatible(共享 `/v1/chat/completions` 协议),
|
||||
//! 但二者 base_url 与 Qwen 需要额外请求头(`X-DashScope-SSE: enable`)。
|
||||
//! 通过 newtype 包装 `GenericOpenaiProvider` 提供:
|
||||
//! - 独立的 `capabilities().provider_name`
|
||||
//! - 未来可独立扩展(如 Qwen 的特殊错误映射、DeepSeek 的特殊响应解析)
|
||||
//!
|
||||
//! 类型别名方案(`pub type DeepSeekProvider = GenericOpenaiProvider`)被否决:类型别名
|
||||
//! 无法在编译期区分 DeepSeek vs OpenAI 调用,编译期安全检查失效。
|
||||
//! (参考 `docs/10b-phase1-provider-adaptation.md` §"OpenAI-compatible 复用策略")
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use futures_core::Stream;
|
||||
|
||||
use super::openai::GenericOpenaiProvider;
|
||||
use super::ProviderCapabilities;
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||
use crate::llm::provider::LlmProvider;
|
||||
|
||||
// =============================================================================
|
||||
// DeepSeek
|
||||
// =============================================================================
|
||||
|
||||
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl DeepSeekProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://api.deepseek.com".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
Self(GenericOpenaiProvider::new_with_name(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"deepseek",
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
impl DeepSeekProvider {
|
||||
/// 测试中(带 mock_client)使用的构造器。
|
||||
pub fn new_with_client(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
client: reqwest::Client,
|
||||
) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://api.deepseek.com".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
let mut inner = GenericOpenaiProvider::new_with_name(url, api_key, model, "deepseek");
|
||||
inner.http_client = client;
|
||||
Self(inner)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for DeepSeekProvider {
|
||||
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
self.0.chat(request).await
|
||||
}
|
||||
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
let mut caps = self.0.capabilities();
|
||||
caps.provider_name = "deepseek";
|
||||
caps
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Qwen
|
||||
// =============================================================================
|
||||
|
||||
pub struct QwenProvider(pub GenericOpenaiProvider);
|
||||
|
||||
impl QwenProvider {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
// Qwen 兼容模式流式需要 DashScope 特定的 SSE 启用头。
|
||||
let inner = GenericOpenaiProvider::new_with_name_and_headers(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"qwen",
|
||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||
);
|
||||
Self(inner)
|
||||
}
|
||||
|
||||
/// 测试构造器。
|
||||
pub fn new_with_client(
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
client: reqwest::Client,
|
||||
) -> Self {
|
||||
let url = if base_url.is_empty() {
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
||||
} else {
|
||||
base_url
|
||||
};
|
||||
let mut inner = GenericOpenaiProvider::new_with_name_and_headers(
|
||||
url,
|
||||
api_key,
|
||||
model,
|
||||
"qwen",
|
||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||
);
|
||||
inner.http_client = client;
|
||||
Self(inner)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for QwenProvider {
|
||||
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
self.0.chat(request).await
|
||||
}
|
||||
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
self.0.chat_stream(request).await
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
let mut caps = self.0.capabilities();
|
||||
caps.provider_name = "qwen";
|
||||
caps
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::message::Message as IrMessage;
|
||||
use serde_json::json;
|
||||
use wiremock::matchers::{method, path};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
|
||||
#[tokio::test]
|
||||
async fn deepseek_chat_basic_text_response() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"id": "ds-1",
|
||||
"object": "chat.completion",
|
||||
"created": 1718000000,
|
||||
"model": "deepseek-chat",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "DeepSeek hi"},
|
||||
"finish_reason": "stop"
|
||||
}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = DeepSeekProvider::new(
|
||||
server.uri(),
|
||||
"sk-test".into(),
|
||||
"deepseek-chat".into(),
|
||||
);
|
||||
let response = provider
|
||||
.chat(MessageRequest {
|
||||
model: "deepseek-chat".into(),
|
||||
messages: vec![IrMessage::user_text("hi")],
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.text(), "DeepSeek hi");
|
||||
assert_eq!(provider.capabilities().provider_name, "deepseek");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn qwen_chat_basic_text_response() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/chat/completions"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"id": "qw-1",
|
||||
"object": "chat.completion",
|
||||
"created": 1718000000,
|
||||
"model": "qwen-plus",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Qwen 你好"},
|
||||
"finish_reason": "stop"
|
||||
}],
|
||||
"usage": {"prompt_tokens": 6, "completion_tokens": 2, "total_tokens": 8}
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into());
|
||||
let response = provider
|
||||
.chat(MessageRequest {
|
||||
model: "qwen-plus".into(),
|
||||
messages: vec![IrMessage::user_text("hi")],
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.text(), "Qwen 你好");
|
||||
assert_eq!(provider.capabilities().provider_name, "qwen");
|
||||
}
|
||||
}
|
||||
@@ -36,9 +36,9 @@ pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage
|
||||
// `StopReason` 是独立类型,由 `Message` / `ContentBlock` / `StopReason` 直接路径访问,
|
||||
// 旧别名(指 `OpenaiChatMessage` / `OpenaiContentPart` / `FinishReason`)已移除,
|
||||
// 避免新类型阴影。Phase 2 完成后再统一收敛。
|
||||
|
||||
/// 旧 wire-format 请求类型别名(Phase 1 重写 Provider 后弃用)。
|
||||
pub type ChatRequest = OpenaiChatRequest;
|
||||
//
|
||||
// Phase 1 起移除 `ChatRequest` 别名 —— 新代码统一使用 `MessageRequest`(v2 IR)。
|
||||
// `ChatResponse` 结构体仍存在,作为 OpenAI `chat_inner()` 内部 wire-format 转换目标。
|
||||
/// 旧 wire-format 响应结构(保留用于 OpenAI 内部转换层)。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ChatResponse {
|
||||
|
||||
@@ -136,7 +136,7 @@ impl PartialUsage {
|
||||
}
|
||||
|
||||
/// 内容块构建器 —— 在流式累积阶段持有单 block 的原始拼接状态。
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ContentBlockBuilder {
|
||||
/// 文本块(按增量拼接)。
|
||||
@@ -217,6 +217,26 @@ pub struct PartialMessageResponse {
|
||||
pub is_complete: bool,
|
||||
}
|
||||
|
||||
// ponytail: Provider 可能在 `finalize()` 时希望保留 partial 快照(例如 chat_stream
|
||||
// 流关闭时需要把最终构造结果暴露为 `MessageComplete.full_response`,但同时 partial
|
||||
// 仍要为消费方后续 apply 提供访问)。当前所有内部字段派生 `Clone` —— 单独 derive 即可。
|
||||
impl Clone for PartialMessageResponse {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
id: self.id.clone(),
|
||||
model: self.model.clone(),
|
||||
blocks: self.blocks.clone(),
|
||||
block_completion: self.block_completion.clone(),
|
||||
last_open_index: self.last_open_index,
|
||||
usage: self.usage.clone(),
|
||||
stop_reason: self.stop_reason,
|
||||
thinking_signature: self.thinking_signature.clone(),
|
||||
is_errored: self.is_errored,
|
||||
is_complete: self.is_complete,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialMessageResponse {
|
||||
/// 创建一个新的空累积状态。
|
||||
pub fn new() -> Self {
|
||||
|
||||
Reference in New Issue
Block a user