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 #[cfg(feature = "vllm")]
72 reasoning: String,
73 refusal: String,
74 function_call: Option<AccumulatedFunctionCall>,
75 tool_calls: BTreeMap<u32, ToolCallAccumulator>,
76 annotations: Vec<crate::chat::Annotation>,
77 audio: Option<crate::chat::ChatCompletionAudio>,
78 finish_reason: Option<FinishReason>,
79}
80
81#[derive(Debug, Clone, Default)]
83struct ToolCallAccumulator {
84 id: Option<String>,
85 type_: Option<ChoiceDeltaToolCallType>,
86 name: String,
87 arguments: String,
88}
89
90impl ChoiceAccumulator {
91 fn push_delta(&mut self, delta: &ChoiceDelta) {
92 if delta.role.is_some() {
93 self.role = delta.role.clone();
94 }
95 if let Some(text) = &delta.content {
96 self.content.push_str(text);
97 }
98 #[cfg(feature = "reasoning")]
99 if let Some(text) = &delta.reasoning_content {
100 self.reasoning_content.push_str(text);
101 }
102 #[cfg(feature = "vllm")]
106 if let Some(text) = &delta.reasoning {
107 self.reasoning.push_str(text);
108 }
109 if let Some(text) = &delta.refusal {
110 self.refusal.push_str(text);
111 }
112 if let Some(function_call) = &delta.function_call {
113 let acc = self.function_call.get_or_insert_with(Default::default);
114 if function_call.name.is_some() {
115 acc.name = function_call.name.clone();
116 }
117 if let Some(arguments) = &function_call.arguments {
118 acc.arguments.push_str(arguments);
119 }
120 }
121 for tool_call in delta.tool_calls.iter().flatten() {
122 let acc = self.tool_calls.entry(tool_call.index).or_default();
123 if tool_call.id.is_some() {
124 acc.id = tool_call.id.clone();
125 }
126 if tool_call.type_.is_some() {
127 acc.type_ = tool_call.type_.clone();
128 }
129 if let Some(function) = &tool_call.function {
130 if function.name.is_some() {
131 acc.name = function.name.clone().unwrap_or_default();
132 }
133 if let Some(arguments) = &function.arguments {
134 acc.arguments.push_str(arguments);
135 }
136 }
137 }
138 if delta.annotations.is_some() {
139 self.annotations = delta.annotations.clone().unwrap_or_default();
140 }
141 if let Some(audio) = &delta.audio {
142 let acc = self
145 .audio
146 .get_or_insert_with(|| crate::chat::ChatCompletionAudio {
147 id: audio.id.clone(),
148 data: String::new(),
149 expires_at: audio.expires_at,
150 transcript: String::new(),
151 });
152 acc.data.push_str(&audio.data);
153 acc.transcript.push_str(&audio.transcript);
154 }
155 }
156
157 fn into_message(self) -> crate::chat::ChatCompletionMessage {
158 crate::chat::ChatCompletionMessage {
159 role: Role::Assistant,
160 audio: self.audio,
161 content: (!self.content.is_empty()).then_some(self.content),
162 #[cfg(feature = "reasoning")]
163 reasoning_content: (!self.reasoning_content.is_empty())
164 .then_some(self.reasoning_content),
165 #[cfg(feature = "vllm")]
166 reasoning: (!self.reasoning.is_empty()).then_some(self.reasoning),
167 tool_calls: (!self.tool_calls.is_empty()).then(|| {
168 self.tool_calls
169 .into_values()
170 .map(|tool_call| match tool_call.type_ {
171 Some(ChoiceDeltaToolCallType::Custom) => {
175 crate::chat::ChatCompletionMessageToolCall::Custom {
176 id: tool_call.id.unwrap_or_default(),
177 custom: crate::chat::MessageToolCallCustom {
178 input: tool_call.arguments,
179 name: tool_call.name,
180 },
181 }
182 }
183 _ => crate::chat::ChatCompletionMessageToolCall::Function {
184 id: tool_call.id.unwrap_or_default(),
185 function: crate::chat::MessageToolCallFunction {
186 arguments: tool_call.arguments,
187 name: tool_call.name,
188 },
189 },
190 })
191 .collect()
192 }),
193 refusal: (!self.refusal.is_empty()).then_some(self.refusal),
194 annotations: (!self.annotations.is_empty()).then_some(self.annotations),
195 }
196 }
197}
198
199#[derive(Debug, Clone, Default)]
211pub struct ChatCompletionAccumulator {
212 choices: BTreeMap<u32, ChoiceAccumulator>,
213 usage: Option<CompletionUsage>,
214}
215
216impl ChatCompletionAccumulator {
217 #[must_use]
219 pub fn new() -> Self {
220 Self::default()
221 }
222
223 pub fn push(&mut self, chunk: &ChatCompletionChunk) {
225 if chunk.usage.is_some() {
226 self.usage = chunk.usage.clone();
227 }
228 for choice in &chunk.choices {
229 let acc = self.choices.entry(choice.index).or_default();
230 acc.push_delta(&choice.delta);
231 if choice.finish_reason.is_some() {
232 acc.finish_reason = choice.finish_reason.clone();
233 }
234 }
235 }
236
237 #[must_use]
240 pub fn content(&self) -> &str {
241 self.choice().map_or("", |acc| acc.content.as_str())
242 }
243
244 #[cfg(feature = "reasoning")]
246 #[must_use]
247 pub fn reasoning_content(&self) -> &str {
248 self.choice()
249 .map_or("", |acc| acc.reasoning_content.as_str())
250 }
251
252 #[cfg(feature = "vllm")]
259 #[must_use]
260 pub fn reasoning(&self) -> &str {
261 self.choice().map_or("", |acc| acc.reasoning.as_str())
262 }
263
264 #[must_use]
266 pub fn finish_reason(&self) -> Option<&FinishReason> {
267 self.choice().and_then(|acc| acc.finish_reason.as_ref())
268 }
269
270 #[must_use]
273 pub fn usage(&self) -> Option<&CompletionUsage> {
274 self.usage.as_ref()
275 }
276
277 #[must_use]
280 pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
281 self.choice().and_then(|acc| acc.function_call.as_ref())
282 }
283
284 #[must_use]
288 pub fn into_message(mut self) -> crate::chat::ChatCompletionMessage {
289 self.choices.remove(&0).unwrap_or_default().into_message()
290 }
291
292 #[must_use]
295 pub fn into_messages(self) -> BTreeMap<u32, crate::chat::ChatCompletionMessage> {
296 self.choices
297 .into_iter()
298 .map(|(index, acc)| (index, acc.into_message()))
299 .collect()
300 }
301
302 fn choice(&self) -> Option<&ChoiceAccumulator> {
303 self.choices.get(&0)
304 }
305}
306
307#[cfg(test)]
308mod test {
309 use std::str::FromStr;
310
311 use super::*;
312
313 fn chunk(json: &str) -> ChatCompletionChunk {
314 ChatCompletionChunk::from_str(json).expect("test chunk must deserialize")
315 }
316
317 #[test]
318 fn assembles_content_and_tool_calls_by_index() {
319 let mut acc = ChatCompletionAccumulator::new();
320 for json in [
321 r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
322 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
323 r#"{"id":"1","choices":[{"index":0,"delta":{"content":", world"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
324 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":"I cannot"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
325 r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":" help with that"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
326 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"}"#,
327 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"}"#,
328 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"}"#,
329 r#"{"id":"1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
330 r#"{"id":"1","choices":[],"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":9,"prompt_tokens":17,"total_tokens":26}}"#,
331 ] {
332 acc.push(&chunk(json));
333 }
334
335 assert_eq!(acc.content(), "Hello, world");
336 assert!(matches!(acc.finish_reason(), Some(FinishReason::ToolCalls)));
337 let message = acc.clone().into_message();
338 assert_eq!(message.refusal.as_deref(), Some("I cannot help with that"));
339 assert_eq!(acc.usage().expect("usage").total_tokens, 26);
340
341 let message = acc.into_message();
342 assert_eq!(message.content.as_deref(), Some("Hello, world"));
343
344 let Some(tool_calls) = message.tool_calls else {
345 panic!("tool calls must be assembled");
346 };
347 assert_eq!(tool_calls.len(), 2);
348 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[0]
349 else {
350 panic!("assembled tool call must be a function call");
351 };
352 assert_eq!(id, "call_1");
353 assert_eq!(function.name, "get_weather");
354 assert_eq!(function.arguments, r#"{"city":"Paris"}"#);
355 let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[1]
356 else {
357 panic!("assembled tool call must be a function call");
358 };
359 assert_eq!(id, "call_2");
360 assert_eq!(function.name, "get_time");
361 assert_eq!(function.arguments, "{}");
362 }
363
364 #[test]
365 fn assembles_deprecated_function_call() {
366 let mut acc = ChatCompletionAccumulator::new();
367 acc.push(&chunk(
368 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"name":"rgb","arguments":"{\"r\":"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
369 ));
370 acc.push(&chunk(
371 r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"arguments":"1,\"g\":2}"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
372 ));
373
374 let function_call = acc.function_call().expect("function call");
375 assert_eq!(function_call.name.as_deref(), Some("rgb"));
376 assert_eq!(function_call.arguments, r#"{"r":1,"g":2}"#);
377 }
378
379 #[test]
383 fn usage_only_chunk_with_null_choices_updates_usage() {
384 let mut acc = ChatCompletionAccumulator::new();
385 acc.push(&chunk(
386 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
387 ));
388 acc.push(&chunk(
389 r#"{"id":"1","choices":null,"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":1,"prompt_tokens":2,"total_tokens":3}}"#,
390 ));
391
392 assert_eq!(acc.content(), "Hi");
393 assert_eq!(acc.usage().expect("usage").total_tokens, 3);
394 }
395
396 #[cfg(feature = "reasoning")]
399 #[test]
400 fn assembles_reasoning_content() {
401 let mut acc = ChatCompletionAccumulator::new();
402 acc.push(&chunk(
403 r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"Think"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
404 ));
405 acc.push(&chunk(
406 r#"{"id":"1","choices":[{"index":0,"delta":{"reasoning_content":"ing…"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
407 ));
408 acc.push(&chunk(
409 r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Answer"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
410 ));
411
412 assert_eq!(acc.reasoning_content(), "Thinking…");
413 assert_eq!(acc.content(), "Answer");
414 let message = acc.into_message();
415 assert_eq!(message.reasoning_content.as_deref(), Some("Thinking…"));
416 }
417
418 #[test]
420 fn assembles_custom_tool_calls() {
421 let mut acc = ChatCompletionAccumulator::new();
422 acc.push(&chunk(
423 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"}"#,
424 ));
425
426 let message = acc.into_message();
427 let tool_calls = message.tool_calls.expect("tool calls");
428 let crate::chat::ChatCompletionMessageToolCall::Custom { id, custom } = &tool_calls[0]
429 else {
430 panic!("custom tool call must keep its type");
431 };
432 assert_eq!(id, "call_1");
433 assert_eq!(custom.name, "get_weather");
434 assert_eq!(custom.input, r#"{"city":"Paris"}"#);
435 }
436
437 #[test]
438 fn accumulates_choices_independently() {
439 let mut acc = ChatCompletionAccumulator::new();
440 acc.push(&chunk(
441 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"}"#,
442 ));
443
444 let messages = acc.into_messages();
445 assert_eq!(messages.len(), 2);
446 assert_eq!(messages[&0].content.as_deref(), Some("zero"));
447 assert_eq!(messages[&1].content.as_deref(), Some("one"));
448 }
449
450 #[test]
451 fn empty_accumulator_yields_empty_assistant_message() {
452 let message = ChatCompletionAccumulator::new().into_message();
453 assert_eq!(message.content, None);
454 assert!(message.tool_calls.is_none());
455 }
456}