behest_runtime/
accumulator.rs1use std::collections::HashMap;
7
8use behest_provider::{ContentPart, Message, ToolCall};
9
10#[derive(Debug, Default)]
14pub struct StreamAccumulator {
15 text: String,
16 tool_calls: HashMap<String, ToolCallAccumulator>,
17}
18
19impl StreamAccumulator {
20 #[must_use]
22 pub fn new() -> Self {
23 Self::default()
24 }
25
26 pub fn append_text(&mut self, delta: &str) {
28 self.text.push_str(delta);
29 }
30
31 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 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 #[must_use]
52 pub fn text(&self) -> &str {
53 &self.text
54 }
55
56 #[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 #[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 pub fn clear(&mut self) {
99 self.text.clear();
100 self.tool_calls.clear();
101 }
102}
103
104#[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}