Files
agcore/src/llm/mock.rs
T

308 lines
11 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 公开的 [`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);
}
}