1use serde_json::{json, Value};
5
6#[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#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct ToolCallOut {
28 pub id: String,
29 pub name: String,
30 pub arguments: String,
32}
33
34#[derive(Debug, Clone, PartialEq)]
37pub enum StreamItem {
38 Delta(String),
40 ToolCallDelta {
44 index: u32,
45 id: Option<String>,
46 name: Option<String>,
47 args_fragment: String,
48 },
49 Done {
51 input_tokens: u64,
52 output_tokens: u64,
53 finish_reason: FinishReason,
54 },
55}
56
57#[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 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 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
152pub 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}