1use std::collections::BTreeMap;
50
51use super::response::streaming::{
52 ChatCompletionChunk, ChoiceDelta, ChoiceDeltaToolCallType, CompletionRole, CompletionUsage,
53 FinishReason,
54};
55
56#[derive(Debug, Clone, Default, PartialEq, Eq)]
59pub struct AccumulatedFunctionCall {
60 pub name: Option<String>,
62 pub arguments: String,
65}
66
67#[derive(Debug, Clone, Default)]
69struct ChoiceAccumulator {
70 role: Option<CompletionRole>,
71 content: String,
72 #[cfg(feature = "deepseek")]
73 reasoning_content: String,
74 refusal: String,
75 function_call: Option<AccumulatedFunctionCall>,
76 tool_calls: BTreeMap<usize, ToolCallAccumulator>,
77 #[cfg(feature = "azure")]
78 annotations: Vec<crate::chat::Annotation>,
79 #[cfg(feature = "azure")]
80 audio: Option<crate::chat::ChatCompletionAudio>,
81 finish_reason: Option<FinishReason>,
82}
83
84#[derive(Debug, Clone, Default)]
86struct ToolCallAccumulator {
87 id: Option<String>,
88 type_: Option<ChoiceDeltaToolCallType>,
89 name: String,
90 arguments: String,
91}
92
93impl ChoiceAccumulator {
94 fn push_delta(&mut self, delta: &ChoiceDelta) {
95 if delta.role.is_some() {
96 self.role = delta.role.clone();
97 }
98 if let Some(text) = &delta.content {
99 self.content.push_str(text);
100 }
101 #[cfg(feature = "deepseek")]
102 if let Some(text) = &delta.reasoning_content {
103 self.reasoning_content.push_str(text);
104 }
105 if let Some(text) = &delta.refusal {
106 self.refusal.push_str(text);
107 }
108 if let Some(function_call) = &delta.function_call {
109 let acc = self.function_call.get_or_insert_with(Default::default);
110 if function_call.name.is_some() {
111 acc.name = function_call.name.clone();
112 }
113 if let Some(arguments) = &function_call.arguments {
114 acc.arguments.push_str(arguments);
115 }
116 }
117 for tool_call in delta.tool_calls.iter().flatten() {
118 let acc = self.tool_calls.entry(tool_call.index).or_default();
119 if tool_call.id.is_some() {
120 acc.id = tool_call.id.clone();
121 }
122 if tool_call.type_.is_some() {
123 acc.type_ = tool_call.type_.clone();
124 }
125 if let Some(function) = &tool_call.function {
126 if function.name.is_some() {
127 acc.name = function.name.clone().unwrap_or_default();
128 }
129 if let Some(arguments) = &function.arguments {
130 acc.arguments.push_str(arguments);
131 }
132 }
133 }
134 #[cfg(feature = "azure")]
135 if delta.annotations.is_some() {
136 self.annotations = delta.annotations.clone().unwrap_or_default();
137 }
138 #[cfg(feature = "azure")]
139 if let Some(audio) = &delta.audio {
140 let acc = self
143 .audio
144 .get_or_insert_with(|| crate::chat::ChatCompletionAudio {
145 id: audio.id.clone(),
146 data: String::new(),
147 expires_at: audio.expires_at,
148 transcript: String::new(),
149 });
150 acc.data.push_str(&audio.data);
151 acc.transcript.push_str(&audio.transcript);
152 }
153 }
154
155 fn into_message(self) -> crate::chat::ChatCompletionMessage {
156 crate::chat::ChatCompletionMessage {
157 role: crate::chat::ResponseRole::Assistant,
158 audio: {
159 #[cfg(feature = "azure")]
160 let audio = self.audio;
161 #[cfg(not(feature = "azure"))]
162 let audio = None;
163 audio
164 },
165 content: (!self.content.is_empty()).then_some(self.content),
166 #[cfg(feature = "deepseek")]
167 reasoning_content: (!self.reasoning_content.is_empty())
168 .then_some(self.reasoning_content),
169 tool_calls: (!self.tool_calls.is_empty()).then(|| {
170 self.tool_calls
171 .into_values()
172 .map(
173 |tool_call| crate::chat::ChatCompletionMessageToolCall::Function {
174 id: tool_call.id.unwrap_or_default(),
175 function: crate::chat::MessageToolCallFunction {
176 arguments: tool_call.arguments,
177 name: tool_call.name,
178 },
179 },
180 )
181 .collect()
182 }),
183 refusal: (!self.refusal.is_empty()).then_some(self.refusal),
184 annotations: {
185 #[cfg(feature = "azure")]
186 let annotations = (!self.annotations.is_empty()).then_some(self.annotations);
187 #[cfg(not(feature = "azure"))]
188 let annotations = None;
189 annotations
190 },
191 }
192 }
193}
194
195#[derive(Debug, Clone, Default)]
207pub struct ChatCompletionAccumulator {
208 choices: BTreeMap<u32, ChoiceAccumulator>,
209 usage: Option<CompletionUsage>,
210}
211
212impl ChatCompletionAccumulator {
213 #[must_use]
215 pub fn new() -> Self {
216 Self::default()
217 }
218
219 pub fn push(&mut self, chunk: &ChatCompletionChunk) {
221 if chunk.usage.is_some() {
222 self.usage = chunk.usage.clone();
223 }
224 for choice in &chunk.choices {
225 let acc = self.choices.entry(choice.index).or_default();
226 acc.push_delta(&choice.delta);
227 if choice.finish_reason.is_some() {
228 acc.finish_reason = choice.finish_reason.clone();
229 }
230 }
231 }
232
233 #[must_use]
236 pub fn content(&self) -> &str {
237 self.choice().map_or("", |acc| acc.content.as_str())
238 }
239
240 #[cfg(feature = "deepseek")]
243 #[must_use]
244 pub fn reasoning_content(&self) -> &str {
245 self.choice()
246 .map_or("", |acc| acc.reasoning_content.as_str())
247 }
248
249 #[must_use]
251 pub fn finish_reason(&self) -> Option<&FinishReason> {
252 self.choice().and_then(|acc| acc.finish_reason.as_ref())
253 }
254
255 #[must_use]
258 pub fn usage(&self) -> Option<&CompletionUsage> {
259 self.usage.as_ref()
260 }
261
262 #[must_use]
265 pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
266 self.choice().and_then(|acc| acc.function_call.as_ref())
267 }
268
269 #[must_use]
273 pub fn into_message(mut self) -> crate::chat::ChatCompletionMessage {
274 self.choices.remove(&0).unwrap_or_default().into_message()
275 }
276
277 #[must_use]
280 pub fn into_messages(self) -> BTreeMap<u32, crate::chat::ChatCompletionMessage> {
281 self.choices
282 .into_iter()
283 .map(|(index, acc)| (index, acc.into_message()))
284 .collect()
285 }
286
287 fn choice(&self) -> Option<&ChoiceAccumulator> {
288 self.choices.get(&0)
289 }
290}
291
292#[cfg(test)]
293mod test {
294 use std::str::FromStr;
295
296 use super::*;
297
298 fn chunk(json: &str) -> ChatCompletionChunk {
299 ChatCompletionChunk::from_str(json).expect("test chunk must deserialize")
300 }
301
302 #[test]
303 fn assembles_content_and_tool_calls_by_index() {
304 let mut acc = ChatCompletionAccumulator::new();
305 for json in [
306 r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
307 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
308 r#"{"id":"1","choices":[{"index":0,"delta":{"content":", world"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
309 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":"I cannot"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
310 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":" help with that"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
311 r#"{"id":"1","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"city"}}]},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
312 r#"{"id":"1","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\":\"Paris\"}"}}]},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
313 r#"{"id":"1","choices":[{"index":0,"delta":{"tool_calls":[{"index":1,"id":"call_2","type":"function","function":{"name":"get_time","arguments":"{}"}}]},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
314 r#"{"id":"1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
315 r#"{"id":"1","choices":[],"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":9,"prompt_tokens":17,"total_tokens":26}}"#,
316 ] {
317 acc.push(&chunk(json));
318 }
319
320 assert_eq!(acc.content(), "Hello, world");
321 assert!(matches!(acc.finish_reason(), Some(FinishReason::ToolCalls)));
322 let message = acc.clone().into_message();
323 assert_eq!(message.refusal.as_deref(), Some("I cannot help with that"));
324 assert_eq!(acc.usage().expect("usage").total_tokens, 26);
325
326 let message = acc.into_message();
327 assert_eq!(message.content.as_deref(), Some("Hello, world"));
328
329 let Some(tool_calls) = message.tool_calls else {
330 panic!("tool calls must be assembled");
331 };
332 assert_eq!(tool_calls.len(), 2);
333 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[0]
334 else {
335 panic!("assembled tool call must be a function call");
336 };
337 assert_eq!(id, "call_1");
338 assert_eq!(function.name, "get_weather");
339 assert_eq!(function.arguments, r#"{"city":"Paris"}"#);
340 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[1]
341 else {
342 panic!("assembled tool call must be a function call");
343 };
344 assert_eq!(id, "call_2");
345 assert_eq!(function.name, "get_time");
346 assert_eq!(function.arguments, "{}");
347 }
348
349 #[test]
350 fn assembles_deprecated_function_call() {
351 let mut acc = ChatCompletionAccumulator::new();
352 acc.push(&chunk(
353 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"name":"rgb","arguments":"{\"r\":"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
354 ));
355 acc.push(&chunk(
356 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"arguments":"1,\"g\":2}"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
357 ));
358
359 let function_call = acc.function_call().expect("function call");
360 assert_eq!(function_call.name.as_deref(), Some("rgb"));
361 assert_eq!(function_call.arguments, r#"{"r":1,"g":2}"#);
362 }
363
364 #[test]
365 fn accumulates_choices_independently() {
366 let mut acc = ChatCompletionAccumulator::new();
367 acc.push(&chunk(
368 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"zero"},"finish_reason":null},{"index":1,"delta":{"content":"one"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
369 ));
370
371 let messages = acc.into_messages();
372 assert_eq!(messages.len(), 2);
373 assert_eq!(messages[&0].content.as_deref(), Some("zero"));
374 assert_eq!(messages[&1].content.as_deref(), Some("one"));
375 }
376
377 #[test]
378 fn empty_accumulator_yields_empty_assistant_message() {
379 let message = ChatCompletionAccumulator::new().into_message();
380 assert_eq!(message.content, None);
381 assert!(message.tool_calls.is_none());
382 }
383}