use std::collections::BTreeMap;
use super::response::streaming::{
ChatCompletionChunk, ChoiceDelta, ChoiceDeltaToolCallType, CompletionUsage, FinishReason,
};
use crate::chat::Role;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AccumulatedFunctionCall {
pub name: Option<String>,
pub arguments: String,
}
#[derive(Debug, Clone, Default)]
struct ChoiceAccumulator {
role: Option<Role>,
content: String,
#[cfg(feature = "reasoning")]
reasoning_content: String,
#[cfg(feature = "vllm")]
reasoning: String,
refusal: String,
function_call: Option<AccumulatedFunctionCall>,
tool_calls: BTreeMap<u32, ToolCallAccumulator>,
annotations: Vec<crate::chat::Annotation>,
audio: Option<crate::chat::ChatCompletionAudio>,
finish_reason: Option<FinishReason>,
}
#[derive(Debug, Clone, Default)]
struct ToolCallAccumulator {
id: Option<String>,
type_: Option<ChoiceDeltaToolCallType>,
name: String,
arguments: String,
}
impl ChoiceAccumulator {
fn push_delta(&mut self, delta: &ChoiceDelta) {
if delta.role.is_some() {
self.role = delta.role.clone();
}
if let Some(text) = &delta.content {
self.content.push_str(text);
}
#[cfg(feature = "reasoning")]
if let Some(text) = &delta.reasoning_content {
self.reasoning_content.push_str(text);
}
#[cfg(feature = "vllm")]
if let Some(text) = &delta.reasoning {
self.reasoning.push_str(text);
}
if let Some(text) = &delta.refusal {
self.refusal.push_str(text);
}
if let Some(function_call) = &delta.function_call {
let acc = self.function_call.get_or_insert_with(Default::default);
if function_call.name.is_some() {
acc.name = function_call.name.clone();
}
if let Some(arguments) = &function_call.arguments {
acc.arguments.push_str(arguments);
}
}
for tool_call in delta.tool_calls.iter().flatten() {
let acc = self.tool_calls.entry(tool_call.index).or_default();
if tool_call.id.is_some() {
acc.id = tool_call.id.clone();
}
if tool_call.type_.is_some() {
acc.type_ = tool_call.type_.clone();
}
if let Some(function) = &tool_call.function {
if function.name.is_some() {
acc.name = function.name.clone().unwrap_or_default();
}
if let Some(arguments) = &function.arguments {
acc.arguments.push_str(arguments);
}
}
}
if delta.annotations.is_some() {
self.annotations = delta.annotations.clone().unwrap_or_default();
}
if let Some(audio) = &delta.audio {
let acc = self
.audio
.get_or_insert_with(|| crate::chat::ChatCompletionAudio {
id: audio.id.clone(),
data: String::new(),
expires_at: audio.expires_at,
transcript: String::new(),
});
acc.data.push_str(&audio.data);
acc.transcript.push_str(&audio.transcript);
}
}
fn into_message(self) -> crate::chat::ChatCompletionMessage {
crate::chat::ChatCompletionMessage {
role: Role::Assistant,
audio: self.audio,
content: (!self.content.is_empty()).then_some(self.content),
#[cfg(feature = "reasoning")]
reasoning_content: (!self.reasoning_content.is_empty())
.then_some(self.reasoning_content),
#[cfg(feature = "vllm")]
reasoning: (!self.reasoning.is_empty()).then_some(self.reasoning),
tool_calls: (!self.tool_calls.is_empty()).then(|| {
self.tool_calls
.into_values()
.map(|tool_call| match tool_call.type_ {
Some(ChoiceDeltaToolCallType::Custom) => {
crate::chat::ChatCompletionMessageToolCall::Custom {
id: tool_call.id.unwrap_or_default(),
custom: crate::chat::MessageToolCallCustom {
input: tool_call.arguments,
name: tool_call.name,
},
}
}
_ => crate::chat::ChatCompletionMessageToolCall::Function {
id: tool_call.id.unwrap_or_default(),
function: crate::chat::MessageToolCallFunction {
arguments: tool_call.arguments,
name: tool_call.name,
},
},
})
.collect()
}),
refusal: (!self.refusal.is_empty()).then_some(self.refusal),
annotations: (!self.annotations.is_empty()).then_some(self.annotations),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ChatCompletionAccumulator {
choices: BTreeMap<u32, ChoiceAccumulator>,
usage: Option<CompletionUsage>,
}
impl ChatCompletionAccumulator {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, chunk: &ChatCompletionChunk) {
if chunk.usage.is_some() {
self.usage = chunk.usage.clone();
}
for choice in &chunk.choices {
let acc = self.choices.entry(choice.index).or_default();
acc.push_delta(&choice.delta);
if choice.finish_reason.is_some() {
acc.finish_reason = choice.finish_reason.clone();
}
}
}
#[must_use]
pub fn content(&self) -> &str {
self.choice().map_or("", |acc| acc.content.as_str())
}
#[cfg(feature = "reasoning")]
#[must_use]
pub fn reasoning_content(&self) -> &str {
self.choice()
.map_or("", |acc| acc.reasoning_content.as_str())
}
#[cfg(feature = "vllm")]
#[must_use]
pub fn reasoning(&self) -> &str {
self.choice().map_or("", |acc| acc.reasoning.as_str())
}
#[must_use]
pub fn finish_reason(&self) -> Option<&FinishReason> {
self.choice().and_then(|acc| acc.finish_reason.as_ref())
}
#[must_use]
pub fn usage(&self) -> Option<&CompletionUsage> {
self.usage.as_ref()
}
#[must_use]
pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
self.choice().and_then(|acc| acc.function_call.as_ref())
}
#[must_use]
pub fn into_message(mut self) -> crate::chat::ChatCompletionMessage {
self.choices.remove(&0).unwrap_or_default().into_message()
}
#[must_use]
pub fn into_messages(self) -> BTreeMap<u32, crate::chat::ChatCompletionMessage> {
self.choices
.into_iter()
.map(|(index, acc)| (index, acc.into_message()))
.collect()
}
fn choice(&self) -> Option<&ChoiceAccumulator> {
self.choices.get(&0)
}
}
#[cfg(test)]
mod test {
use std::str::FromStr;
use super::*;
fn chunk(json: &str) -> ChatCompletionChunk {
ChatCompletionChunk::from_str(json).expect("test chunk must deserialize")
}
#[test]
fn assembles_content_and_tool_calls_by_index() {
let mut acc = ChatCompletionAccumulator::new();
for json in [
r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
r#"{"id":"1","choices":[{"index":0,"delta":{"content":", world"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":"I cannot"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
r#"{"id":"1","choices":[{"index":0,"delta":{"refusal":" help with that"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
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"}"#,
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"}"#,
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"}"#,
r#"{"id":"1","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
r#"{"id":"1","choices":[],"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":9,"prompt_tokens":17,"total_tokens":26}}"#,
] {
acc.push(&chunk(json));
}
assert_eq!(acc.content(), "Hello, world");
assert!(matches!(acc.finish_reason(), Some(FinishReason::ToolCalls)));
let message = acc.clone().into_message();
assert_eq!(message.refusal.as_deref(), Some("I cannot help with that"));
assert_eq!(acc.usage().expect("usage").total_tokens, 26);
let message = acc.into_message();
assert_eq!(message.content.as_deref(), Some("Hello, world"));
let Some(tool_calls) = message.tool_calls else {
panic!("tool calls must be assembled");
};
assert_eq!(tool_calls.len(), 2);
let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[0]
else {
panic!("assembled tool call must be a function call");
};
assert_eq!(id, "call_1");
assert_eq!(function.name, "get_weather");
assert_eq!(function.arguments, r#"{"city":"Paris"}"#);
let crate::chat::ChatCompletionMessageToolCall::Function { id, function } = &tool_calls[1]
else {
panic!("assembled tool call must be a function call");
};
assert_eq!(id, "call_2");
assert_eq!(function.name, "get_time");
assert_eq!(function.arguments, "{}");
}
#[test]
fn assembles_deprecated_function_call() {
let mut acc = ChatCompletionAccumulator::new();
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"name":"rgb","arguments":"{\"r\":"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"function_call":{"arguments":"1,\"g\":2}"}},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
let function_call = acc.function_call().expect("function call");
assert_eq!(function_call.name.as_deref(), Some("rgb"));
assert_eq!(function_call.arguments, r#"{"r":1,"g":2}"#);
}
#[test]
fn usage_only_chunk_with_null_choices_updates_usage() {
let mut acc = ChatCompletionAccumulator::new();
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
acc.push(&chunk(
r#"{"id":"1","choices":null,"created":1,"model":"m","object":"chat.completion.chunk","usage":{"completion_tokens":1,"prompt_tokens":2,"total_tokens":3}}"#,
));
assert_eq!(acc.content(), "Hi");
assert_eq!(acc.usage().expect("usage").total_tokens, 3);
}
#[cfg(feature = "reasoning")]
#[test]
fn assembles_reasoning_content() {
let mut acc = ChatCompletionAccumulator::new();
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"Think"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"reasoning_content":"ing…"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
acc.push(&chunk(
r#"{"id":"1","choices":[{"index":0,"delta":{"content":"Answer"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#,
));
assert_eq!(acc.reasoning_content(), "Thinking…");
assert_eq!(acc.content(), "Answer");
let message = acc.into_message();
assert_eq!(message.reasoning_content.as_deref(), Some("Thinking…"));
}
#[test]
fn assembles_custom_tool_calls() {
let mut acc = ChatCompletionAccumulator::new();
acc.push(&chunk(
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"}"#,
));
let message = acc.into_message();
let tool_calls = message.tool_calls.expect("tool calls");
let crate::chat::ChatCompletionMessageToolCall::Custom { id, custom } = &tool_calls[0]
else {
panic!("custom tool call must keep its type");
};
assert_eq!(id, "call_1");
assert_eq!(custom.name, "get_weather");
assert_eq!(custom.input, r#"{"city":"Paris"}"#);
}
#[test]
fn accumulates_choices_independently() {
let mut acc = ChatCompletionAccumulator::new();
acc.push(&chunk(
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"}"#,
));
let messages = acc.into_messages();
assert_eq!(messages.len(), 2);
assert_eq!(messages[&0].content.as_deref(), Some("zero"));
assert_eq!(messages[&1].content.as_deref(), Some("one"));
}
#[test]
fn empty_accumulator_yields_empty_assistant_message() {
let message = ChatCompletionAccumulator::new().into_message();
assert_eq!(message.content, None);
assert!(message.tool_calls.is_none());
}
}