Skip to main content

agentic_core/executor/
gateway_accumulator.rs

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}