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:
+2
-44
@@ -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>,
|
||||
|
||||
@@ -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
@@ -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
@@ -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};
|
||||
|
||||
Reference in New Issue
Block a user