deepseek_sdk/chat/
stream.rs1use 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;
9pub type ChatStreamItem = Result<ChatStream, DeepSeekError>;
11
12pub 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 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}