Skip to main content

behest_runtime/
accumulator.rs

1//! Streaming accumulator for text and tool calls.
2//!
3//! Maintains state for accumulating streaming deltas from provider
4//! into complete assistant messages and tool calls.
5
6use std::collections::HashMap;
7
8use behest_provider::{ContentPart, Message, ToolCall};
9
10/// Accumulates streaming deltas from a provider into a complete assistant
11/// message and zero or more tool calls. Tracks both text content and
12/// partial tool call arguments until all deltas are received.
13#[derive(Debug, Default)]
14pub struct StreamAccumulator {
15    text: String,
16    tool_calls: HashMap<String, ToolCallAccumulator>,
17}
18
19impl StreamAccumulator {
20    /// Creates a new accumulator.
21    #[must_use]
22    pub fn new() -> Self {
23        Self::default()
24    }
25
26    /// Appends a text delta chunk to the internal buffer.
27    pub fn append_text(&mut self, delta: &str) {
28        self.text.push_str(delta);
29    }
30
31    /// Starts a tool call.
32    pub fn start_tool_call(&mut self, id: String, name: String) {
33        self.tool_calls.insert(
34            id.clone(),
35            ToolCallAccumulator {
36                id,
37                name,
38                arguments: String::new(),
39            },
40        );
41    }
42
43    /// Appends arguments to a tool call.
44    pub fn append_tool_arguments(&mut self, id: &str, delta: &str) {
45        if let Some(tc) = self.tool_calls.get_mut(id) {
46            tc.arguments.push_str(delta);
47        }
48    }
49
50    /// Returns the accumulated text.
51    #[must_use]
52    pub fn text(&self) -> &str {
53        &self.text
54    }
55
56    /// Parses accumulated tool call arguments into [`ToolCall`] values.
57    /// Unparseable JSON arguments produce a [`serde_json::Value::Null`] fallback.
58    #[must_use]
59    pub fn tool_calls(&self) -> Vec<ToolCall> {
60        self.tool_calls
61            .values()
62            .map(|tc| {
63                let arguments =
64                    serde_json::from_str(&tc.arguments).unwrap_or(serde_json::Value::Null);
65                ToolCall::new(tc.id.clone(), tc.name.clone(), arguments)
66            })
67            .collect()
68    }
69
70    /// Assembles an assistant [`Message`] from the accumulated text and
71    /// tool calls. Handles three cases:
72    /// - Neither text nor tool calls → empty assistant message.
73    /// - Text only → [`Message::assistant_text`].
74    /// - Tool calls present (with optional text) → structured [`Message::Assistant`].
75    #[must_use]
76    pub fn to_message(&self) -> Message {
77        let tool_calls = self.tool_calls();
78        if tool_calls.is_empty() && self.text.is_empty() {
79            Message::Assistant {
80                content: vec![],
81                tool_calls: vec![],
82            }
83        } else if tool_calls.is_empty() {
84            Message::assistant_text(&self.text)
85        } else {
86            Message::Assistant {
87                content: if self.text.is_empty() {
88                    vec![]
89                } else {
90                    vec![ContentPart::text(&self.text)]
91                },
92                tool_calls,
93            }
94        }
95    }
96
97    /// Clears the accumulator.
98    pub fn clear(&mut self) {
99        self.text.clear();
100        self.tool_calls.clear();
101    }
102}
103
104/// Accumulator for a single tool call.
105#[derive(Debug)]
106struct ToolCallAccumulator {
107    id: String,
108    name: String,
109    arguments: String,
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115    use serde_json::Value;
116
117    #[test]
118    fn accumulate_text() {
119        let mut acc = StreamAccumulator::new();
120        acc.append_text("Hello");
121        acc.append_text(" ");
122        acc.append_text("World");
123        assert_eq!(acc.text(), "Hello World");
124    }
125
126    #[test]
127    fn accumulate_tool_call() {
128        let mut acc = StreamAccumulator::new();
129        acc.start_tool_call("call_1".to_string(), "get_weather".to_string());
130        acc.append_tool_arguments("call_1", r#"{"location":"#);
131        acc.append_tool_arguments("call_1", r#""Paris"}"#);
132
133        let calls = acc.tool_calls();
134        assert_eq!(calls.len(), 1);
135        assert_eq!(calls[0].id, "call_1");
136        assert_eq!(calls[0].name, "get_weather");
137        assert_eq!(calls[0].arguments["location"], "Paris");
138    }
139
140    #[test]
141    fn to_message_text_only() {
142        let mut acc = StreamAccumulator::new();
143        acc.append_text("Response");
144        let msg = acc.to_message();
145        match msg {
146            Message::Assistant {
147                content,
148                tool_calls,
149            } => {
150                assert!(tool_calls.is_empty());
151                assert!(!content.is_empty());
152            }
153            _ => panic!("Expected Assistant message"),
154        }
155    }
156
157    #[test]
158    fn to_message_with_tool_calls() {
159        let mut acc = StreamAccumulator::new();
160        acc.append_text("Thinking...");
161        acc.start_tool_call("call_1".to_string(), "tool".to_string());
162        acc.append_tool_arguments("call_1", "{}");
163
164        let msg = acc.to_message();
165        match msg {
166            Message::Assistant {
167                content,
168                tool_calls,
169            } => {
170                assert_eq!(tool_calls.len(), 1);
171                assert!(!content.is_empty());
172            }
173            _ => panic!("Expected Assistant message with tool calls"),
174        }
175    }
176
177    #[test]
178    fn to_message_empty_accumulator_returns_empty_assistant() {
179        let acc = StreamAccumulator::new();
180
181        let msg = acc.to_message();
182        match msg {
183            Message::Assistant {
184                content,
185                tool_calls,
186            } => {
187                assert!(content.is_empty());
188                assert!(tool_calls.is_empty());
189            }
190            _ => panic!("Expected empty Assistant message"),
191        }
192    }
193
194    #[test]
195    fn invalid_tool_arguments_fall_back_to_null() {
196        let mut acc = StreamAccumulator::new();
197        acc.start_tool_call("call_1".to_string(), "tool".to_string());
198        acc.append_tool_arguments("call_1", "{invalid json");
199
200        let calls = acc.tool_calls();
201        assert_eq!(calls.len(), 1);
202        assert_eq!(calls[0].arguments, Value::Null);
203    }
204
205    #[test]
206    fn clear_resets_text_and_tool_calls() {
207        let mut acc = StreamAccumulator::new();
208        acc.append_text("partial");
209        acc.start_tool_call("call_1".to_string(), "tool".to_string());
210        acc.append_tool_arguments("call_1", "{}");
211
212        acc.clear();
213
214        assert_eq!(acc.text(), "");
215        assert!(acc.tool_calls().is_empty());
216        match acc.to_message() {
217            Message::Assistant {
218                content,
219                tool_calls,
220            } => {
221                assert!(content.is_empty());
222                assert!(tool_calls.is_empty());
223            }
224            _ => panic!("Expected empty Assistant message after clear"),
225        }
226    }
227
228    #[test]
229    fn starting_same_tool_call_id_replaces_previous_state() {
230        let mut acc = StreamAccumulator::new();
231        acc.start_tool_call("call_1".to_string(), "first_tool".to_string());
232        acc.append_tool_arguments("call_1", r#"{"old":"value"}"#);
233
234        acc.start_tool_call("call_1".to_string(), "second_tool".to_string());
235        acc.append_tool_arguments("call_1", r#"{"fresh":true}"#);
236
237        let calls = acc.tool_calls();
238        assert_eq!(calls.len(), 1);
239        assert_eq!(calls[0].name, "second_tool");
240        assert_eq!(calls[0].arguments["fresh"], true);
241        assert_eq!(calls[0].arguments.get("old"), None);
242    }
243}