Skip to main content

openai_interface/chat/create/
accumulator.rs

1//! Accumulates streamed chat completion chunks into complete messages.
2//!
3//! Streaming responses deliver a message piecewise: text arrives as
4//! `content` fragments spread over many chunks, and a tool call arrives as
5//! fragments whose `arguments` must be concatenated in order, keyed by the
6//! tool call's `index` (the official SDKs perform this assembly for you;
7//! this accumulator does the same for this crate).
8//!
9//! Feed every [`ChatCompletionChunk`] you receive into
10//! [`ChatCompletionAccumulator::push`], then read the result with
11//! [`ChatCompletionAccumulator::into_message`].
12//!
13//! # Example
14//!
15//! ```rust,no_run
16//! use futures_util::StreamExt;
17//! use openai_interface::chat::create::accumulator::ChatCompletionAccumulator;
18//! use openai_interface::chat::create::request::{Message, RequestBody};
19//! use openai_interface::rest::{default_client, post::PostStream, RequestOptions};
20//!
21//! # async fn example(api_key: String) -> Result<(), Box<dyn std::error::Error>> {
22//! let request = RequestBody {
23//!     messages: vec![Message::User {
24//!         content: "What's the weather in Paris?".into(),
25//!         name: None,
26//!     }],
27//!     model: "deepseek-chat".to_string(),
28//!     stream: Some(true),
29//!     ..Default::default()
30//! };
31//!
32//! let stream = request
33//!     .get_stream_response(&default_client(), "https://api.deepseek.com", &RequestOptions::bearer(api_key))
34//!     .await?;
35//!
36//! let mut accumulator = ChatCompletionAccumulator::new();
37//! let mut stream = stream;
38//! while let Some(chunk) = stream.next().await {
39//!     accumulator.push(&chunk?);
40//! }
41//!
42//! let message = accumulator.into_message();
43//! println!("content: {:?}", message.content);
44//! println!("tool calls: {:?}", message.tool_calls);
45//! # Ok(())
46//! # }
47//! ```
48
49use std::collections::BTreeMap;
50
51use super::response::streaming::{
52    ChatCompletionChunk, ChoiceDelta, ChoiceDeltaToolCallType, CompletionRole, CompletionUsage,
53    FinishReason,
54};
55
56/// A deprecated `function_call` (replaced by `tool_calls`) assembled from
57/// its streamed fragments.
58#[derive(Debug, Clone, Default, PartialEq, Eq)]
59pub struct AccumulatedFunctionCall {
60    /// The name of the function to call, from the fragment that carried it.
61    pub name: Option<String>,
62    /// The arguments to call the function with, concatenated across all
63    /// fragments.
64    pub arguments: String,
65}
66
67/// Accumulated state of one streamed choice.
68#[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/// Accumulated state of one streamed tool call.
85#[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            // The first fragment carries the identity; the byte and
141            // transcript payloads arrive incrementally.
142            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/// Assembles streamed [`ChatCompletionChunk`]s into complete messages.
196///
197/// All choices are accumulated, keyed by their chunk index; the accessors
198/// and [`ChatCompletionAccumulator::into_message`] default to choice `0`
199/// (the only choice unless the request set `n > 1`).
200///
201/// Fragment semantics follow the official SDKs: `content` / `refusal`
202/// / tool call `arguments` concatenate; `role`, tool call `id` / `type` /
203/// `name` and `function_call.name` are overwritten whenever a fragment
204/// carries them; `finish_reason` and `usage` are kept from the last chunk
205/// that provided them.
206#[derive(Debug, Clone, Default)]
207pub struct ChatCompletionAccumulator {
208    choices: BTreeMap<u32, ChoiceAccumulator>,
209    usage: Option<CompletionUsage>,
210}
211
212impl ChatCompletionAccumulator {
213    /// Creates an empty accumulator.
214    #[must_use]
215    pub fn new() -> Self {
216        Self::default()
217    }
218
219    /// Merges one chunk into the accumulated state.
220    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    /// The accumulated text content of choice `0`, empty before any content
234    /// chunk arrives.
235    #[must_use]
236    pub fn content(&self) -> &str {
237        self.choice().map_or("", |acc| acc.content.as_str())
238    }
239
240    /// The accumulated `reasoning_content` of choice `0` (DeepSeek thinking
241    /// models).
242    #[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    /// The `finish_reason` of choice `0`, once a chunk carries it.
250    #[must_use]
251    pub fn finish_reason(&self) -> Option<&FinishReason> {
252        self.choice().and_then(|acc| acc.finish_reason.as_ref())
253    }
254
255    /// The token usage, once a chunk carries it (the final chunk when the
256    /// request set `stream_options: {"include_usage": true}`).
257    #[must_use]
258    pub fn usage(&self) -> Option<&CompletionUsage> {
259        self.usage.as_ref()
260    }
261
262    /// The deprecated `function_call` of choice `0`, assembled from its
263    /// fragments, if the model used the legacy path.
264    #[must_use]
265    pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
266        self.choice().and_then(|acc| acc.function_call.as_ref())
267    }
268
269    /// Consumes the accumulator, producing the assembled message of choice
270    /// `0`. An accumulator that never saw a choice-`0` chunk yields an
271    /// empty assistant message.
272    #[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    /// Consumes the accumulator, producing every assembled message keyed by
278    /// its choice index.
279    #[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}