feat(llm): 公开 MockProvider 并修复测试与 API 一致性

- 新增 src/llm/mock.rs:公开 MockProvider,实现 chat 与 chat_stream
  (chat_stream 从 MessageResponse 拆解为 StreamEvent 序列),含 5 个内联测试
- src/llm.rs 注册 pub mod mock,供示例与外部 crate 引用
- src/agent/session.rs 测试改用 crate::llm::mock::MockProvider,
  删除内嵌私有副本,统一公开路径
- src/prompt.rs re-export validate_messages,与 PromptComposer 对齐公开路径
This commit is contained in:
徐涛
2026-07-03 15:50:39 +08:00
parent c084c57e2c
commit 7b2d2db322
4 changed files with 309 additions and 45 deletions
+2 -44
View File
@@ -175,16 +175,12 @@ impl AgentSession {
mod tests {
use super::*;
use crate::agent::builder::AgentBuilder;
use crate::llm::error::LlmError;
use crate::llm::hooks::{Hook, HookContext, HookExecutor, HookResult};
use crate::llm::provider::{LlmProvider, ProviderCapabilities, ProviderFeatures};
use crate::llm::mock::MockProvider;
use crate::llm::types::message::ContentBlock;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
use crate::llm::types::response_v2::{MessageResponse, StopReason};
use crate::tools::ToolRegistry;
use async_trait::async_trait;
use futures_core::Stream;
use std::pin::Pin;
use std::sync::atomic::{AtomicU32, Ordering};
/// 计数 hook —— 每被调用一次 +1。
@@ -208,44 +204,6 @@ mod tests {
}
}
/// MockProvider:按调用顺序返回预设响应。
struct MockProvider {
responses: std::sync::Mutex<Vec<MessageResponse>>,
}
impl MockProvider {
fn new(responses: Vec<MessageResponse>) -> Self {
Self {
responses: std::sync::Mutex::new(responses),
}
}
}
#[async_trait]
impl LlmProvider for MockProvider {
async fn chat(&self, _request: MessageRequest) -> Result<MessageResponse, LlmError> {
let mut responses = self.responses.lock().unwrap();
if responses.is_empty() {
return Err(LlmError::Other("no more mock responses".into()));
}
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(),
}
}
}
struct StubAgent {
name: String,
prompt: Option<String>,
+1
View File
@@ -5,6 +5,7 @@ pub mod convert;
pub mod cycle;
pub mod error;
pub mod hooks;
pub mod mock;
pub mod provider;
pub mod stream;
pub mod types;
+305
View File
@@ -0,0 +1,305 @@
//! 公开的 [`MockProvider`] —— 在示例与测试中按顺序返回预设响应。
//!
//! 提供与真实 LLM Provider 一致的 `chat` / `chat_stream` 行为,
//! 但不发起任何网络请求,因此适合做离线示例、回归测试和 CI gate。
//!
//! # 用法
//!
//! ```no_run
//! use std::sync::Arc;
//! use agcore::llm::mock::MockProvider;
//! use agcore::llm::provider::LlmProvider;
//! use agcore::llm::types::message::{ContentBlock, Message};
//! use agcore::llm::types::response_v2::{MessageResponse, StopReason};
//! use agcore::llm::types::Usage;
//!
//! let response = MessageResponse {
//! id: "resp-1".into(),
//! model: "mock".into(),
//! message: Message::Assistant {
//! content: vec![ContentBlock::Text { text: "hi".into() }],
//! },
//! usage: Usage::from_input_output(2, 1),
//! stop_reason: StopReason::Stop,
//! extra: Default::default(),
//! };
//!
//! let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![response]));
//! ```
//!
//! # 流式行为
//!
//! `chat_stream` 把当前响应的 `ContentBlock` 序列拆解为标准流事件:
//!
//! - `Text` → `ContentBlockStart(Text)` + `TextDelta` + `ContentBlockEnd`
//! - `Thinking` → `ContentBlockStart(Thinking)` + `ThinkingDelta` + `ContentBlockEnd`
//! - `ToolUse` → `ContentBlockStart(ToolUse)` + `ToolCallArgumentsDelta` + `ToolCallEnd`
//! - 其他变体(Image/Audio/File/Extension)—— 不在流路径中模拟,跳过
use std::pin::Pin;
use std::sync::Mutex;
use async_stream::stream;
use futures_core::Stream;
use crate::llm::error::LlmError;
use crate::llm::provider::{LlmProvider, ProviderCapabilities, ProviderFeatures};
use crate::llm::types::message::{ContentBlock, ContentBlockType, Message};
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, PartialUsage, StreamEvent};
/// 按调用顺序返回预设响应的 [`LlmProvider`]。
///
/// 内部用 `Mutex<Vec<MessageResponse>>` 存储队列;`chat` / `chat_stream` 均弹出队首响应。
/// 当队列耗尽时返回 `LlmError::Other("MockProvider: 预设响应已用完")`。
pub struct MockProvider {
responses: Mutex<Vec<MessageResponse>>,
}
impl MockProvider {
/// 用预设响应列表创建。
pub fn new(responses: Vec<MessageResponse>) -> Self {
Self {
responses: Mutex::new(responses),
}
}
/// 创建空队列的 Provider(后续可用 `extend` 注入响应)。
pub fn empty() -> Self {
Self::new(Vec::new())
}
/// 追加更多响应到队列末尾。
pub fn extend(&self, responses: Vec<MessageResponse>) {
self.responses.lock().unwrap().extend(responses);
}
/// 队列中剩余响应数量。
pub fn remaining(&self) -> usize {
self.responses.lock().unwrap().len()
}
fn pop(&self) -> Result<MessageResponse, LlmError> {
let mut guard = self.responses.lock().unwrap();
if guard.is_empty() {
return Err(LlmError::Other("MockProvider: 预设响应已用完".into()));
}
Ok(guard.remove(0))
}
}
#[async_trait::async_trait]
impl LlmProvider for MockProvider {
async fn chat(&self, _request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.pop()
}
async fn chat_stream(
&self,
_request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
let response = self.pop()?;
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
let id = response.id.clone();
let model = response.model.clone();
let usage = response.usage;
let blocks: Vec<ContentBlock> = match &response.message {
Message::Assistant { content } => content.clone(),
_ => Vec::new(),
};
let stream = stream! {
yield Ok(StreamEvent::MessageStart {
id: id.clone(),
model: model.clone(),
});
for (i, block) in blocks.iter().enumerate() {
let index = i as u32;
match block {
ContentBlock::Text { text } => {
yield Ok(StreamEvent::ContentBlockStart {
index,
block_type: ContentBlockType::Text,
});
if !text.is_empty() {
yield Ok(StreamEvent::TextDelta { text: text.clone() });
}
yield Ok(StreamEvent::ContentBlockEnd { index });
}
ContentBlock::Thinking { text, .. } => {
yield Ok(StreamEvent::ContentBlockStart {
index,
block_type: ContentBlockType::Thinking,
});
if !text.is_empty() {
yield Ok(StreamEvent::ThinkingDelta { text: text.clone() });
}
yield Ok(StreamEvent::ContentBlockEnd { index });
}
ContentBlock::ToolUse { id: tu_id, name, input } => {
yield Ok(StreamEvent::ContentBlockStart {
index,
block_type: ContentBlockType::ToolUse {
id: tu_id.clone(),
name: name.clone(),
},
});
yield Ok(StreamEvent::ToolCallArgumentsDelta {
index,
arguments: serde_json::to_string(input)
.unwrap_or_else(|_| "{}".into()),
});
yield Ok(StreamEvent::ToolCallEnd { index });
}
// 其他变体(Image / Audio / File / ToolResult / Extension / Refusal
// 在 mock 流路径下跳过:示例仅需 text/thinking/tool 三种主流路径。
_ => {}
}
}
yield Ok(StreamEvent::CostUpdate {
usage: PartialUsage {
prompt_tokens: Some(usage.prompt_tokens),
completion_tokens: Some(usage.completion_tokens),
total_tokens: Some(usage.total_tokens),
..Default::default()
},
});
yield Ok(StreamEvent::MessageComplete { full_response: response });
};
Ok(Box::pin(stream))
}
fn capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities {
provider_name: "mock",
supported_models: None,
features: ProviderFeatures::default(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::types::{StopReason, Usage};
use futures_util::StreamExt;
use serde_json::json;
fn text_response(text: &str) -> MessageResponse {
MessageResponse {
id: "r".into(),
model: "mock".into(),
message: Message::Assistant {
content: vec![ContentBlock::Text { text: text.into() }],
},
usage: Usage::from_input_output(3, 7),
stop_reason: StopReason::Stop,
extra: Default::default(),
}
}
#[tokio::test]
async fn chat_returns_queued_response() {
let provider = MockProvider::new(vec![text_response("hello")]);
let resp = provider
.chat(MessageRequest::default())
.await
.unwrap();
assert_eq!(resp.text(), "hello");
assert_eq!(provider.remaining(), 0);
}
#[tokio::test]
async fn chat_errors_when_empty() {
let provider = MockProvider::empty();
let err = provider.chat(MessageRequest::default()).await.unwrap_err();
assert!(matches!(err, LlmError::Other(_)));
}
#[tokio::test]
async fn extend_appends_responses() {
let provider = MockProvider::new(vec![text_response("a")]);
provider.extend(vec![text_response("b"), text_response("c")]);
assert_eq!(provider.remaining(), 3);
}
#[tokio::test]
async fn chat_stream_emits_text_delta_sequence() {
let provider = MockProvider::new(vec![text_response("hi")]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut seen_start = false;
let mut seen_block_start = false;
let mut seen_delta = false;
let mut seen_block_end = false;
let mut seen_complete = false;
while let Some(event) = stream.next().await {
match event.unwrap() {
StreamEvent::MessageStart { .. } => seen_start = true,
StreamEvent::ContentBlockStart {
block_type: ContentBlockType::Text,
..
} => seen_block_start = true,
StreamEvent::TextDelta { text } => {
assert_eq!(text, "hi");
seen_delta = true;
}
StreamEvent::ContentBlockEnd { .. } => seen_block_end = true,
StreamEvent::MessageComplete { .. } => {
seen_complete = true;
break;
}
_ => {}
}
}
assert!(seen_start);
assert!(seen_block_start);
assert!(seen_delta);
assert!(seen_block_end);
assert!(seen_complete);
}
#[tokio::test]
async fn chat_stream_emits_tool_call_arguments() {
let response = MessageResponse {
id: "r".into(),
model: "mock".into(),
message: Message::Assistant {
content: vec![ContentBlock::ToolUse {
id: "call_1".into(),
name: "search".into(),
input: json!({"q": "rust"}),
}],
},
usage: Usage::from_input_output(2, 5),
stop_reason: StopReason::ToolUse,
extra: Default::default(),
};
let provider = MockProvider::new(vec![response]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut saw_tool_args = false;
let mut saw_tool_end = false;
while let Some(event) = stream.next().await {
match event.unwrap() {
StreamEvent::ToolCallArgumentsDelta { arguments, .. } => {
let v: serde_json::Value = serde_json::from_str(&arguments).unwrap();
assert_eq!(v, json!({"q": "rust"}));
saw_tool_args = true;
}
StreamEvent::ToolCallEnd { .. } => saw_tool_end = true,
StreamEvent::MessageComplete { .. } => break,
_ => {}
}
}
assert!(saw_tool_args);
assert!(saw_tool_end);
}
}
+1 -1
View File
@@ -4,4 +4,4 @@ pub mod composer;
pub use error::PromptError;
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
pub use composer::PromptComposer;
pub use composer::{validate_messages, PromptComposer};