Skip to main content

lc_agents/streaming/
tool_call_stream.rs

1//! StreamingFunctionCallingAgent - 流式输出 Agent
2
3use std::pin::Pin;
4use std::sync::Arc;
5
6use futures_util::{Stream, StreamExt};
7use tokio::sync::mpsc;
8use tokio_stream::wrappers::ReceiverStream;
9
10use lc_core::language_models::BaseChatModel;
11use lc_providers::ProviderError;
12use lc_schema::Message;
13
14use super::state::AgentStreamEvent;
15
16/// 流式 Function Calling Agent
17///
18/// 流式输出 LLM 文本(token),结束后发 FinalAnswer。
19/// 工具调用状态通过 `AgentStreamEvent::ToolCall` 暴露。
20/// 支持任何实现了 `BaseChatModel` 的 LLM Provider。
21pub struct StreamingFunctionCallingAgent {
22    llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
23}
24
25impl StreamingFunctionCallingAgent {
26    /// 创建新的流式 Function Calling Agent
27    ///
28    /// # 向后兼容
29    /// 旧代码 `StreamingFunctionCallingAgent::new(openai_chat)` 仍然可用。
30    pub fn new<L>(llm: L) -> Self
31    where
32        L: BaseChatModel + Send + Sync + 'static,
33        L::Error: Into<ProviderError>,
34    {
35        Self {
36            llm: lc_providers::wrap_chat_model(llm),
37        }
38    }
39
40    /// 从已包装的 `Arc<dyn BaseChatModel>` 创建 Agent
41    pub fn from_arc(llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
42        Self { llm }
43    }
44
45    /// 流式执行:返回事件流
46    pub async fn invoke_stream(
47        &self,
48        input: String,
49    ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send>> {
50        let (tx, rx) = mpsc::channel(32);
51        let llm = self.llm.clone();
52        let messages = vec![Message::human(input)];
53
54        tokio::spawn(async move {
55            let mut stream = match llm.stream_chat(messages, None).await {
56                Ok(s) => s,
57                Err(e) => {
58                    let _ = tx
59                        .send(AgentStreamEvent::Error {
60                            message: format!("Stream initialization failed: {}", e),
61                        })
62                        .await;
63                    return;
64                }
65            };
66
67            let mut full = String::new();
68            while let Some(chunk) = stream.next().await {
69                match chunk {
70                    Ok(token) => {
71                        full.push_str(&token);
72                        if tx
73                            .send(AgentStreamEvent::Text { content: token })
74                            .await
75                            .is_err()
76                        {
77                            break;
78                        }
79                    }
80                    Err(e) => {
81                        let _ = tx
82                            .send(AgentStreamEvent::Error {
83                                message: format!("Stream error: {}", e),
84                            })
85                            .await;
86                        break;
87                    }
88                }
89            }
90
91            let _ = tx
92                .send(AgentStreamEvent::FinalAnswer { content: full })
93                .await;
94        });
95
96        Box::pin(ReceiverStream::new(rx))
97    }
98}
99
100#[cfg(test)]
101mod tests {
102    use super::*;
103    use lc_providers::{OpenAIChat, OpenAIConfig};
104
105    #[test]
106    fn test_new() {
107        let llm = OpenAIChat::new(OpenAIConfig::default());
108        let _agent = StreamingFunctionCallingAgent::new(llm);
109    }
110}