1use super::{
4 Content, ContentDelta, ContentStart, GenerateError, GenerateOutput, MessagePart, TokenUsage,
5 ToolUse, ToolUseDelta, ToolUseStart, TurnEndReason,
6};
7
8#[derive(Default)]
9pub struct MessageAccumulator {
10 started: bool,
11 content: Vec<Option<Content>>,
12 tool_uses: Vec<Option<IncompleteToolUse>>,
13 turn_end_reason: Option<TurnEndReason>,
14 usage: Option<TokenUsage>,
15}
16
17impl MessageAccumulator {
18 pub fn push(&mut self, part: &MessagePart) -> Result<(), GenerateError> {
19 if !self.started {
20 return self.start(part);
21 }
22 match part {
23 MessagePart::MessageStart => return Err(malformed("unexpected MessageStart")),
24 MessagePart::ContentStart(start) => {
25 let (index, block) = start_block(start.clone());
26 ensure_slot(&mut self.content, index, block);
27 }
28 MessagePart::ContentDelta(delta) => apply_content_delta(&mut self.content, delta)?,
29 MessagePart::ToolUseStart(ToolUseStart { index, name }) => {
30 ensure_slot(
31 &mut self.tool_uses,
32 *index,
33 IncompleteToolUse {
34 name: name.clone(),
35 input_json: String::new(),
36 },
37 );
38 }
39 MessagePart::ToolUseDelta(delta) => self.append_tool_input(delta)?,
40 MessagePart::TurnEndReason(reason) => self.turn_end_reason = Some(reason.clone()),
41 MessagePart::Usage(usage) => {
42 self.usage = Some(self.usage.map_or(*usage, |prev| prev.merge(*usage)))
43 }
44 }
45 Ok(())
46 }
47
48 fn start(&mut self, part: &MessagePart) -> Result<(), GenerateError> {
49 if !matches!(part, MessagePart::MessageStart) {
50 return Err(malformed(
51 "first item is not MessageStart. Did you accidentally drain the stream already?",
52 ));
53 }
54 self.started = true;
55 Ok(())
56 }
57
58 fn append_tool_input(&mut self, delta: &ToolUseDelta) -> Result<(), GenerateError> {
59 let index = delta.index;
60 let tool = self
61 .tool_uses
62 .get_mut(index)
63 .and_then(Option::as_mut)
64 .ok_or_else(|| malformed(format!("tool use delta index {index} is out of bounds")))?;
65 tool.input_json.push_str(&delta.input_json_delta);
66 Ok(())
67 }
68
69 pub fn finish(self) -> Result<GenerateOutput, GenerateError> {
70 if !self.started {
71 return Err(malformed(
72 "empty stream. Did you accidentally drain the stream already?",
73 ));
74 }
75 let content = filled_slots(self.content, "content block")?;
76 let tool_uses = filled_slots(self.tool_uses, "tool use")?
77 .into_iter()
78 .map(IncompleteToolUse::finish)
79 .collect::<Result<_, _>>()?;
80 let turn_end_reason = self
81 .turn_end_reason
82 .ok_or_else(|| malformed("no turn end reason provided"))?;
83 Ok(GenerateOutput {
84 content,
85 tool_uses,
86 turn_end_reason,
87 usage: self.usage,
88 })
89 }
90}
91
92struct IncompleteToolUse {
93 name: String,
94 input_json: String,
95}
96
97impl IncompleteToolUse {
98 fn finish(self) -> Result<ToolUse, GenerateError> {
99 let json = if self.input_json.is_empty() {
100 "{}"
101 } else {
102 &self.input_json
103 };
104 let input = serde_json::from_str(json)
105 .map_err(|error| malformed(format!("tool use input JSON is invalid: {error}")))?;
106 Ok(ToolUse {
107 name: self.name,
108 input,
109 })
110 }
111}
112
113fn malformed(message: impl std::fmt::Display) -> GenerateError {
114 GenerateError::MalformedResponseError(format!("Malformed stream: {message}"))
115}
116
117fn filled_slots<T>(slots: Vec<Option<T>>, name: &str) -> Result<Vec<T>, GenerateError> {
118 slots
119 .into_iter()
120 .enumerate()
121 .map(|(index, slot)| {
122 slot.ok_or_else(|| malformed(format!("missing {name} at index {index}")))
123 })
124 .collect()
125}
126
127fn ensure_slot<T>(slots: &mut Vec<Option<T>>, index: usize, value: T) {
128 while slots.len() <= index {
129 slots.push(None);
130 }
131 slots[index] = Some(value);
132}
133
134fn start_block(start: ContentStart) -> (usize, Content) {
136 match start {
137 ContentStart::Text { index } => (
138 index,
139 Content::Text {
140 text: String::new(),
141 },
142 ),
143 ContentStart::Image { index } => (
144 index,
145 Content::Image {
146 source: String::new(),
147 },
148 ),
149 ContentStart::Thinking {
150 index,
151 signature,
152 redacted,
153 } => (
154 index,
155 Content::Thinking {
156 text: String::new(),
157 signature,
158 redacted,
159 },
160 ),
161 }
162}
163
164fn apply_content_delta(
167 content: &mut [Option<Content>],
168 delta: &ContentDelta,
169) -> Result<(), GenerateError> {
170 let index = match delta {
171 ContentDelta::Text { index, .. }
172 | ContentDelta::Image { index, .. }
173 | ContentDelta::Thinking { index, .. } => *index,
174 };
175 match (content.get_mut(index).and_then(Option::as_mut), delta) {
176 (Some(Content::Text { text }), ContentDelta::Text { delta, .. }) => text.push_str(delta),
177 (Some(Content::Image { source }), ContentDelta::Image { delta, .. }) => {
178 source.push_str(delta);
179 }
180 (Some(Content::Thinking { text, redacted, .. }), ContentDelta::Thinking { delta, .. }) => {
181 if !*redacted {
182 text.push_str(delta);
183 }
184 }
185 _ => {
186 return Err(GenerateError::MalformedResponseError(format!(
187 "Malformed stream: content delta at index {index}: out of bounds or wrong kind"
188 )));
189 }
190 }
191 Ok(())
192}
193
194#[cfg(test)]
195mod tests {
196 use super::*;
197
198 fn assemble(parts: &[MessagePart]) -> Result<GenerateOutput, GenerateError> {
199 let mut accumulator = MessageAccumulator::default();
200 for part in parts {
201 accumulator.push(part)?;
202 }
203 accumulator.finish()
204 }
205
206 #[test]
207 fn tool_input_fragments_form_one_json_value() {
208 let output = assemble(&[
209 MessagePart::MessageStart,
210 MessagePart::ToolUseStart(ToolUseStart {
211 index: 0,
212 name: "bash".into(),
213 }),
214 MessagePart::ToolUseDelta(ToolUseDelta {
215 index: 0,
216 input_json_delta: "{\"command\":".into(),
217 }),
218 MessagePart::ToolUseDelta(ToolUseDelta {
219 index: 0,
220 input_json_delta: "\"pwd\"}".into(),
221 }),
222 MessagePart::TurnEndReason(TurnEndReason::ToolUse),
223 ])
224 .unwrap();
225 assert_eq!(output.tool_uses.len(), 1);
226 assert_eq!(output.tool_uses[0].name, "bash");
227 assert_eq!(
228 output.tool_uses[0].input,
229 serde_json::json!({"command": "pwd"})
230 );
231 }
232
233 #[test]
234 fn incomplete_and_malformed_messages_are_rejected() {
235 let cases = [
236 vec![],
237 vec![MessagePart::TurnEndReason(TurnEndReason::EndTurn)],
238 vec![MessagePart::MessageStart],
239 vec![MessagePart::MessageStart, MessagePart::MessageStart],
240 vec![
241 MessagePart::MessageStart,
242 MessagePart::ContentDelta(ContentDelta::Text {
243 index: 0,
244 delta: "unopened".into(),
245 }),
246 ],
247 vec![
248 MessagePart::MessageStart,
249 MessagePart::ContentStart(ContentStart::Text { index: 1 }),
250 MessagePart::TurnEndReason(TurnEndReason::EndTurn),
251 ],
252 vec![
253 MessagePart::MessageStart,
254 MessagePart::ToolUseDelta(ToolUseDelta {
255 index: 0,
256 input_json_delta: "{}".into(),
257 }),
258 ],
259 vec![
260 MessagePart::MessageStart,
261 MessagePart::ToolUseStart(ToolUseStart {
262 index: 0,
263 name: "bash".into(),
264 }),
265 MessagePart::ToolUseDelta(ToolUseDelta {
266 index: 0,
267 input_json_delta: "{".into(),
268 }),
269 MessagePart::TurnEndReason(TurnEndReason::ToolUse),
270 ],
271 ];
272 for parts in cases {
273 assert!(
274 matches!(
275 assemble(&parts),
276 Err(GenerateError::MalformedResponseError(_))
277 ),
278 "{parts:?}"
279 );
280 }
281 }
282}