1use crate::events::{EventFrame, EventPayload, SSEEventType, WireEvent, normalize_sse_line};
2use crate::executor::error::{ExecutorError, ExecutorResult};
3use crate::types::request_response::ResponsePayload;
4use crate::utils::common::{serialize_to_string, serialize_to_value};
5use serde_json::Value;
6
7#[derive(Clone)]
8pub struct GatewayStreamAccumulator {
9 next_sequence_number: u64,
10 emitted_created: bool,
11 emitted_in_progress: bool,
12}
13
14pub(super) struct StreamEvent {
15 pub(super) content: String,
16 pub(super) sequence_number: u64,
17}
18
19impl GatewayStreamAccumulator {
20 #[must_use]
21 pub fn new() -> Self {
22 Self {
23 next_sequence_number: 0,
24 emitted_created: false,
25 emitted_in_progress: false,
26 }
27 }
28
29 pub fn process_sse_line(&mut self, line: &str, output_offset: usize) -> Option<EventFrame> {
30 let mut frame = normalize_sse_line(line)?;
31 self.process_event(&mut frame, output_offset).then_some(frame)
32 }
33
34 #[must_use]
35 pub fn process_event(&mut self, frame: &mut EventFrame, output_offset: usize) -> bool {
36 if !self.should_emit_lifecycle(frame.event_type) {
37 return false;
38 }
39 self.stamp_event(frame, output_offset);
40 true
41 }
42
43 fn stamp_event(&mut self, frame: &mut EventFrame, output_offset: usize) {
44 frame.wire.sequence_number = Some(self.take_sequence_number());
45 rebase_output_index(&mut frame.wire, output_offset);
46 }
47
48 pub(crate) fn terminal_response_chunk(&mut self, payload: &ResponsePayload) -> ExecutorResult<String> {
49 let mut frame = terminal_response_frame(payload)?;
50 self.stamp_event(&mut frame, 0);
51 serialize_sse_frame(&frame)
52 }
53
54 pub(crate) fn executor_error_chunk(&mut self, error: &ExecutorError) -> String {
55 let mut frame = executor_error_frame(error);
56 self.stamp_event(&mut frame, 0);
57 serialize_sse_frame(&frame)
58 .unwrap_or_else(|_| error_sse_chunk(&error.to_string(), frame.sequence_number().unwrap_or(0)))
59 }
60
61 fn should_emit_lifecycle(&mut self, event_type: SSEEventType) -> bool {
62 match event_type {
63 SSEEventType::ResponseCreated => take_once(&mut self.emitted_created),
64 SSEEventType::ResponseInProgress => take_once(&mut self.emitted_in_progress),
65 _ => true,
66 }
67 }
68
69 fn take_sequence_number(&mut self) -> u64 {
70 let sequence_number = self.next_sequence_number;
71 self.next_sequence_number = self.next_sequence_number.saturating_add(1);
72 sequence_number
73 }
74}
75
76impl Default for GatewayStreamAccumulator {
77 fn default() -> Self {
78 Self::new()
79 }
80}
81
82fn take_once(already_taken: &mut bool) -> bool {
83 if *already_taken {
84 false
85 } else {
86 *already_taken = true;
87 true
88 }
89}
90
91fn rebase_output_index(wire: &mut WireEvent, output_offset: usize) {
92 let Some(offset) = u64::try_from(output_offset).ok().filter(|offset| *offset > 0) else {
93 return;
94 };
95 if let Some(index) = wire.output_index {
96 wire.output_index = Some(index.saturating_add(offset));
97 }
98}
99
100fn terminal_response_frame(payload: &ResponsePayload) -> ExecutorResult<EventFrame> {
101 let event_type = match payload.terminal_event_type() {
102 "response.incomplete" => SSEEventType::ResponseIncomplete,
103 "response.failed" => SSEEventType::ResponseFailed,
104 "response.in_progress" => SSEEventType::ResponseInProgress,
105 _ => SSEEventType::ResponseCompleted,
106 };
107 let mut rest = serde_json::Map::new();
108 rest.insert(
109 "response".to_owned(),
110 serialize_to_value(payload).map_err(ExecutorError::JsonError)?,
111 );
112 EventFrame::synthetic(event_type, rest)
113 .ok_or_else(|| ExecutorError::StreamError("terminal response event has no wire representation".to_owned()))
114}
115
116fn executor_error_frame(error: &ExecutorError) -> EventFrame {
117 let error_type = error.error_type();
118 let code = error.error_code();
119 let mut wire = WireEvent::new("error");
120 wire.rest
121 .insert("status".to_owned(), serde_json::json!(error.http_status().as_u16()));
122 let mut error_details = serde_json::Map::new();
123 error_details.insert("message".to_owned(), serde_json::json!(error.error_message()));
124 error_details.insert("type".to_owned(), serde_json::json!(error_type));
125 error_details.insert("code".to_owned(), serde_json::json!(code));
126 if let Some(param) = error.error_param() {
127 error_details.insert("param".to_owned(), serde_json::json!(param));
128 }
129 wire.rest
130 .insert("error".to_owned(), serde_json::Value::Object(error_details));
131 EventFrame {
132 event_type: SSEEventType::Other,
133 payload: EventPayload::None,
134 wire,
135 }
136}
137
138pub(super) fn error_sse_chunk(message: &str, sequence_number: u64) -> String {
139 let code = "server_error";
140 let mut wire = WireEvent::new("error");
141 wire.sequence_number = Some(sequence_number);
142 wire.rest.insert("status".to_owned(), serde_json::json!(500));
143 wire.rest.insert(
144 "error".to_owned(),
145 serde_json::json!({
146 "message": message,
147 "type": code,
148 "code": code,
149 }),
150 );
151 let frame = EventFrame {
152 event_type: SSEEventType::Other,
153 payload: EventPayload::None,
154 wire,
155 };
156 serialize_sse_frame(&frame).unwrap_or_else(|_| {
157 format!("event: error\ndata: {{\"type\":\"error\",\"status\":500,\"sequence_number\":{sequence_number}}}\n\n")
158 })
159}
160
161pub(super) fn synthetic_event(
162 event_type: SSEEventType,
163 rest: impl IntoIterator<Item = (String, Value)>,
164) -> ExecutorResult<EventFrame> {
165 EventFrame::synthetic(event_type, rest.into_iter().collect())
166 .ok_or_else(|| ExecutorError::StreamError("synthetic event has no wire representation".to_owned()))
167}
168
169pub(super) fn emit_sse_frame(
170 sender: &tokio::sync::mpsc::UnboundedSender<StreamEvent>,
171 frame: &EventFrame,
172) -> ExecutorResult<()> {
173 let sequence_number = frame
174 .sequence_number()
175 .ok_or_else(|| ExecutorError::StreamError("stream event has no sequence number".to_owned()))?;
176 sender
177 .send(StreamEvent {
178 content: serialize_sse_frame(frame)?,
179 sequence_number,
180 })
181 .map_err(|_| ExecutorError::StreamError("stream receiver closed while emitting gateway event".to_owned()))
182}
183
184fn serialize_sse_frame(frame: &EventFrame) -> ExecutorResult<String> {
185 let event_json = serialize_to_string(&frame.wire).map_err(ExecutorError::JsonError)?;
186 let event_name =
187 frame.wire.event_type.as_deref().filter(|event_name| {
188 !event_name.is_empty() && !event_name.bytes().any(|byte| matches!(byte, b'\r' | b'\n'))
189 });
190 let event_name_len = event_name.map_or(0, str::len);
191 let mut chunk = String::with_capacity(event_name_len + event_json.len() + 16);
192 if let Some(event_name) = event_name {
193 chunk.push_str("event: ");
194 chunk.push_str(event_name);
195 chunk.push('\n');
196 }
197 chunk.push_str("data: ");
198 chunk.push_str(&event_json);
199 chunk.push_str("\n\n");
200 Ok(chunk)
201}
202
203#[cfg(test)]
204mod tests {
205 use super::*;
206 use crate::StorageError;
207
208 fn parse_named_sse_event(chunk: &str) -> (&str, serde_json::Value) {
209 let body = chunk.strip_suffix("\n\n").expect("SSE event terminator");
210 let (event_line, data_line) = body.split_once('\n').expect("named SSE event and data lines");
211 let event_name = event_line.strip_prefix("event: ").expect("SSE event prefix");
212 let data = data_line.strip_prefix("data: ").expect("SSE data prefix");
213 let event = serde_json::from_str(data).expect("valid event JSON");
214 (event_name, event)
215 }
216
217 #[test]
218 fn process_sse_line_numbers_and_rebases_output_index() {
219 let mut accumulator = GatewayStreamAccumulator::new();
220 let frame = accumulator
221 .process_sse_line(
222 r#"data: {"type":"response.output_text.delta","output_index":2,"delta":"hi"}"#,
223 3,
224 )
225 .expect("line should normalize");
226
227 assert_eq!(frame.sequence_number(), Some(0));
228 assert_eq!(frame.wire.sequence_number, Some(0));
229 assert_eq!(frame.wire.output_index, Some(5));
230 assert_eq!(frame.wire.rest["delta"], "hi");
231 }
232
233 #[test]
234 fn error_sse_chunk_escapes_error_messages() {
235 let chunk = error_sse_chunk("task failed: \"unexpected\"\nretry", 7);
236 let (event_name, event) = parse_named_sse_event(&chunk);
237
238 assert_eq!(event_name, "error");
239 assert_eq!(event["type"], "error");
240 assert_eq!(event["sequence_number"], 7);
241 assert_eq!(event["error"]["message"], "task failed: \"unexpected\"\nretry");
242 }
243
244 #[test]
245 fn executor_conflict_sse_chunk_uses_client_conflict_contract() {
246 let mut accumulator = GatewayStreamAccumulator::new();
247 let error = ExecutorError::Persistence(Box::new(ExecutorError::ConversationLocked {
248 source: StorageError::ConversationConflict {
249 conversation_id: "conv_test".to_owned(),
250 },
251 }));
252 let chunk = accumulator.executor_error_chunk(&error);
253 let (event_name, event) = parse_named_sse_event(&chunk);
254
255 assert_eq!(event_name, "error");
256 assert_eq!(event["type"], "error");
257 assert_eq!(event["status"], 400);
258 assert_eq!(
259 event["error"],
260 serde_json::json!({
261 "message": "conversation changed while the response was being generated; retry the request",
262 "type": "invalid_request_error",
263 "code": "conversation_locked",
264 "param": "conversation"
265 })
266 );
267 }
268
269 #[test]
270 fn emits_in_progress_terminal_event_after_lifecycle_event() {
271 let mut accumulator = GatewayStreamAccumulator::new();
272 accumulator
273 .process_sse_line(r#"data: {"type":"response.in_progress"}"#, 0)
274 .expect("first lifecycle event should be emitted");
275
276 let payload: ResponsePayload = serde_json::from_value(serde_json::json!({
277 "id": "resp_1",
278 "object": "response",
279 "created_at": 0,
280 "model": "test",
281 "status": "in_progress",
282 "output": [],
283 "usage": null,
284 "incomplete_details": null,
285 "error": null,
286 "previous_response_id": null,
287 "conversation_id": null,
288 "instructions": null
289 }))
290 .expect("valid response payload");
291
292 let chunk = accumulator
293 .terminal_response_chunk(&payload)
294 .expect("terminal event serializes");
295 assert!(chunk.starts_with("event: response.in_progress\n"));
296 assert!(chunk.contains("\"type\":\"response.in_progress\""));
297 assert!(chunk.contains("\"sequence_number\":1"));
298 }
299
300 #[test]
301 fn serialize_sse_frame_uses_wire_event_type_as_event_name() {
302 let frame = synthetic_event(SSEEventType::ResponseCreated, []).expect("synthetic event");
303 let chunk = serialize_sse_frame(&frame).expect("event serializes");
304 let (event_name, event) = parse_named_sse_event(&chunk);
305
306 assert_eq!(event_name, "response.created");
307 assert_eq!(event["type"], event_name);
308 }
309
310 #[test]
311 fn serialize_sse_frame_omits_invalid_event_name() {
312 let frame = EventFrame {
313 event_type: SSEEventType::Other,
314 payload: EventPayload::None,
315 wire: WireEvent::new("error\nevent: injected"),
316 };
317 let chunk = serialize_sse_frame(&frame).expect("event serializes");
318
319 assert!(chunk.starts_with("data: "));
320 assert!(!chunk.contains("\nevent: injected\n"));
321 }
322}