Skip to main content

synapse/routing/
stream.rs

1//! Lane-agnostic streaming primitives shared by both lanes and both consumers.
2//! Pure: no I/O, no provider types. Heavily unit-tested.
3
4use serde_json::{json, Value};
5
6/// Why the model stopped. Maps to the OpenAI `finish_reason` string.
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
8pub enum FinishReason {
9    #[default]
10    Stop,
11    ToolCalls,
12    Length,
13}
14
15impl FinishReason {
16    pub fn as_str(self) -> &'static str {
17        match self {
18            FinishReason::Stop => "stop",
19            FinishReason::ToolCalls => "tool_calls",
20            FinishReason::Length => "length",
21        }
22    }
23}
24
25/// A fully-reassembled tool call (buffered output / final message).
26#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct ToolCallOut {
28    pub id: String,
29    pub name: String,
30    /// JSON arguments as a STRING (OpenAI wire shape), reassembled from fragments.
31    pub arguments: String,
32}
33
34/// One lane-agnostic streaming item. Both lanes normalize their upstream
35/// events into this; both consumers read it.
36#[derive(Debug, Clone, PartialEq)]
37pub enum StreamItem {
38    /// A text content delta.
39    Delta(String),
40    /// A tool-call delta. `index` is the stable position within the (possibly
41    /// parallel) tool-call set. `id`/`name` appear on the first delta for an
42    /// index; later deltas carry only `args_fragment`.
43    ToolCallDelta {
44        index: u32,
45        id: Option<String>,
46        name: Option<String>,
47        args_fragment: String,
48    },
49    /// Terminal item: usage + finish reason. Exactly one per successful stream.
50    Done {
51        input_tokens: u64,
52        output_tokens: u64,
53        finish_reason: FinishReason,
54    },
55}
56
57/// Folds a stream of `StreamItem` into final content + tool calls + usage.
58/// Used by the buffered consumer and (for the running totals) by the stream guard.
59#[derive(Debug, Default)]
60pub struct Accumulator {
61    pub content: String,
62    pub tool_calls: Vec<ToolCallOut>,
63    pub input_tokens: u64,
64    pub output_tokens: u64,
65    pub finish_reason: FinishReason,
66    pub got_done: bool,
67}
68
69impl Accumulator {
70    /// Fold one item into the running result.
71    ///
72    /// Assumes `ToolCallDelta.index` values are **dense and 0-based** — which both
73    /// lanes guarantee by construction (the standard lane assigns indices by
74    /// first-seen call id, the native lane by a running counter). A sparse index
75    /// would leave empty `ToolCallOut` slots; we never produce one.
76    pub fn push(&mut self, item: StreamItem) {
77        match item {
78            StreamItem::Delta(t) => self.content.push_str(&t),
79            StreamItem::ToolCallDelta {
80                index,
81                id,
82                name,
83                args_fragment,
84            } => {
85                let i = index as usize;
86                if self.tool_calls.len() <= i {
87                    self.tool_calls.resize(
88                        i + 1,
89                        ToolCallOut {
90                            id: String::new(),
91                            name: String::new(),
92                            arguments: String::new(),
93                        },
94                    );
95                }
96                let slot = &mut self.tool_calls[i];
97                if let Some(id) = id {
98                    slot.id = id;
99                }
100                if let Some(name) = name {
101                    slot.name = name;
102                }
103                slot.arguments.push_str(&args_fragment);
104            }
105            StreamItem::Done {
106                input_tokens,
107                output_tokens,
108                finish_reason,
109            } => {
110                self.input_tokens = input_tokens;
111                self.output_tokens = output_tokens;
112                self.finish_reason = finish_reason;
113                self.got_done = true;
114            }
115        }
116    }
117
118    pub fn has_tool_calls(&self) -> bool {
119        !self.tool_calls.is_empty()
120    }
121
122    /// Serialize the accumulated result as an OpenAI `chat.completion` object.
123    pub fn to_openai_response(&self, id: &str, model: &str) -> Value {
124        let message = if self.has_tool_calls() {
125            json!({
126                "role": "assistant",
127                "content": Value::Null,
128                "tool_calls": self.tool_calls.iter().map(|c| json!({
129                    "id": c.id,
130                    "type": "function",
131                    "function": { "name": c.name, "arguments": c.arguments },
132                })).collect::<Vec<_>>(),
133            })
134        } else {
135            json!({ "role": "assistant", "content": self.content })
136        };
137        json!({
138            "id": format!("chatcmpl-{id}"),
139            "object": "chat.completion",
140            "created": chrono::Utc::now().timestamp(),
141            "model": model,
142            "choices": [{ "index": 0, "message": message, "finish_reason": self.finish_reason.as_str() }],
143            "usage": {
144                "prompt_tokens": self.input_tokens,
145                "completion_tokens": self.output_tokens,
146                "total_tokens": self.input_tokens + self.output_tokens
147            }
148        })
149    }
150}
151
152/// Render one `StreamItem` as an OpenAI `chat.completion.chunk` object.
153/// Tool-call deltas omit `id`/`function.name` when absent (continuation fragments).
154pub fn stream_item_to_sse_json(item: &StreamItem, id: &str, model: &str) -> Value {
155    let base = |delta: Value, finish: Value| {
156        json!({
157            "id": format!("chatcmpl-{id}"),
158            "object": "chat.completion.chunk",
159            "created": 0,
160            "model": model,
161            "choices": [{ "index": 0, "delta": delta, "finish_reason": finish }]
162        })
163    };
164    match item {
165        StreamItem::Delta(t) => base(json!({ "content": t }), Value::Null),
166        StreamItem::ToolCallDelta {
167            index,
168            id: cid,
169            name,
170            args_fragment,
171        } => {
172            let mut func = json!({ "arguments": args_fragment });
173            if let Some(name) = name {
174                func["name"] = json!(name);
175            }
176            let mut call = json!({ "index": index, "type": "function", "function": func });
177            if let Some(cid) = cid {
178                call["id"] = json!(cid);
179            }
180            base(json!({ "tool_calls": [call] }), Value::Null)
181        }
182        StreamItem::Done { finish_reason, .. } => base(json!({}), json!(finish_reason.as_str())),
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189
190    #[test]
191    fn finish_reason_strings() {
192        assert_eq!(FinishReason::Stop.as_str(), "stop");
193        assert_eq!(FinishReason::ToolCalls.as_str(), "tool_calls");
194        assert_eq!(FinishReason::Length.as_str(), "length");
195    }
196
197    #[test]
198    fn accumulates_text_and_usage() {
199        let mut acc = Accumulator::default();
200        acc.push(StreamItem::Delta("Hel".into()));
201        acc.push(StreamItem::Delta("lo".into()));
202        acc.push(StreamItem::Done {
203            input_tokens: 3,
204            output_tokens: 2,
205            finish_reason: FinishReason::Stop,
206        });
207        assert_eq!(acc.content, "Hello");
208        assert!(acc.tool_calls.is_empty());
209        assert_eq!(acc.input_tokens, 3);
210        assert_eq!(acc.output_tokens, 2);
211        assert_eq!(acc.finish_reason, FinishReason::Stop);
212    }
213
214    #[test]
215    fn accumulates_parallel_tool_calls_from_fragments() {
216        let mut acc = Accumulator::default();
217        acc.push(StreamItem::ToolCallDelta {
218            index: 0,
219            id: Some("call_a".into()),
220            name: Some("f".into()),
221            args_fragment: "{\"x\":".into(),
222        });
223        acc.push(StreamItem::ToolCallDelta {
224            index: 1,
225            id: Some("call_b".into()),
226            name: Some("g".into()),
227            args_fragment: "{\"y\":2}".into(),
228        });
229        acc.push(StreamItem::ToolCallDelta {
230            index: 0,
231            id: None,
232            name: None,
233            args_fragment: "1}".into(),
234        });
235        acc.push(StreamItem::Done {
236            input_tokens: 5,
237            output_tokens: 9,
238            finish_reason: FinishReason::ToolCalls,
239        });
240        assert_eq!(acc.tool_calls.len(), 2);
241        assert_eq!(
242            acc.tool_calls[0],
243            ToolCallOut {
244                id: "call_a".into(),
245                name: "f".into(),
246                arguments: "{\"x\":1}".into()
247            }
248        );
249        assert_eq!(
250            acc.tool_calls[1],
251            ToolCallOut {
252                id: "call_b".into(),
253                name: "g".into(),
254                arguments: "{\"y\":2}".into()
255            }
256        );
257        assert_eq!(acc.finish_reason, FinishReason::ToolCalls);
258    }
259
260    #[test]
261    fn buffered_json_text_response() {
262        let mut acc = Accumulator::default();
263        acc.push(StreamItem::Delta("hi".into()));
264        acc.push(StreamItem::Done {
265            input_tokens: 1,
266            output_tokens: 1,
267            finish_reason: FinishReason::Stop,
268        });
269        let v = acc.to_openai_response("abc", "gemini-3-pro");
270        assert_eq!(v["object"], "chat.completion");
271        assert_eq!(v["choices"][0]["message"]["content"], "hi");
272        assert_eq!(v["choices"][0]["finish_reason"], "stop");
273        assert_eq!(v["usage"]["total_tokens"], 2);
274        assert!(v["choices"][0]["message"].get("tool_calls").is_none());
275    }
276
277    #[test]
278    fn buffered_json_tool_call_response() {
279        let mut acc = Accumulator::default();
280        acc.push(StreamItem::ToolCallDelta {
281            index: 0,
282            id: Some("call_0".into()),
283            name: Some("f".into()),
284            args_fragment: "{}".into(),
285        });
286        acc.push(StreamItem::Done {
287            input_tokens: 4,
288            output_tokens: 2,
289            finish_reason: FinishReason::ToolCalls,
290        });
291        let v = acc.to_openai_response("abc", "m");
292        assert_eq!(v["choices"][0]["finish_reason"], "tool_calls");
293        assert!(v["choices"][0]["message"]["content"].is_null());
294        let tc = &v["choices"][0]["message"]["tool_calls"][0];
295        assert_eq!(tc["id"], "call_0");
296        assert_eq!(tc["type"], "function");
297        assert_eq!(tc["function"]["name"], "f");
298        assert_eq!(tc["function"]["arguments"], "{}");
299    }
300
301    #[test]
302    fn sse_text_chunk() {
303        let v = stream_item_to_sse_json(&StreamItem::Delta("hi".into()), "abc", "m");
304        assert_eq!(v["object"], "chat.completion.chunk");
305        assert_eq!(v["choices"][0]["delta"]["content"], "hi");
306        assert!(v["choices"][0]["finish_reason"].is_null());
307    }
308
309    #[test]
310    fn sse_tool_call_chunk() {
311        let item = StreamItem::ToolCallDelta {
312            index: 2,
313            id: Some("call_2".into()),
314            name: Some("f".into()),
315            args_fragment: "{\"a\":1}".into(),
316        };
317        let v = stream_item_to_sse_json(&item, "abc", "m");
318        let tc = &v["choices"][0]["delta"]["tool_calls"][0];
319        assert_eq!(tc["index"], 2);
320        assert_eq!(tc["id"], "call_2");
321        assert_eq!(tc["type"], "function");
322        assert_eq!(tc["function"]["name"], "f");
323        assert_eq!(tc["function"]["arguments"], "{\"a\":1}");
324    }
325
326    #[test]
327    fn sse_done_chunk_sets_finish_reason() {
328        let item = StreamItem::Done {
329            input_tokens: 1,
330            output_tokens: 1,
331            finish_reason: FinishReason::ToolCalls,
332        };
333        let v = stream_item_to_sse_json(&item, "abc", "m");
334        assert_eq!(v["choices"][0]["finish_reason"], "tool_calls");
335        assert_eq!(v["choices"][0]["delta"], serde_json::json!({}));
336    }
337}