1use serde::de::{self, Deserializer, MapAccess, Visitor};
7use serde::{Deserialize, Serialize};
8use serde_json::Value;
9use std::fmt;
10
11#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
17pub enum StreamEventType {
18 #[serde(rename = "response.created")]
20 ResponseCreated,
21 #[serde(rename = "response.in_progress")]
22 ResponseInProgress,
23 #[serde(rename = "response.completed")]
24 ResponseCompleted,
25 #[serde(rename = "response.failed")]
26 ResponseFailed,
27 #[serde(rename = "response.incomplete")]
28 ResponseIncomplete,
29
30 #[serde(rename = "response.output_item.added")]
32 OutputItemAdded,
33 #[serde(rename = "response.output_item.done")]
34 OutputItemDone,
35
36 #[serde(rename = "response.output_text.delta")]
38 OutputTextDelta,
39 #[serde(rename = "response.output_text.done")]
40 OutputTextDone,
41
42 #[serde(rename = "response.content_part.added")]
44 ContentPartAdded,
45 #[serde(rename = "response.content_part.done")]
46 ContentPartDone,
47
48 #[serde(rename = "response.function_call_arguments.delta")]
50 FunctionCallArgumentsDelta,
51 #[serde(rename = "response.function_call_arguments.done")]
52 FunctionCallArgumentsDone,
53
54 #[serde(rename = "response.reasoning_summary_text.delta")]
56 ReasoningSummaryTextDelta,
57 #[serde(rename = "response.reasoning_summary_text.done")]
58 ReasoningSummaryTextDone,
59
60 #[serde(rename = "response.reasoning_content.delta")]
62 ReasoningContentDelta,
63 #[serde(rename = "response.reasoning_content.done")]
64 ReasoningContentDone,
65
66 #[serde(rename = "error")]
68 Error,
69 #[serde(other)]
71 Unknown,
72}
73
74#[derive(Debug, Clone, Serialize)]
76pub struct StreamEvent {
77 #[serde(rename = "type")]
78 event_type: String,
79 #[serde(default)]
80 sequence_number: u32,
81 #[serde(flatten)]
82 data: StreamEventData,
83}
84
85impl<'de> Deserialize<'de> for StreamEvent {
86 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
87 where
88 D: Deserializer<'de>,
89 {
90 let wire = StreamEventWire::deserialize(deserializer)?;
91 let (event_type, sequence_number, data) = wire.into_event().map_err(de::Error::custom)?;
92 Ok(Self { event_type, sequence_number, data })
93 }
94}
95
96#[derive(Debug, Default)]
104struct StreamEventWire {
105 event_type: Option<String>,
106 sequence_number: u32,
107 response: Option<Value>,
108 item: Option<Value>,
109 output_index: Option<u32>,
110 item_id: Option<String>,
111 content_index: Option<u32>,
112 call_id: Option<String>,
113 delta: Option<String>,
114 error: Option<StreamError>,
115 extra: Option<serde_json::Map<String, Value>>,
116}
117
118impl StreamEventWire {
119 fn into_event(mut self) -> Result<(String, u32, StreamEventData), String> {
120 let event_type = required_field(self.event_type.take(), "type", "stream event")?;
124 let sequence_number = self.sequence_number;
125 let data = self.into_data(&event_type)?;
126 Ok((event_type, sequence_number, data))
127 }
128
129 fn into_data(self, event_type: &str) -> Result<StreamEventData, String> {
130 match event_type {
131 "response.created"
132 | "response.in_progress"
133 | "response.completed"
134 | "response.failed"
135 | "response.incomplete" => Ok(StreamEventData::Response(ResponseEventData { response: self.response })),
136 "response.output_item.added" | "response.output_item.done" => {
137 Ok(StreamEventData::OutputItem(OutputItemEventData {
138 item: self.item,
139 output_index: self.output_index,
140 item_id: self.item_id,
141 }))
142 }
143 "response.output_text.delta"
144 | "response.output_text.done"
145 | "response.reasoning_summary_text.delta"
146 | "response.reasoning_summary_text.done" => Ok(StreamEventData::TextDelta(TextDeltaEventData {
147 delta: required_field(self.delta, "delta", event_type)?,
148 item_id: self.item_id,
149 output_index: self.output_index,
150 content_index: self.content_index,
151 })),
152 "response.function_call_arguments.delta" | "response.function_call_arguments.done" => {
153 Ok(StreamEventData::FunctionCallDelta(FunctionCallDeltaEventData {
154 delta: required_field(self.delta, "delta", event_type)?,
155 item_id: self.item_id,
156 output_index: self.output_index,
157 call_id: self.call_id,
158 }))
159 }
160 "response.reasoning_content.delta" | "response.reasoning_content.done" => {
161 Ok(StreamEventData::ReasoningContentDelta(ReasoningContentDeltaEventData {
162 delta: required_field(self.delta, "delta", event_type)?,
163 item_id: self.item_id,
164 output_index: self.output_index,
165 }))
166 }
167 "error" => Ok(StreamEventData::Error(ErrorEventData {
168 error: required_field(self.error, "error", event_type)?,
169 })),
170 _ => Ok(StreamEventData::Generic(self.into_generic_value())),
171 }
172 }
173
174 fn into_generic_value(self) -> Value {
175 let Self {
176 response,
177 item,
178 output_index,
179 item_id,
180 content_index,
181 call_id,
182 delta,
183 error,
184 extra,
185 ..
186 } = self;
187 let mut object = extra.unwrap_or_default();
188 insert_optional(&mut object, "response", response);
189 insert_optional(&mut object, "item", item);
190 insert_optional(&mut object, "output_index", output_index);
191 insert_optional(&mut object, "item_id", item_id);
192 insert_optional(&mut object, "content_index", content_index);
193 insert_optional(&mut object, "call_id", call_id);
194 insert_optional(&mut object, "delta", delta);
195 if let Some(error) = error {
196 if let Ok(value) = serde_json::to_value(error) {
197 object.insert("error".to_string(), value);
198 }
199 }
200 Value::Object(object)
201 }
202}
203
204impl<'de> Deserialize<'de> for StreamEventWire {
205 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
206 where
207 D: Deserializer<'de>,
208 {
209 struct StreamEventWireVisitor;
210
211 impl<'de> Visitor<'de> for StreamEventWireVisitor {
212 type Value = StreamEventWire;
213
214 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
215 formatter.write_str("an OpenResponses streaming event object")
216 }
217
218 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
219 where
220 A: MapAccess<'de>,
221 {
222 let mut wire = StreamEventWire::default();
223 while let Some(key) = map.next_key::<&str>()? {
224 match key {
225 "type" => wire.event_type = map.next_value()?,
226 "sequence_number" => wire.sequence_number = map.next_value()?,
227 "response" => wire.response = map.next_value()?,
228 "item" => wire.item = map.next_value()?,
229 "output_index" => wire.output_index = map.next_value()?,
230 "item_id" => wire.item_id = map.next_value()?,
231 "content_index" => wire.content_index = map.next_value()?,
232 "call_id" => wire.call_id = map.next_value()?,
233 "delta" => wire.delta = map.next_value()?,
234 "error" => wire.error = map.next_value()?,
235 _ => {
236 wire.extra
237 .get_or_insert_with(serde_json::Map::new)
238 .insert(key.to_string(), map.next_value()?);
239 }
240 }
241 }
242 if wire.event_type.is_none() {
243 return Err(de::Error::missing_field("type"));
244 }
245 Ok(wire)
246 }
247 }
248
249 deserializer.deserialize_map(StreamEventWireVisitor)
250 }
251}
252
253fn required_field<T>(value: Option<T>, field: &str, event_type: &str) -> Result<T, String> {
254 value.ok_or_else(|| format!("OpenResponses event {event_type:?} is missing {field:?}"))
255}
256
257fn insert_optional<T: Serialize>(object: &mut serde_json::Map<String, Value>, key: &str, value: Option<T>) {
258 if let Some(value) = value
259 && let Ok(value) = serde_json::to_value(value)
260 {
261 object.insert(key.to_string(), value);
262 }
263}
264
265#[derive(Debug, Clone, Serialize, Deserialize)]
267#[serde(untagged)]
268pub enum StreamEventData {
269 Response(ResponseEventData),
271 OutputItem(OutputItemEventData),
273 TextDelta(TextDeltaEventData),
275 FunctionCallDelta(FunctionCallDeltaEventData),
277 ReasoningContentDelta(ReasoningContentDeltaEventData),
279 Error(ErrorEventData),
281 Generic(Value),
283}
284
285#[derive(Debug, Clone, Serialize, Deserialize)]
287pub struct ResponseEventData {
288 #[serde(skip_serializing_if = "Option::is_none")]
289 response: Option<Value>,
290}
291
292#[derive(Debug, Clone, Serialize, Deserialize)]
294pub struct OutputItemEventData {
295 #[serde(skip_serializing_if = "Option::is_none")]
296 item: Option<Value>,
297 #[serde(skip_serializing_if = "Option::is_none")]
298 output_index: Option<u32>,
299 #[serde(skip_serializing_if = "Option::is_none")]
300 item_id: Option<String>,
301}
302
303#[derive(Debug, Clone, Serialize, Deserialize)]
305pub struct TextDeltaEventData {
306 delta: String,
307 #[serde(skip_serializing_if = "Option::is_none")]
308 item_id: Option<String>,
309 #[serde(skip_serializing_if = "Option::is_none")]
310 output_index: Option<u32>,
311 #[serde(skip_serializing_if = "Option::is_none")]
312 content_index: Option<u32>,
313}
314
315#[derive(Debug, Clone, Serialize, Deserialize)]
317pub struct FunctionCallDeltaEventData {
318 delta: String,
319 #[serde(skip_serializing_if = "Option::is_none")]
320 item_id: Option<String>,
321 #[serde(skip_serializing_if = "Option::is_none")]
322 output_index: Option<u32>,
323 #[serde(skip_serializing_if = "Option::is_none")]
324 call_id: Option<String>,
325}
326
327#[derive(Debug, Clone, Serialize, Deserialize)]
329pub struct ReasoningContentDeltaEventData {
330 delta: String,
331 #[serde(skip_serializing_if = "Option::is_none")]
332 item_id: Option<String>,
333 #[serde(skip_serializing_if = "Option::is_none")]
334 output_index: Option<u32>,
335}
336
337#[derive(Debug, Clone, Serialize, Deserialize)]
339pub struct ErrorEventData {
340 error: StreamError,
341}
342
343#[derive(Debug, Clone, Serialize, Deserialize)]
345pub struct StreamError {
346 code: String,
347 message: String,
348 #[serde(skip_serializing_if = "Option::is_none")]
349 param: Option<String>,
350}
351
352fn parse_sse_event(line: &str) -> Option<StreamEvent> {
358 let line = line.trim();
360 if line.is_empty() || line == "[DONE]" {
361 return None;
362 }
363
364 if let Some(data) = line.strip_prefix("data: ") {
365 if data == "[DONE]" {
366 return None;
367 }
368 serde_json::from_str(data).ok()
369 } else if line.starts_with('{') {
370 serde_json::from_str(line).ok()
372 } else {
373 None
374 }
375}
376
377pub fn extract_event_type(line: &str) -> Option<String> {
379 let line = line.trim();
380 line.strip_prefix("event: ").map(|event_type| event_type.to_string())
381}
382
383#[derive(Debug, Default)]
385pub struct StreamAccumulator {
386 text_content: String,
387 reasoning_content: String,
388 reasoning_summary: String,
389 function_calls: Vec<AccumulatedFunctionCall>,
390 current_function_call: Option<AccumulatingFunctionCall>,
391 output_items: Vec<Value>,
392 response_id: Option<String>,
393 model: Option<String>,
394 usage: Option<Value>,
395 is_complete: bool,
396 error: Option<StreamError>,
397}
398
399#[derive(Debug, Clone, Default)]
401pub struct AccumulatingFunctionCall {
402 id: String,
403 call_id: String,
404 name: String,
405 arguments: String,
406}
407
408#[derive(Debug, Clone)]
410pub struct AccumulatedFunctionCall {
411 id: String,
412 call_id: String,
413 name: String,
414 arguments: String,
415}
416
417impl StreamAccumulator {
418 fn new() -> Self {
419 Self::default()
420 }
421
422 fn process_event(&mut self, event: &StreamEvent) {
424 match event.event_type.as_str() {
425 "response.created" | "response.in_progress" => {
426 if let StreamEventData::Response(data) = &event.data
427 && let Some(response) = &data.response
428 {
429 self.response_id = response.get("id").and_then(|v| v.as_str()).map(String::from);
430 self.model = response.get("model").and_then(|v| v.as_str()).map(String::from);
431 }
432 }
433 "response.output_text.delta" => {
434 if let StreamEventData::TextDelta(data) = &event.data {
435 self.text_content.push_str(&data.delta);
436 }
437 }
438 "response.reasoning_summary_text.delta" => {
439 if let StreamEventData::TextDelta(data) = &event.data {
441 self.reasoning_summary.push_str(&data.delta);
442 }
443 }
444 "response.reasoning_content.delta" => {
445 if let StreamEventData::ReasoningContentDelta(data) = &event.data {
447 self.reasoning_content.push_str(&data.delta);
448 }
449 }
450 "response.function_call_arguments.delta" => {
451 if let StreamEventData::FunctionCallDelta(data) = &event.data
452 && let Some(ref mut fc) = self.current_function_call
453 {
454 fc.arguments.push_str(&data.delta);
455 }
456 }
457 "response.output_item.added" => {
458 if let StreamEventData::OutputItem(data) = &event.data
459 && let Some(item) = &data.item
460 {
461 if item.get("type").and_then(|v| v.as_str()) == Some("function_call") {
463 let fc = AccumulatingFunctionCall {
464 id: item.get("id").and_then(|v| v.as_str()).unwrap_or_default().to_string(),
465 call_id: item.get("call_id").and_then(|v| v.as_str()).unwrap_or_default().to_string(),
466 name: item.get("name").and_then(|v| v.as_str()).unwrap_or_default().to_string(),
467 arguments: String::new(),
468 };
469 self.current_function_call = Some(fc);
470 }
471 self.output_items.push(item.clone());
472 }
473 }
474 "response.output_item.done" => {
475 if let Some(fc) = self.current_function_call.take() {
477 self.function_calls.push(AccumulatedFunctionCall {
478 id: fc.id,
479 call_id: fc.call_id,
480 name: fc.name,
481 arguments: fc.arguments,
482 });
483 }
484 }
485 "response.completed" => {
486 self.is_complete = true;
487 if let StreamEventData::Response(data) = &event.data
488 && let Some(response) = &data.response
489 {
490 self.usage = response.get("usage").cloned();
491 }
492 }
493 "response.failed" => {
494 self.is_complete = true;
495 }
496 "error" => {
497 if let StreamEventData::Error(data) = &event.data {
498 self.error = Some(data.error.clone());
499 }
500 self.is_complete = true;
501 }
502 _ => {}
503 }
504 }
505}
506
507#[cfg(test)]
508mod tests {
509 use super::*;
510
511 #[test]
512 fn test_parse_sse_text_delta() {
513 let line = r#"data: {"type":"response.output_text.delta","sequence_number":1,"delta":"Hello"}"#;
514 let event = parse_sse_event(line).unwrap();
515 assert_eq!(event.event_type, "response.output_text.delta");
516 assert!(matches!(
517 event.data,
518 StreamEventData::TextDelta(TextDeltaEventData { delta, .. }) if delta == "Hello"
519 ));
520 }
521
522 #[test]
523 fn test_parse_sse_dispatches_payload_by_event_type() {
524 let cases = [
525 (r#"data: {"type":"response.created","response":{"id":"resp_1"}}"#, "response"),
526 (r#"data: {"type":"response.output_item.added","item":{"type":"message"}}"#, "output_item"),
527 (
528 r#"data: {"type":"response.function_call_arguments.delta","delta":"{}","call_id":"call_1"}"#,
529 "function_call",
530 ),
531 (r#"data: {"type":"response.reasoning_content.delta","delta":"think"}"#, "reasoning"),
532 (r#"data: {"type":"error","error":{"code":"bad_request","message":"nope"}}"#, "error"),
533 ];
534
535 for (line, expected) in cases {
536 let event = parse_sse_event(line).expect("valid streaming event");
537 let actual = match event.data {
538 StreamEventData::Response(_) => "response",
539 StreamEventData::OutputItem(_) => "output_item",
540 StreamEventData::FunctionCallDelta(_) => "function_call",
541 StreamEventData::ReasoningContentDelta(_) => "reasoning",
542 StreamEventData::Error(_) => "error",
543 _ => "other",
544 };
545 assert_eq!(actual, expected, "event line: {line}");
546 }
547 }
548
549 #[test]
550 fn test_parse_sse_preserves_unknown_payload_fields() {
551 let line = r#"data: {"type":"response.future_event","sequence_number":7,"custom":{"value":true}}"#;
552 let event = parse_sse_event(line).expect("valid unknown streaming event");
553 assert_eq!(event.event_type, "response.future_event");
554 assert_eq!(event.sequence_number, 7);
555 assert!(matches!(
556 event.data,
557 StreamEventData::Generic(Value::Object(ref object))
558 if object.get("custom") == Some(&serde_json::json!({"value": true}))
559 ));
560 }
561
562 #[test]
563 fn test_parse_done_signal() {
564 assert!(parse_sse_event("[DONE]").is_none());
565 assert!(parse_sse_event("data: [DONE]").is_none());
566 }
567
568 #[test]
569 fn test_stream_accumulator_text() {
570 let mut acc = StreamAccumulator::new();
571
572 let event1 = StreamEvent {
573 event_type: "response.output_text.delta".to_string(),
574 sequence_number: 1,
575 data: StreamEventData::TextDelta(TextDeltaEventData {
576 delta: "Hello, ".to_string(),
577 item_id: None,
578 output_index: None,
579 content_index: None,
580 }),
581 };
582
583 let event2 = StreamEvent {
584 event_type: "response.output_text.delta".to_string(),
585 sequence_number: 2,
586 data: StreamEventData::TextDelta(TextDeltaEventData {
587 delta: "world!".to_string(),
588 item_id: None,
589 output_index: None,
590 content_index: None,
591 }),
592 };
593
594 acc.process_event(&event1);
595 acc.process_event(&event2);
596
597 assert_eq!(acc.text_content, "Hello, world!");
598 }
599}