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:
+283
-177
@@ -16,11 +16,16 @@ use crate::llm::compact::{should_compact, microcompact, CompactConfig, CompactSt
|
||||
use crate::llm::cycle::retry::should_retry;
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::hooks::{HookContext, HookExecutor};
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::provider::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::openai_message::{
|
||||
ContentField, OpenaiChatMessage, OpenaiContentPart,
|
||||
};
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use crate::llm::types::{
|
||||
ChatRequest, ChatResponse, FinishReason, OpenaiChatMessage, OpenaiTool, OpenaiToolCall,
|
||||
ToolChoice, ToolDefinition,
|
||||
FinishReason, FunctionCall, OpenaiToolCall, ToolChoice, ToolDefinition,
|
||||
};
|
||||
|
||||
/// LLM 调用周期配置。
|
||||
@@ -155,27 +160,18 @@ impl LlmCycle {
|
||||
&mut self,
|
||||
messages: Vec<OpenaiChatMessage>,
|
||||
tools: Vec<ToolDefinition>,
|
||||
) -> Result<ChatResponse, LlmError> {
|
||||
let openai_tools: Option<Vec<OpenaiTool>> = if tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(
|
||||
tools
|
||||
.iter()
|
||||
.map(|t| OpenaiTool::Function {
|
||||
function: t.clone(),
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
let ir_messages: Vec<Message> =
|
||||
messages.iter().map(chat_message_to_ir_message).collect();
|
||||
let ir_tools: Vec<ToolDefinition> = tools.clone();
|
||||
|
||||
let request = ChatRequest {
|
||||
let request = MessageRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages,
|
||||
messages: ir_messages,
|
||||
tools: ir_tools,
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: self.config.max_tokens,
|
||||
temperature: self.config.temperature,
|
||||
tools: openai_tools,
|
||||
tool_choice: Some(ToolChoice::Auto),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
@@ -198,11 +194,7 @@ impl LlmCycle {
|
||||
match self.provider.chat(request).await {
|
||||
Ok(response) => {
|
||||
if let Some(ref executor) = self.hook_executor {
|
||||
let post_request = ChatRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages: vec![],
|
||||
..Default::default()
|
||||
};
|
||||
let post_request = MessageRequest::default();
|
||||
let ctx = HookContext::new(crate::llm::hooks::HookEvent::PostRequest)
|
||||
.with_request(&post_request);
|
||||
executor
|
||||
@@ -229,7 +221,7 @@ impl LlmCycle {
|
||||
&mut self,
|
||||
prompt: String,
|
||||
tools: Vec<ToolDefinition>,
|
||||
) -> Result<ChatResponse, LlmError> {
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
self.messages.push(OpenaiChatMessage::user_text(prompt));
|
||||
|
||||
if let Some(ref config) = self.compact_config
|
||||
@@ -273,7 +265,9 @@ impl LlmCycle {
|
||||
.await;
|
||||
}
|
||||
|
||||
self.messages.push(response.message.clone());
|
||||
// ponytail: Phase 0 中内部消息存储仍为 Vec<OpenaiChatMessage>,
|
||||
// 把 IR Message 转回旧 wire 格式以便存储。Phase 2 切换后移除。
|
||||
self.messages.push(message_to_chat_message(&response.message));
|
||||
self.usage.add(&response.usage);
|
||||
|
||||
return Ok(response);
|
||||
@@ -346,59 +340,28 @@ impl LlmCycle {
|
||||
}
|
||||
}
|
||||
|
||||
let chunk_stream = self.provider.chat_stream(request).await?;
|
||||
let ir_event_stream = self.provider.chat_stream(request).await?;
|
||||
let hook_executor = self.hook_executor.clone();
|
||||
let post_request = self.build_request(&tools);
|
||||
|
||||
// ponytail: Phase 0 中 Provider 临时桥接层把 IR 事件再回填成本 phase 的
|
||||
// 信息冗余(OpenaiProvider 的 chat_stream 现在已输出 IR 事件)。当前实现
|
||||
// 直接转发 IR 事件,仅在 `MessageComplete` 到达时收尾。
|
||||
Ok(Box::pin(stream! {
|
||||
use futures_util::StreamExt;
|
||||
let mut chunk_stream = chunk_stream;
|
||||
let mut ir_event_stream = ir_event_stream;
|
||||
|
||||
while let Some(result) = chunk_stream.next().await {
|
||||
while let Some(result) = ir_event_stream.next().await {
|
||||
match result {
|
||||
Ok(chunk) => {
|
||||
let mut assistant_text = String::new();
|
||||
let mut tool_started: Option<(String, String, String)> = None;
|
||||
|
||||
for choice in &chunk.choices {
|
||||
let delta = &choice.delta;
|
||||
|
||||
if let Some(content) = &delta.content {
|
||||
assistant_text.push_str(content);
|
||||
}
|
||||
|
||||
if let Some(tool_calls) = &delta.tool_calls
|
||||
&& let Some(tc) = tool_calls.first()
|
||||
{
|
||||
let crate::llm::types::OpenaiToolCall::Function { id, function } = tc;
|
||||
tool_started = Some((id.clone(), function.name.clone(), function.arguments.clone()));
|
||||
}
|
||||
Ok(event) => {
|
||||
let is_terminal = matches!(event, StreamEvent::MessageComplete { .. });
|
||||
let is_error = matches!(event, StreamEvent::Error { .. });
|
||||
yield event;
|
||||
if is_terminal {
|
||||
break;
|
||||
}
|
||||
|
||||
if !assistant_text.is_empty() {
|
||||
yield StreamEvent::AssistantTextDelta { text: assistant_text };
|
||||
}
|
||||
|
||||
if let Some((tool_call_id, tool_name, arguments)) = tool_started {
|
||||
let args: serde_json::Value = serde_json::from_str(&arguments)
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
yield StreamEvent::ToolExecutionStarted {
|
||||
tool_name,
|
||||
input: args,
|
||||
tool_call_id,
|
||||
};
|
||||
}
|
||||
|
||||
for choice in &chunk.choices {
|
||||
if let Some(finish_reason) = &choice.finish_reason {
|
||||
yield StreamEvent::TurnComplete {
|
||||
reason: *finish_reason,
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(usage_info) = &chunk.usage {
|
||||
yield StreamEvent::CostUpdate { usage: *usage_info };
|
||||
if is_error {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -429,37 +392,30 @@ impl LlmCycle {
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_request(&self, tools: &[ToolDefinition]) -> ChatRequest {
|
||||
let mut messages = self.messages.clone();
|
||||
fn build_request(&self, tools: &[ToolDefinition]) -> MessageRequest {
|
||||
let messages = self.messages.clone();
|
||||
|
||||
let ir_messages: Vec<Message> = messages
|
||||
.iter()
|
||||
.map(chat_message_to_ir_message)
|
||||
.collect();
|
||||
|
||||
let mut ir_messages = ir_messages;
|
||||
if let Some(sys_prompt) = &self.system_prompt
|
||||
&& !messages
|
||||
&& !ir_messages
|
||||
.iter()
|
||||
.any(|m| matches!(m, OpenaiChatMessage::System { .. }))
|
||||
.any(|m| matches!(m, Message::System { .. }))
|
||||
{
|
||||
messages.insert(0, OpenaiChatMessage::system_text(sys_prompt));
|
||||
ir_messages.insert(0, Message::system(sys_prompt.as_str()));
|
||||
}
|
||||
|
||||
let openai_tools: Option<Vec<OpenaiTool>> = if tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(
|
||||
tools
|
||||
.iter()
|
||||
.map(|t| OpenaiTool::Function {
|
||||
function: t.clone(),
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
};
|
||||
|
||||
ChatRequest {
|
||||
MessageRequest {
|
||||
model: self.config.model.clone(),
|
||||
messages,
|
||||
messages: ir_messages,
|
||||
tools: tools.to_vec(),
|
||||
tool_choice: ToolChoice::Auto,
|
||||
max_tokens: self.config.max_tokens,
|
||||
temperature: self.config.temperature,
|
||||
tools: openai_tools,
|
||||
tool_choice: Some(ToolChoice::Auto),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
@@ -470,7 +426,7 @@ impl LlmCycle {
|
||||
async fn submit_request(
|
||||
&mut self,
|
||||
tools: &[ToolDefinition],
|
||||
) -> Result<ChatResponse, LlmError> {
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
let mut attempts = 0;
|
||||
|
||||
loop {
|
||||
@@ -548,7 +504,7 @@ impl LlmCycle {
|
||||
&mut self,
|
||||
prompt: String,
|
||||
registry: &crate::tools::ToolRegistry,
|
||||
) -> Result<ChatResponse, LlmError> {
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
let tools = registry.definitions();
|
||||
let max_turns = self.config.max_tool_turns.unwrap_or(10);
|
||||
let tool_timeout = self.config.tool_timeout_secs;
|
||||
@@ -570,24 +526,25 @@ impl LlmCycle {
|
||||
let response = self.submit_request(&tools).await?;
|
||||
|
||||
// 判断是否需要执行工具
|
||||
let should_execute = matches!(response.stop_reason, Some(FinishReason::ToolCalls))
|
||||
&& has_tool_calls_in_message(&response.message);
|
||||
let should_execute = matches!(response.stop_reason, StopReason::ToolUse)
|
||||
&& has_tool_calls_in_response(&response);
|
||||
|
||||
// 将 Assistant 响应(含 tool_calls 或最终文本)追加到消息历史
|
||||
self.messages.push(response.message.clone());
|
||||
// ponytail: Phase 0 中内部消息存储仍为 Vec<OpenaiChatMessage>,
|
||||
// 把 IR Message 转回旧 wire 格式以便存储。Phase 2 切换后移除。
|
||||
self.messages.push(message_to_chat_message(&response.message));
|
||||
|
||||
if !should_execute {
|
||||
return Ok(response);
|
||||
}
|
||||
|
||||
// 解析 tool_calls 并执行
|
||||
let tool_calls = extract_tool_calls_from_message(&response.message);
|
||||
let tool_calls = extract_tool_calls_from_response(&response);
|
||||
let calls: Vec<(String, serde_json::Value)> = tool_calls
|
||||
.into_iter()
|
||||
.map(|(_id, name, args)| {
|
||||
let args: serde_json::Value =
|
||||
let value: serde_json::Value =
|
||||
serde_json::from_str(&args).unwrap_or(serde_json::Value::Null);
|
||||
(name, args)
|
||||
(name, value)
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -642,37 +599,185 @@ impl LlmCycle {
|
||||
}
|
||||
|
||||
/// 判断 Assistant 消息是否包含 tool_calls。
|
||||
fn has_tool_calls_in_message(msg: &OpenaiChatMessage) -> bool {
|
||||
matches!(
|
||||
msg,
|
||||
OpenaiChatMessage::Assistant {
|
||||
tool_calls: Some(calls),
|
||||
..
|
||||
} if !calls.is_empty()
|
||||
)
|
||||
fn has_tool_calls_in_response(response: &MessageResponse) -> bool {
|
||||
match &response.message {
|
||||
Message::Assistant { content } => content
|
||||
.iter()
|
||||
.any(|b| matches!(b, ContentBlock::ToolUse { .. })),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// 提取 Assistant 消息中的 tool_calls。
|
||||
///
|
||||
/// 返回 `(tool_call_id, tool_name, arguments_json_string)` 列表。
|
||||
fn extract_tool_calls_from_message(
|
||||
msg: &OpenaiChatMessage,
|
||||
///
|
||||
/// ponytail: 当前 Phase 0 实现,`arguments_json_string` 内含 JSON 序列化的 input。
|
||||
/// 消费方在调用 `registry.invoke_all()` 时反序列化一次。该小段冗余序列化
|
||||
/// 在 Phase 2 切换为 `Vec<Message>` 后可整体消除。
|
||||
fn extract_tool_calls_from_response(
|
||||
response: &MessageResponse,
|
||||
) -> Vec<(String, String, String)> {
|
||||
if let OpenaiChatMessage::Assistant {
|
||||
tool_calls: Some(calls),
|
||||
..
|
||||
} = msg
|
||||
{
|
||||
calls
|
||||
.iter()
|
||||
.map(|c| match c {
|
||||
OpenaiToolCall::Function { id, function } => {
|
||||
(id.clone(), function.name.clone(), function.arguments.clone())
|
||||
let mut out = Vec::new();
|
||||
if let Message::Assistant { content } = &response.message {
|
||||
for block in content {
|
||||
if let ContentBlock::ToolUse { id, name, input } = block {
|
||||
let args = serde_json::to_string(input).unwrap_or_else(|_| "null".to_string());
|
||||
out.push((id.clone(), name.clone(), args));
|
||||
}
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// 把 `OpenaiChatMessage` 转换为 IR `Message`。
|
||||
///
|
||||
/// ponytail: Phase 0 转换层 —— build_request 中需要把存储用的
|
||||
/// `Vec<OpenaiChatMessage>` 转为 `MessageRequest.messages`。Phase 2 中
|
||||
/// `self.messages` 直接改为 `Vec<Message>` 时此函数整体删除。
|
||||
fn chat_message_to_ir_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: serde_json::Value = serde_json::from_str(&function.arguments)
|
||||
.unwrap_or(serde_json::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,
|
||||
},
|
||||
OpenaiChatMessage::Function { content, name } => Message::ToolResult {
|
||||
tool_call_id: name.clone(),
|
||||
content: content_field_to_blocks(content),
|
||||
is_error: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/// 把 IR `Message` 转换为 `OpenaiChatMessage`(用于 `self.messages` 持久化)。
|
||||
///
|
||||
/// ponytail: Phase 0 转换层 —— submit 后 `response.message: Message` 需要
|
||||
/// 转回 `OpenaiChatMessage` 才能追加到旧 store。Phase 2 中移除。
|
||||
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 切换后此函数整体移除。
|
||||
Message::UserImage { .. } => OpenaiChatMessage::User {
|
||||
content: ContentField::Array(vec![]),
|
||||
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: FunctionCall {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input)
|
||||
.unwrap_or_else(|_| "null".to_string()),
|
||||
},
|
||||
});
|
||||
}
|
||||
// ponytail: thinking/refusal/image/audio/file 在 Phase 0 持久化时被丢弃。
|
||||
// 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_field(content),
|
||||
tool_call_id: tool_call_id.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn blocks_to_content_field(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() }),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
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() })
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect()
|
||||
} else {
|
||||
Vec::new()
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -691,19 +796,20 @@ fn truncate_tool_result(s: &str, max_bytes: usize) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::llm::types::{ContentField, OpenaiContentPart};
|
||||
use crate::tools::{BaseTool, ToolRegistry};
|
||||
use async_trait::async_trait;
|
||||
use futures_core::Stream;
|
||||
use serde_json::{json, Value};
|
||||
use std::pin::Pin;
|
||||
|
||||
/// 模拟 Provider —— 预定义响应序列,按调用顺序返回。
|
||||
struct MockProvider {
|
||||
responses: std::sync::Mutex<Vec<ChatResponse>>,
|
||||
responses: std::sync::Mutex<Vec<MessageResponse>>,
|
||||
call_count: std::sync::Mutex<u32>,
|
||||
}
|
||||
|
||||
impl MockProvider {
|
||||
fn new(responses: Vec<ChatResponse>) -> Self {
|
||||
fn new(responses: Vec<MessageResponse>) -> Self {
|
||||
Self {
|
||||
responses: std::sync::Mutex::new(responses),
|
||||
call_count: std::sync::Mutex::new(0),
|
||||
@@ -713,7 +819,7 @@ mod tests {
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for MockProvider {
|
||||
async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
|
||||
async fn chat(&self, _request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||
let mut count = self.call_count.lock().unwrap();
|
||||
*count += 1;
|
||||
let mut responses = self.responses.lock().unwrap();
|
||||
@@ -722,45 +828,61 @@ mod tests {
|
||||
}
|
||||
Ok(responses.remove(0))
|
||||
}
|
||||
async fn chat_stream(
|
||||
&self,
|
||||
_request: MessageRequest,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
||||
{
|
||||
unimplemented!()
|
||||
}
|
||||
fn capabilities(&self) -> ProviderCapabilities {
|
||||
ProviderCapabilities {
|
||||
provider_name: "mock",
|
||||
supported_models: None,
|
||||
features: ProviderFeatures::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn empty_usage() -> crate::llm::types::Usage {
|
||||
crate::llm::types::Usage::default()
|
||||
}
|
||||
|
||||
fn assistant_text_response(text: &str) -> ChatResponse {
|
||||
ChatResponse {
|
||||
message: OpenaiChatMessage::assistant_text(text),
|
||||
fn assistant_text_response(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text {
|
||||
text: text.into(),
|
||||
}],
|
||||
},
|
||||
usage: empty_usage(),
|
||||
stop_reason: Some(FinishReason::Stop),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: std::collections::HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_tool_call_response(
|
||||
calls: Vec<(&str, &str, &str)>,
|
||||
) -> ChatResponse {
|
||||
use crate::llm::types::{OpenaiToolCall, FunctionCall};
|
||||
let tool_calls: Vec<OpenaiToolCall> = calls
|
||||
fn assistant_tool_call_response(calls: Vec<(&str, &str, &str)>) -> MessageResponse {
|
||||
let tool_blocks: Vec<ContentBlock> = calls
|
||||
.into_iter()
|
||||
.map(|(id, name, args)| OpenaiToolCall::Function {
|
||||
id: id.to_string(),
|
||||
function: FunctionCall {
|
||||
.map(|(id, name, args)| {
|
||||
let input: serde_json::Value =
|
||||
serde_json::from_str(args).unwrap_or(serde_json::Value::Null);
|
||||
ContentBlock::ToolUse {
|
||||
id: id.to_string(),
|
||||
name: name.to_string(),
|
||||
arguments: args.to_string(),
|
||||
},
|
||||
input,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
ChatResponse {
|
||||
message: OpenaiChatMessage::Assistant {
|
||||
content: ContentField::Array(vec![OpenaiContentPart::Text {
|
||||
text: String::new(),
|
||||
}]),
|
||||
refusal: None,
|
||||
name: None,
|
||||
tool_calls: Some(tool_calls),
|
||||
},
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant { content: tool_blocks },
|
||||
usage: empty_usage(),
|
||||
stop_reason: Some(FinishReason::ToolCalls),
|
||||
stop_reason: StopReason::ToolUse,
|
||||
extra: std::collections::HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -790,7 +912,6 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_submit_with_tools_single_turn() {
|
||||
// 第一轮:返回 tool_call;第二轮:返回最终文本
|
||||
let responses = vec![
|
||||
assistant_tool_call_response(vec![("call_1", "add", r#"{"a":1,"b":2}"#)]),
|
||||
assistant_text_response("答案是 3"),
|
||||
@@ -805,13 +926,11 @@ mod tests {
|
||||
.submit_with_tools("1+2=?".to_string(), ®istry)
|
||||
.await
|
||||
.unwrap();
|
||||
// 验证最终响应是文本响应
|
||||
assert!(matches!(
|
||||
response.message,
|
||||
OpenaiChatMessage::Assistant { .. }
|
||||
));
|
||||
// 最终响应是 Assistant 消息
|
||||
assert!(matches!(response.message, Message::Assistant { .. }));
|
||||
|
||||
// 验证消息历史:user, assistant(tool_calls), tool, assistant(text)
|
||||
// 验证消息历史(Vec<OpenaiChatMessage>,Phase 2 才切到 Vec<Message>):
|
||||
// user, assistant(tool_calls → 转回后含 tool_use), tool, assistant(text)
|
||||
let messages = cycle.messages();
|
||||
assert_eq!(messages.len(), 4);
|
||||
assert!(matches!(messages[0], OpenaiChatMessage::User { .. }));
|
||||
@@ -822,7 +941,6 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_submit_with_tools_multi_turn() {
|
||||
// 3 轮 tool 调用后给出最终答案
|
||||
let responses = vec![
|
||||
assistant_tool_call_response(vec![("call_1", "add", r#"{"a":1,"b":2}"#)]),
|
||||
assistant_tool_call_response(vec![("call_2", "add", r#"{"a":3,"b":4}"#)]),
|
||||
@@ -839,10 +957,7 @@ mod tests {
|
||||
.submit_with_tools("计算总和".to_string(), ®istry)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
response.message,
|
||||
OpenaiChatMessage::Assistant { .. }
|
||||
));
|
||||
assert!(matches!(response.message, Message::Assistant { .. }));
|
||||
|
||||
// user + 3*(assistant + tool) + final assistant = 8
|
||||
let messages = cycle.messages();
|
||||
@@ -851,10 +966,8 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_submit_with_tools_max_turns_exceeded() {
|
||||
// 配置 max_tool_turns = 2
|
||||
let mut config = CycleConfig::default();
|
||||
config.max_tool_turns = Some(2);
|
||||
// 4 轮 tool 调用 + 终止
|
||||
let responses = vec![
|
||||
assistant_tool_call_response(vec![("c1", "add", r#"{"a":1,"b":1}"#)]),
|
||||
assistant_tool_call_response(vec![("c2", "add", r#"{"a":1,"b":1}"#)]),
|
||||
@@ -867,15 +980,12 @@ mod tests {
|
||||
let mut registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(AddTool)).unwrap();
|
||||
|
||||
let result = cycle
|
||||
.submit_with_tools("test".to_string(), ®istry)
|
||||
.await;
|
||||
let result = cycle.submit_with_tools("test".to_string(), ®istry).await;
|
||||
assert!(matches!(result, Err(LlmError::Other(msg)) if msg.contains("达到最大工具循环轮次")));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_submit_with_tools_no_tool_call_response() {
|
||||
// LLM 直接给出最终响应(不调用工具)
|
||||
let responses = vec![assistant_text_response("直接回答")];
|
||||
let provider = Box::new(MockProvider::new(responses));
|
||||
let mut cycle = LlmCycle::new(provider, CycleConfig::default());
|
||||
@@ -887,10 +997,7 @@ mod tests {
|
||||
.submit_with_tools("直接回答".to_string(), ®istry)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
response.message,
|
||||
OpenaiChatMessage::Assistant { .. }
|
||||
));
|
||||
assert!(matches!(response.message, Message::Assistant { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -911,7 +1018,6 @@ mod tests {
|
||||
fn test_truncate_tool_result_chinese_chars() {
|
||||
let s = "中".repeat(100);
|
||||
let truncated = truncate_tool_result(&s, 50);
|
||||
// 不会在字符中间截断
|
||||
assert!(truncated.starts_with("中"));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user