//! 公开的 [`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 = 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>` 存储队列;`chat` / `chat_stream` 均弹出队首响应。 /// 当队列耗尽时返回 `LlmError::Other("MockProvider: 预设响应已用完")`。 pub struct MockProvider { responses: Mutex>, } impl MockProvider { /// 用预设响应列表创建。 pub fn new(responses: Vec) -> Self { Self { responses: Mutex::new(responses), } } /// 创建空队列的 Provider(后续可用 `extend` 注入响应)。 pub fn empty() -> Self { Self::new(Vec::new()) } /// 追加更多响应到队列末尾。 pub fn extend(&self, responses: Vec) { self.responses.lock().unwrap().extend(responses); } /// 队列中剩余响应数量。 pub fn remaining(&self) -> usize { self.responses.lock().unwrap().len() } fn pop(&self) -> Result { 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 { self.pop() } async fn chat_stream( &self, _request: MessageRequest, ) -> Result> + 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 = 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); } }