feat(llm): 新增 IR 类型层 + 切换 LlmProvider trait 签名
引入跨 Provider 统一的新类型系统: - Message 扁平大枚举(System/User/UserImage/Assistant/ToolResult) - MessageRequest + MessageResponse + StreamEvent(高精度 IR) - PartialMessageResponse 含 apply_to/finalize 汇聚算法 - ProviderType enum 扩展 + ProviderCapabilities 切换 LlmProvider trait 签名为 chat(MessageRequest) → MessageResponse, chat_stream 返回 Stream<Item = Result<StreamEvent, LlmError>>,新增 capabilities()。 StreamEvent 命名冲突处理:旧变体迁移至 old_stream::LegacyStreamEvent, stream.rs 通过 pub use 重新导出新高精度事件。 OpenaiProvider 添加 Phase 0 临时桥接(MessageRequest ↔ OpenaiChatRequest 转换), 标注 ponytail: Phase 0 临时桥接 标记,Phase 1 重写时移除。 测试 147 passed / 0 failed(新增 31 个新类型测试)。
This commit is contained in:
+323
-13
@@ -6,22 +6,32 @@ use bytes::Bytes;
|
||||
use futures_core::stream::Stream;
|
||||
use futures_util::StreamExt;
|
||||
use reqwest::Client;
|
||||
use serde_json::Value;
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
use super::LlmProvider;
|
||||
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::stream::parse_chunk_stream;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::request::{OpenaiChatRequest, StreamOptions};
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
|
||||
use crate::llm::types::usage::Usage;
|
||||
use crate::llm::types::{
|
||||
ChatRequest, ChatResponse, OpenaiChatChunk, OpenaiChatResponse, StreamOptions,
|
||||
ChatResponse, ContentField, FinishReason, FunctionCall, OpenaiChatChunk, OpenaiChatMessage,
|
||||
OpenaiChatResponse, OpenaiContentPart, OpenaiTool, OpenaiToolCall, OpenaiToolDefinition,
|
||||
StopSequence, ToolChoice,
|
||||
};
|
||||
|
||||
pub struct OpenaiProvider {
|
||||
http_client: Client,
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
model: String,
|
||||
}
|
||||
|
||||
impl OpenaiProvider {
|
||||
pub fn new(base_url: String, api_key: String, _model: String) -> Self {
|
||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||
let http_client = Client::builder()
|
||||
.timeout(Duration::from_secs(120))
|
||||
.build()
|
||||
@@ -31,6 +41,7 @@ impl OpenaiProvider {
|
||||
http_client,
|
||||
base_url,
|
||||
api_key,
|
||||
model,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -50,11 +61,63 @@ impl OpenaiProvider {
|
||||
LlmError::Other(format!("请求失败: {}", e))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for OpenaiProvider {
|
||||
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError> {
|
||||
fn map_stop_reason(r: Option<FinishReason>) -> StopReason {
|
||||
r.map(StopReason::from).unwrap_or(StopReason::Stop)
|
||||
}
|
||||
|
||||
/// ponytail: Phase 0 临时桥接 —— 把 IR `MessageRequest` 转换为 OpenAI wire 格式。
|
||||
fn convert_to_chat_request(&self, request: MessageRequest) -> OpenaiChatRequest {
|
||||
let messages = request.messages.iter().map(message_to_chat_message).collect();
|
||||
|
||||
let tools = if request.tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(
|
||||
request
|
||||
.tools
|
||||
.into_iter()
|
||||
.map(|t| OpenaiTool::Function { function: t })
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
|
||||
OpenaiChatRequest {
|
||||
model: request.model,
|
||||
messages,
|
||||
max_tokens: request.max_tokens,
|
||||
temperature: request.temperature,
|
||||
top_p: request.top_p,
|
||||
stop: if request.stop_sequences.is_empty() {
|
||||
None
|
||||
} else if request.stop_sequences.len() == 1 {
|
||||
Some(StopSequence::Single(request.stop_sequences[0].clone()))
|
||||
} else {
|
||||
Some(StopSequence::Multiple(request.stop_sequences))
|
||||
},
|
||||
tools,
|
||||
tool_choice: Some(request.tool_choice),
|
||||
stream: Some(request.stream),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// ponytail: Phase 0 临时桥接 —— 把 OpenAI `ChatResponse` 转换为 IR `MessageResponse`。
|
||||
fn convert_to_message_response(&self, response: ChatResponse) -> MessageResponse {
|
||||
let mut extra = std::collections::HashMap::new();
|
||||
let stop_reason = Self::map_stop_reason(response.stop_reason);
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: self.model.clone(),
|
||||
message: chat_message_to_message(&response.message),
|
||||
usage: response.usage,
|
||||
stop_reason,
|
||||
extra,
|
||||
}
|
||||
}
|
||||
|
||||
/// ponytail: Phase 0 临时桥接 —— 调用 OpenAI `/chat/completions` 拿到旧 wire 格式响应。
|
||||
async fn chat_inner(&self, request: OpenaiChatRequest) -> Result<ChatResponse, LlmError> {
|
||||
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
|
||||
|
||||
info!(model = %request.model, max_tokens = request.max_tokens, temperature = request.temperature, "发送 LLM 请求");
|
||||
@@ -118,9 +181,10 @@ impl LlmProvider for OpenaiProvider {
|
||||
Ok(ChatResponse::from(chat_response))
|
||||
}
|
||||
|
||||
async fn chat_stream(
|
||||
/// ponytail: Phase 0 临时桥接 —— 流式调用 `/chat/completions` 返回 chunk 流。
|
||||
async fn chat_stream_inner(
|
||||
&self,
|
||||
mut request: ChatRequest,
|
||||
mut request: OpenaiChatRequest,
|
||||
) -> Result<
|
||||
Pin<Box<dyn Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
|
||||
LlmError,
|
||||
@@ -159,14 +223,251 @@ impl LlmProvider for OpenaiProvider {
|
||||
});
|
||||
}
|
||||
|
||||
let byte_stream = response.bytes_stream().map(|r| {
|
||||
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
|
||||
});
|
||||
let byte_stream =
|
||||
response.bytes_stream().map(|r| {
|
||||
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
|
||||
});
|
||||
|
||||
Ok(Box::pin(SseChunkStream::new(byte_stream)))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for OpenaiProvider {
|
||||
// ponytail: Phase 0 临时桥接,Phase 1 重写时移除。
|
||||
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
let chat_req = self.convert_to_chat_request(request);
|
||||
let chat_resp = self.chat_inner(chat_req).await?;
|
||||
Ok(self.convert_to_message_response(chat_resp))
|
||||
}
|
||||
|
||||
// ponytail: Phase 0 临时桥接,Phase 1 重写时移除。
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
let chat_req = self.convert_to_chat_request(request);
|
||||
let chunk_stream = self.chat_stream_inner(chat_req).await?;
|
||||
// ponytail: 复用旧 parse_chunk_stream 把 chunk 流映射为 IR StreamEvent 流。
|
||||
// Phase 1 中 OpenAI Provider 直接产出 IR 事件后整体删除此调用。
|
||||
Ok(parse_chunk_stream(chunk_stream))
|
||||
}
|
||||
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
ProviderCapabilities {
|
||||
provider_name: "openai",
|
||||
supported_models: None,
|
||||
features: ProviderFeatures::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ===== 转换函数 =====
|
||||
|
||||
/// IR `Message` → OpenAI `OpenaiChatMessage`。
|
||||
///
|
||||
/// ponytail: Phase 0 临时转换层;Phase 2 中 LlmCycle 切换为 `Vec<Message>` 后移除。
|
||||
fn message_to_chat_message(msg: &Message) -> OpenaiChatMessage {
|
||||
match msg {
|
||||
Message::System { content } => OpenaiChatMessage::System {
|
||||
content: blocks_to_content_field(content),
|
||||
name: None,
|
||||
},
|
||||
Message::User { content } => OpenaiChatMessage::User {
|
||||
content: blocks_to_content_field(content),
|
||||
name: None,
|
||||
},
|
||||
// ponytail: UserImage 当前无对应 OpenAI 形态,回退为空的 User 消息。
|
||||
// Phase 2 切到 Vec<Message> 后此函数整体删除。
|
||||
Message::UserImage { data: _, mime_type: _, detail: _ } => OpenaiChatMessage::User {
|
||||
content: ContentField::Array(vec![]),
|
||||
name: None,
|
||||
},
|
||||
Message::Assistant { content } => {
|
||||
let mut text_blocks = Vec::new();
|
||||
let mut tool_call_blocks = 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: FunctionCall {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input).unwrap_or_default(),
|
||||
},
|
||||
});
|
||||
}
|
||||
// ponytail: thinking/refusal/image/audio/file 在回传时被截断,
|
||||
// Phase 2 切到 Vec<Message> 整体修复。
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
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_field(content),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/// OpenAI `OpenaiChatMessage` → IR `Message`。
|
||||
fn chat_message_to_message(msg: &OpenaiChatMessage) -> Message {
|
||||
match msg {
|
||||
OpenaiChatMessage::Developer { content, .. }
|
||||
| OpenaiChatMessage::System { content, .. } => Message::System {
|
||||
content: content_field_to_blocks(content),
|
||||
},
|
||||
OpenaiChatMessage::User { content, .. } => Message::User {
|
||||
content: content_field_to_blocks(content),
|
||||
},
|
||||
OpenaiChatMessage::Assistant {
|
||||
content,
|
||||
tool_calls,
|
||||
..
|
||||
} => {
|
||||
let mut blocks = content_field_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_field_to_blocks(content),
|
||||
is_error: false,
|
||||
},
|
||||
// ponytail: OpenAI 已废弃 Function 消息变体;映射为带 name 假 tool_call_id 的 ToolResult。
|
||||
OpenaiChatMessage::Function { content, name } => Message::ToolResult {
|
||||
tool_call_id: name.clone(),
|
||||
content: content_field_to_blocks(content),
|
||||
is_error: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn blocks_to_content_field(blocks: &[ContentBlock]) -> ContentField {
|
||||
if blocks.len() == 1
|
||||
&& let ContentBlock::Text { text } = &blocks[0]
|
||||
{
|
||||
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: crate::llm::types::openai_message::ImageURL {
|
||||
url: if source.is_url {
|
||||
source.data.clone()
|
||||
} else {
|
||||
format!("data:{};base64,{}", source.mime_type, source.data)
|
||||
},
|
||||
detail: Some(match source.mime_type.as_str() {
|
||||
"image/jpeg" | "image/png" | "image/webp" => {
|
||||
crate::llm::types::shared::ImageDetail::Auto
|
||||
}
|
||||
_ => crate::llm::types::shared::ImageDetail::Auto,
|
||||
}),
|
||||
},
|
||||
detail: None,
|
||||
}),
|
||||
ContentBlock::Audio { source } => Some(OpenaiContentPart::InputAudio {
|
||||
input_audio: crate::llm::types::openai_message::InputAudio {
|
||||
data: source.data.clone(),
|
||||
format: source.format,
|
||||
},
|
||||
}),
|
||||
// ponytail: File/Thinking/Extension/ToolUse/ToolResult 在回传时被截断。
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
if parts.is_empty() {
|
||||
ContentField::Array(vec![])
|
||||
} else {
|
||||
ContentField::Array(parts)
|
||||
}
|
||||
}
|
||||
|
||||
fn content_field_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, .. } => {
|
||||
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
|
||||
.detail
|
||||
.map(|_| "image/url".to_string())
|
||||
.unwrap_or_else(|| "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,
|
||||
},
|
||||
}),
|
||||
OpenaiContentPart::File { .. } => None,
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
struct SseChunkStream<S> {
|
||||
inner: S,
|
||||
buffer: String,
|
||||
@@ -184,7 +485,10 @@ impl<S: Stream<Item = Result<Bytes, LlmError>> + Unpin> SseChunkStream<S> {
|
||||
impl<S: Stream<Item = Result<Bytes, LlmError>> + Unpin> Stream for SseChunkStream<S> {
|
||||
type Item = Result<OpenaiChatChunk, LlmError>;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> std::task::Poll<Option<Self::Item>> {
|
||||
fn poll_next(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Option<Self::Item>> {
|
||||
loop {
|
||||
if let Some(pos) = self.buffer.find("\n") {
|
||||
let line = self.buffer.drain(..pos + 1).collect::<String>();
|
||||
@@ -233,3 +537,9 @@ impl<S: Stream<Item = Result<Bytes, LlmError>> + Unpin> Stream for SseChunkStrea
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Suppress unused-import warnings for symbols kept for Phase 1 parity.
|
||||
#[allow(dead_code)]
|
||||
fn _phase0_reexports() {
|
||||
let _ = Usage::default;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user