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:
徐涛
2026-07-02 21:58:07 +08:00
parent 925c8f9729
commit 7c299f1cfd
18 changed files with 2605 additions and 441 deletions
+323 -13
View File
@@ -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;
}