Skip to main content

myco_model/
accumulator.rs

1//! Assemble validated message parts without owning a stream, retry policy, or event sink.
2
3use 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
134/// The empty [`Content`] block a [`ContentStart`] opens, with its index.
135fn 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
164/// Append a [`ContentDelta`] to its opened block; the slot must exist and be
165/// the matching kind (redacted thinking swallows its deltas).
166fn 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}