Skip to main content

deepseek_sdk/chat/
stream.rs

1//! Chat streaming and request implementation for `/chat/completions`.
2use crate::DeepSeekRequest;
3use crate::error::DeepSeekError;
4use crate::{api_post, api_request_stream, consume_sse, spawn_blocking_stream};
5
6use super::{Chat, ChatStream, request::*};
7use reqwest::Method;
8use tokio::sync::mpsc;
9/// Stream item produced by chat streaming.
10pub type ChatStreamItem = Result<ChatStream, DeepSeekError>;
11
12/// Blocking iterator over streaming chat chunks.
13pub struct ChatStreamBlocking {
14    pub rx: std::sync::mpsc::Receiver<ChatStreamItem>,
15}
16
17impl Iterator for ChatStreamBlocking {
18    type Item = ChatStreamItem;
19
20    fn next(&mut self) -> Option<Self::Item> {
21        self.rx.recv().ok()
22    }
23}
24
25impl DeepSeekRequest for ChatRequest {
26    type Response = Chat;
27    type StreamItem = ChatStreamItem;
28    type BlockingStream = ChatStreamBlocking;
29
30    async fn send(self) -> Result<Chat, DeepSeekError> {
31        let client = self.client.clone();
32        api_post("/chat/completions", &self, client).await
33    }
34
35    async fn stream(self) -> Result<mpsc::Receiver<ChatStreamItem>, DeepSeekError> {
36        let mut request = self;
37        request.stream = Some(true);
38
39        let client = request.client.clone();
40        let event_source = api_request_stream(
41            Method::POST,
42            "/chat/completions",
43            |builder| builder.json(&request),
44            client,
45        )
46        .await?;
47
48        Ok(consume_sse(event_source, |data| {
49            serde_json::from_str::<ChatStream>(&data)
50                .map(Some)
51                .map_err(|err| DeepSeekError::decode(err.to_string(), data))
52        }))
53    }
54
55    fn stream_blocking(self) -> Result<ChatStreamBlocking, DeepSeekError> {
56        let rx = spawn_blocking_stream(self.stream())?;
57        Ok(ChatStreamBlocking { rx })
58    }
59}
60
61#[cfg(test)]
62mod tests {
63    use super::*;
64    use crate::{DEFAULT_BASE_URL, DeepSeekClient};
65
66    fn get_client() -> DeepSeekClient {
67        DeepSeekClient::new(
68            std::env::var("DEEPSEEK_API_KEY").expect("DEEPSEEK_API_KEY is not set"),
69            DEFAULT_BASE_URL.clone(),
70        )
71    }
72
73    fn get_builder() -> ChatRequestBuilder {
74        ChatRequestBuilder::default()
75            .client(get_client())
76            .model("deepseek-v4-flash")
77            .thinking(Thinking::disabled())
78    }
79
80    #[tokio::test]
81    async fn chat() {
82        let req = get_builder()
83            .message(ChatMessage::User {
84                content: "Hi".into(),
85                name: None,
86            })
87            .max_tokens(5_u32)
88            .logprobs(true)
89            .top_logprobs(2_u32)
90            .build()
91            .unwrap();
92        let response = req.send().await.unwrap();
93        println!("{:#?}", response);
94    }
95
96    #[tokio::test]
97    async fn api_error() {
98        // Send a request with a deliberately invalid model to verify API error handling
99        let req = get_builder()
100            .model("invalid-model-name")
101            .message(ChatMessage::User {
102                content: "Hi".into(),
103                name: None,
104            })
105            .build()
106            .unwrap();
107        let response = req.send().await;
108        assert!(response.is_err());
109        if let Err(err) = response {
110            assert!(matches!(err, DeepSeekError::Api { .. }));
111            if let DeepSeekError::Api {
112                error,
113                status,
114                body,
115            } = err
116            {
117                assert_eq!(status, Some(400));
118                assert!(body.is_some());
119                assert_eq!(error.error_type, "invalid_request_error");
120                assert_eq!(error.code.as_deref(), Some("invalid_request_error"));
121            } else {
122                panic!("Expected DeepSeekError::Api");
123            }
124        }
125    }
126
127    #[tokio::test]
128    async fn chat_tool_call() {
129        let mut messages = vec![ChatMessage::User {
130            content: "How's the weather in Hangzhou, Zhejiang?".into(),
131            name: None,
132        }];
133        let req_tool = Tool::new(
134            "get_weather",
135            "Get weather of a location, the user should supply a location first.",
136            Some(serde_json::json!({
137                "type": "object",
138                "properties": {
139                    "location": {
140                        "type": "string",
141                        "description": "The city and state, e.g. San Francisco, CA"
142                    },
143                },
144                "required": ["location"]
145            })),
146        );
147        let req = get_builder()
148            .tool(req_tool.clone())
149            .messages(messages.clone())
150            .build()
151            .unwrap();
152        let message = req.send().await.unwrap().choices[0].clone().message;
153        let Some(tool_calls) = message.tool_calls.clone() else {
154            return;
155        };
156        let tool_call = tool_calls[0].clone();
157        messages.push(ChatMessage::Assistant {
158            content: message.content,
159            name: None,
160            tool_calls: Some(tool_calls),
161        });
162        messages.push(ChatMessage::Tool {
163            tool_call_id: tool_call.id,
164            content: "24°C".to_string(),
165        });
166
167        let req2 = get_builder()
168            .tool(req_tool)
169            .messages(messages)
170            .build()
171            .unwrap();
172        let response = req2.send().await.unwrap();
173        println!("{:#?}", response);
174        assert!(
175            response.choices[0]
176                .message
177                .content
178                .as_ref()
179                .unwrap()
180                .contains("24°C")
181        );
182    }
183
184    #[tokio::test]
185    async fn chat_stream_async() {
186        let req = get_builder()
187            .message(ChatMessage::User {
188                content: "Hi".into(),
189                name: None,
190            })
191            .max_tokens(16_u32)
192            .build()
193            .unwrap();
194
195        let mut rx = req.stream().await.unwrap();
196        while let Some(item) = rx.recv().await {
197            match item {
198                Ok(chunk) => println!("Model>\t {:#?}", chunk),
199                Err(err) => eprintln!("Error>\t {:#?}", err),
200            }
201        }
202    }
203
204    #[test]
205    fn chat_stream_blocking() {
206        let req = get_builder()
207            .message(ChatMessage::User {
208                content: "Hi".into(),
209                name: None,
210            })
211            .max_tokens(16_u32)
212            .build()
213            .unwrap();
214
215        let mut stream = req.stream_blocking().unwrap();
216        let mut content = String::new();
217
218        for item in stream.by_ref().take(50) {
219            let chunk = item.unwrap();
220            for choice in chunk.choices {
221                if let Some(delta_content) = choice.delta.content {
222                    content.push_str(&delta_content);
223                }
224            }
225        }
226
227        println!("Model>\t {}", content);
228        assert!(!content.is_empty());
229    }
230}