1use std::collections::BTreeMap;
47
48use super::response::streaming::{
49 ChatCompletionChunk, ChoiceDelta, ChoiceDeltaToolCallType, CompletionUsage, FinishReason,
50};
51use crate::chat::Role;
52
53#[derive(Debug, Clone, Default, PartialEq, Eq)]
56pub struct AccumulatedFunctionCall {
57 pub name: Option<String>,
59 pub arguments: String,
62}
63
64#[derive(Debug, Clone, Default)]
66struct ChoiceAccumulator {
67 role: Option<Role>,
68 content: String,
69 #[cfg(feature = "reasoning")]
70 reasoning_content: String,
71 refusal: String,
72 function_call: Option<AccumulatedFunctionCall>,
73 tool_calls: BTreeMap<u32, ToolCallAccumulator>,
74 annotations: Vec<crate::chat::Annotation>,
75 audio: Option<crate::chat::ChatCompletionAudio>,
76 finish_reason: Option<FinishReason>,
77}
78
79#[derive(Debug, Clone, Default)]
81struct ToolCallAccumulator {
82 id: Option<String>,
83 type_: Option<ChoiceDeltaToolCallType>,
84 name: String,
85 arguments: String,
86}
87
88impl ChoiceAccumulator {
89 fn push_delta(&mut self, delta: &ChoiceDelta) {
90 if delta.role.is_some() {
91 self.role = delta.role.clone();
92 }
93 if let Some(text) = &delta.content {
94 self.content.push_str(text);
95 }
96 #[cfg(feature = "reasoning")]
97 if let Some(text) = &delta.reasoning_content {
98 self.reasoning_content.push_str(text);
99 }
100 if let Some(text) = &delta.refusal {
101 self.refusal.push_str(text);
102 }
103 if let Some(function_call) = &delta.function_call {
104 let acc = self.function_call.get_or_insert_with(Default::default);
105 if function_call.name.is_some() {
106 acc.name = function_call.name.clone();
107 }
108 if let Some(arguments) = &function_call.arguments {
109 acc.arguments.push_str(arguments);
110 }
111 }
112 for tool_call in delta.tool_calls.iter().flatten() {
113 let acc = self.tool_calls.entry(tool_call.index).or_default();
114 if tool_call.id.is_some() {
115 acc.id = tool_call.id.clone();
116 }
117 if tool_call.type_.is_some() {
118 acc.type_ = tool_call.type_.clone();
119 }
120 if let Some(function) = &tool_call.function {
121 if function.name.is_some() {
122 acc.name = function.name.clone().unwrap_or_default();
123 }
124 if let Some(arguments) = &function.arguments {
125 acc.arguments.push_str(arguments);
126 }
127 }
128 }
129 if delta.annotations.is_some() {
130 self.annotations = delta.annotations.clone().unwrap_or_default();
131 }
132 if let Some(audio) = &delta.audio {
133 let acc = self
136 .audio
137 .get_or_insert_with(|| crate::chat::ChatCompletionAudio {
138 id: audio.id.clone(),
139 data: String::new(),
140 expires_at: audio.expires_at,
141 transcript: String::new(),
142 });
143 acc.data.push_str(&audio.data);
144 acc.transcript.push_str(&audio.transcript);
145 }
146 }
147
148 fn into_message(self) -> crate::chat::ChatCompletionMessage {
149 crate::chat::ChatCompletionMessage {
150 role: Role::Assistant,
151 audio: self.audio,
152 content: (!self.content.is_empty()).then_some(self.content),
153 #[cfg(feature = "reasoning")]
154 reasoning_content: (!self.reasoning_content.is_empty())
155 .then_some(self.reasoning_content),
156 tool_calls: (!self.tool_calls.is_empty()).then(|| {
157 self.tool_calls
158 .into_values()
159 .map(|tool_call| match tool_call.type_ {
160 Some(ChoiceDeltaToolCallType::Custom) => {
164 crate::chat::ChatCompletionMessageToolCall::Custom {
165 id: tool_call.id.unwrap_or_default(),
166 custom: crate::chat::MessageToolCallCustom {
167 input: tool_call.arguments,
168 name: tool_call.name,
169 },
170 }
171 }
172 _ => crate::chat::ChatCompletionMessageToolCall::Function {
173 id: tool_call.id.unwrap_or_default(),
174 function: crate::chat::MessageToolCallFunction {
175 arguments: tool_call.arguments,
176 name: tool_call.name,
177 },
178 },
179 })
180 .collect()
181 }),
182 refusal: (!self.refusal.is_empty()).then_some(self.refusal),
183 annotations: (!self.annotations.is_empty()).then_some(self.annotations),
184 }
185 }
186}
187
188#[derive(Debug, Clone, Default)]
200pub struct ChatCompletionAccumulator {
201 choices: BTreeMap<u32, ChoiceAccumulator>,
202 usage: Option<CompletionUsage>,
203}
204
205impl ChatCompletionAccumulator {
206 #[must_use]
208 pub fn new() -> Self {
209 Self::default()
210 }
211
212 pub fn push(&mut self, chunk: &ChatCompletionChunk) {
214 if chunk.usage.is_some() {
215 self.usage = chunk.usage.clone();
216 }
217 for choice in &chunk.choices {
218 let acc = self.choices.entry(choice.index).or_default();
219 acc.push_delta(&choice.delta);
220 if choice.finish_reason.is_some() {
221 acc.finish_reason = choice.finish_reason.clone();
222 }
223 }
224 }
225
226 #[must_use]
229 pub fn content(&self) -> &str {
230 self.choice().map_or("", |acc| acc.content.as_str())
231 }
232
233 #[cfg(feature = "reasoning")]
235 #[must_use]
236 pub fn reasoning_content(&self) -> &str {
237 self.choice()
238 .map_or("", |acc| acc.reasoning_content.as_str())
239 }
240
241 #[must_use]
243 pub fn finish_reason(&self) -> Option<&FinishReason> {
244 self.choice().and_then(|acc| acc.finish_reason.as_ref())
245 }
246
247 #[must_use]
250 pub fn usage(&self) -> Option<&CompletionUsage> {
251 self.usage.as_ref()
252 }
253
254 #[must_use]
257 pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
258 self.choice().and_then(|acc| acc.function_call.as_ref())
259 }
260
261 #[must_use]
265 pub fn into_message(mut self) -> crate::chat::ChatCompletionMessage {
266 self.choices.remove(&0).unwrap_or_default().into_message()
267 }
268
269 #[must_use]
272 pub fn into_messages(self) -> BTreeMap<u32, crate::chat::ChatCompletionMessage> {
273 self.choices
274 .into_iter()
275 .map(|(index, acc)| (index, acc.into_message()))
276 .collect()
277 }
278
279 fn choice(&self) -> Option<&ChoiceAccumulator> {
280 self.choices.get(&0)
281 }
282}
283
284#[cfg(test)]
285mod test {
286 use std::str::FromStr;
287
288 use super::*;
289
290 fn chunk(json: &str) -> ChatCompletionChunk {
291 ChatCompletionChunk::from_str(json).expect("test chunk must deserialize")
292 }
293
294 #[test]
295 fn assembles_content_and_tool_calls_by_index() {
296 let mut acc = ChatCompletionAccumulator::new();
297 for json in [
298 r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
299 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
300 r#"{"id":"1","choices":[{"index":0,"delta":{"content":", world"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
301 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":"I cannot"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
302 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":" help with that"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
303 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"}"#,
304 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"}"#,
305 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"}"#,
306 r#"{"id":"1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
307 r#"{"id":"1","choices":[],"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":9,"prompt_tokens":17,"total_tokens":26}}"#,
308 ] {
309 acc.push(&chunk(json));
310 }
311
312 assert_eq!(acc.content(), "Hello, world");
313 assert!(matches!(acc.finish_reason(), Some(FinishReason::ToolCalls)));
314 let message = acc.clone().into_message();
315 assert_eq!(message.refusal.as_deref(), Some("I cannot help with that"));
316 assert_eq!(acc.usage().expect("usage").total_tokens, 26);
317
318 let message = acc.into_message();
319 assert_eq!(message.content.as_deref(), Some("Hello, world"));
320
321 let Some(tool_calls) = message.tool_calls else {
322 panic!("tool calls must be assembled");
323 };
324 assert_eq!(tool_calls.len(), 2);
325 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[0]
326 else {
327 panic!("assembled tool call must be a function call");
328 };
329 assert_eq!(id, "call_1");
330 assert_eq!(function.name, "get_weather");
331 assert_eq!(function.arguments, r#"{"city":"Paris"}"#);
332 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[1]
333 else {
334 panic!("assembled tool call must be a function call");
335 };
336 assert_eq!(id, "call_2");
337 assert_eq!(function.name, "get_time");
338 assert_eq!(function.arguments, "{}");
339 }
340
341 #[test]
342 fn assembles_deprecated_function_call() {
343 let mut acc = ChatCompletionAccumulator::new();
344 acc.push(&chunk(
345 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"name":"rgb","arguments":"{\"r\":"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
346 ));
347 acc.push(&chunk(
348 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"arguments":"1,\"g\":2}"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
349 ));
350
351 let function_call = acc.function_call().expect("function call");
352 assert_eq!(function_call.name.as_deref(), Some("rgb"));
353 assert_eq!(function_call.arguments, r#"{"r":1,"g":2}"#);
354 }
355
356 #[test]
360 fn usage_only_chunk_with_null_choices_updates_usage() {
361 let mut acc = ChatCompletionAccumulator::new();
362 acc.push(&chunk(
363 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
364 ));
365 acc.push(&chunk(
366 r#"{"id":"1","choices":null,"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":1,"prompt_tokens":2,"total_tokens":3}}"#,
367 ));
368
369 assert_eq!(acc.content(), "Hi");
370 assert_eq!(acc.usage().expect("usage").total_tokens, 3);
371 }
372
373 #[cfg(feature = "reasoning")]
376 #[test]
377 fn assembles_reasoning_content() {
378 let mut acc = ChatCompletionAccumulator::new();
379 acc.push(&chunk(
380 r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"Think"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
381 ));
382 acc.push(&chunk(
383 r#"{"id":"1","choices":[{"index":0,"delta":{"reasoning_content":"ing…"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
384 ));
385 acc.push(&chunk(
386 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Answer"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
387 ));
388
389 assert_eq!(acc.reasoning_content(), "Thinking…");
390 assert_eq!(acc.content(), "Answer");
391 let message = acc.into_message();
392 assert_eq!(message.reasoning_content.as_deref(), Some("Thinking…"));
393 }
394
395 #[test]
397 fn assembles_custom_tool_calls() {
398 let mut acc = ChatCompletionAccumulator::new();
399 acc.push(&chunk(
400 r#"{"id":"1","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"custom","function":{"name":"get_weather","arguments":"{\"city\":\"Paris\"}"}}]},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
401 ));
402
403 let message = acc.into_message();
404 let tool_calls = message.tool_calls.expect("tool calls");
405 let crate::chat::ChatCompletionMessageToolCall::Custom { id, custom } = &tool_calls[0]
406 else {
407 panic!("custom tool call must keep its type");
408 };
409 assert_eq!(id, "call_1");
410 assert_eq!(custom.name, "get_weather");
411 assert_eq!(custom.input, r#"{"city":"Paris"}"#);
412 }
413
414 #[test]
415 fn accumulates_choices_independently() {
416 let mut acc = ChatCompletionAccumulator::new();
417 acc.push(&chunk(
418 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"}"#,
419 ));
420
421 let messages = acc.into_messages();
422 assert_eq!(messages.len(), 2);
423 assert_eq!(messages[&0].content.as_deref(), Some("zero"));
424 assert_eq!(messages[&1].content.as_deref(), Some("one"));
425 }
426
427 #[test]
428 fn empty_accumulator_yields_empty_assistant_message() {
429 let message = ChatCompletionAccumulator::new().into_message();
430 assert_eq!(message.content, None);
431 assert!(message.tool_calls.is_none());
432 }
433}