openai-interface 0.11.0

A low-level Rust interface for the OpenAI API
Documentation
//! Accumulates streamed chat completion chunks into complete messages.
//!
//! Streaming responses deliver a message piecewise: text arrives as
//! `content` fragments spread over many chunks, and a tool call arrives as
//! fragments whose `arguments` must be concatenated in order, keyed by the
//! tool call's `index` (the official SDKs perform this assembly for you;
//! this accumulator does the same for this crate).
//!
//! Feed every [`ChatCompletionChunk`] you receive into
//! [`ChatCompletionAccumulator::push`], then read the result with
//! [`ChatCompletionAccumulator::into_message`].
//!
//! # Example
//!
//! ```rust,no_run
//! use futures_util::StreamExt;
//! use openai_interface::chat::create::accumulator::ChatCompletionAccumulator;
//! use openai_interface::chat::create::request::{Message, RequestBody};
//! use openai_interface::rest::{default_client, post::PostStream, RequestOptions};
//!
//! # async fn example(api_key: String) -> Result<(), Box<dyn std::error::Error>> {
//! let request = RequestBody {
//!     messages: vec![Message::User {
//!         content: "What's the weather in Paris?".into(),
//!         name: None,
//!     }],
//!     model: "deepseek-chat".to_string(),
//!     stream: Some(true),
//!     ..Default::default()
//! };
//!
//! let stream = request
//!     .get_stream_response(&default_client(), "https://api.deepseek.com", &RequestOptions::bearer(api_key))
//!     .await?;
//!
//! let mut accumulator = ChatCompletionAccumulator::new();
//! let mut stream = stream;
//! while let Some(chunk) = stream.next().await {
//!     accumulator.push(&chunk?);
//! }
//!
//! let message = accumulator.into_message();
//! println!("content: {:?}", message.content);
//! println!("tool calls: {:?}", message.tool_calls);
//! # Ok(())
//! # }
//! ```

use std::collections::BTreeMap;

use super::response::streaming::{
    ChatCompletionChunk, ChoiceDelta, ChoiceDeltaToolCallType, CompletionRole, CompletionUsage,
    FinishReason,
};

/// A deprecated `function_call` (replaced by `tool_calls`) assembled from
/// its streamed fragments.
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AccumulatedFunctionCall {
    /// The name of the function to call, from the fragment that carried it.
    pub name: Option<String>,
    /// The arguments to call the function with, concatenated across all
    /// fragments.
    pub arguments: String,
}

/// Accumulated state of one streamed choice.
#[derive(Debug, Clone, Default)]
struct ChoiceAccumulator {
    role: Option<CompletionRole>,
    content: String,
    #[cfg(feature = "deepseek")]
    reasoning_content: String,
    refusal: String,
    function_call: Option<AccumulatedFunctionCall>,
    tool_calls: BTreeMap<usize, ToolCallAccumulator>,
    #[cfg(feature = "azure")]
    annotations: Vec<crate::chat::Annotation>,
    #[cfg(feature = "azure")]
    audio: Option<crate::chat::ChatCompletionAudio>,
    finish_reason: Option<FinishReason>,
}

/// Accumulated state of one streamed tool call.
#[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 = "deepseek")]
        if let Some(text) = &delta.reasoning_content {
            self.reasoning_content.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);
                }
            }
        }
        #[cfg(feature = "azure")]
        if delta.annotations.is_some() {
            self.annotations = delta.annotations.clone().unwrap_or_default();
        }
        #[cfg(feature = "azure")]
        if let Some(audio) = &delta.audio {
            // The first fragment carries the identity; the byte and
            // transcript payloads arrive incrementally.
            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: crate::chat::ResponseRole::Assistant,
            audio: {
                #[cfg(feature = "azure")]
                let audio = self.audio;
                #[cfg(not(feature = "azure"))]
                let audio = None;
                audio
            },
            content: (!self.content.is_empty()).then_some(self.content),
            #[cfg(feature = "deepseek")]
            reasoning_content: (!self.reasoning_content.is_empty())
                .then_some(self.reasoning_content),
            tool_calls: (!self.tool_calls.is_empty()).then(|| {
                self.tool_calls
                    .into_values()
                    .map(
                        |tool_call| 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: {
                #[cfg(feature = "azure")]
                let annotations = (!self.annotations.is_empty()).then_some(self.annotations);
                #[cfg(not(feature = "azure"))]
                let annotations = None;
                annotations
            },
        }
    }
}

/// Assembles streamed [`ChatCompletionChunk`]s into complete messages.
///
/// All choices are accumulated, keyed by their chunk index; the accessors
/// and [`ChatCompletionAccumulator::into_message`] default to choice `0`
/// (the only choice unless the request set `n > 1`).
///
/// Fragment semantics follow the official SDKs: `content` / `refusal`
/// / tool call `arguments` concatenate; `role`, tool call `id` / `type` /
/// `name` and `function_call.name` are overwritten whenever a fragment
/// carries them; `finish_reason` and `usage` are kept from the last chunk
/// that provided them.
#[derive(Debug, Clone, Default)]
pub struct ChatCompletionAccumulator {
    choices: BTreeMap<u32, ChoiceAccumulator>,
    usage: Option<CompletionUsage>,
}

impl ChatCompletionAccumulator {
    /// Creates an empty accumulator.
    #[must_use]
    pub fn new() -> Self {
        Self::default()
    }

    /// Merges one chunk into the accumulated state.
    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();
            }
        }
    }

    /// The accumulated text content of choice `0`, empty before any content
    /// chunk arrives.
    #[must_use]
    pub fn content(&self) -> &str {
        self.choice().map_or("", |acc| acc.content.as_str())
    }

    /// The accumulated `reasoning_content` of choice `0` (DeepSeek thinking
    /// models).
    #[cfg(feature = "deepseek")]
    #[must_use]
    pub fn reasoning_content(&self) -> &str {
        self.choice()
            .map_or("", |acc| acc.reasoning_content.as_str())
    }

    /// The `finish_reason` of choice `0`, once a chunk carries it.
    #[must_use]
    pub fn finish_reason(&self) -> Option<&FinishReason> {
        self.choice().and_then(|acc| acc.finish_reason.as_ref())
    }

    /// The token usage, once a chunk carries it (the final chunk when the
    /// request set `stream_options: {"include_usage": true}`).
    #[must_use]
    pub fn usage(&self) -> Option<&CompletionUsage> {
        self.usage.as_ref()
    }

    /// The deprecated `function_call` of choice `0`, assembled from its
    /// fragments, if the model used the legacy path.
    #[must_use]
    pub fn function_call(&self) -> Option<&AccumulatedFunctionCall> {
        self.choice().and_then(|acc| acc.function_call.as_ref())
    }

    /// Consumes the accumulator, producing the assembled message of choice
    /// `0`. An accumulator that never saw a choice-`0` chunk yields an
    /// empty assistant message.
    #[must_use]
    pub fn into_message(mut self) -> crate::chat::ChatCompletionMessage {
        self.choices.remove(&0).unwrap_or_default().into_message()
    }

    /// Consumes the accumulator, producing every assembled message keyed by
    /// its choice index.
    #[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 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());
    }
}