308 lines
11 KiB
Rust
308 lines
11 KiB
Rust
//! 公开的 [`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);
|
||
}
|
||
}
|