Skip to main content

gproxy_transform/transform/stream_adapter/
responses.rs

1//! Typed aggregation state machine for Responses SSE streams: fills in
2//! synthetic lifecycle events (item_added / *_done / response.completed) that
3//! sparse upstreams omit, so the inbound side always sees a complete stream.
4
5mod text;
6mod tool;
7
8use std::collections::BTreeMap;
9
10use crate::protocol::openai::{
11    Extra, KnownResponseStreamEvent as KnownEvent, ResponseItem, ResponseMessageItem,
12    ResponseObject, ResponseObjectType, ResponseOutputItem, ResponseStatus, ResponseStreamEvent,
13    TypedResponseItem,
14};
15
16use super::{SseDecoder, SseFrame, encode_responses_event};
17use crate::transform::TransformError;
18use text::ResponsesTextItemState;
19use tool::{ResponsesToolItemState, ResponsesToolKind};
20
21/// Stateful normalizer for an upstream that already speaks Responses SSE.
22/// Frames outside the currently modeled typed schema are forwarded unchanged.
23#[derive(Default)]
24pub struct ResponsesStreamNormalizer {
25    decoder: SseDecoder,
26    responses: ResponsesStreamState,
27    done_seen: bool,
28    failed: bool,
29}
30
31impl ResponsesStreamNormalizer {
32    pub fn new() -> Self {
33        Self::default()
34    }
35
36    pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
37        if self.failed {
38            return Err(TransformError::InvalidInput {
39                reason: "Responses stream normalizer is failed".to_owned(),
40            });
41        }
42        let mut out = Vec::new();
43        let frames = self
44            .decoder
45            .push(chunk)
46            .inspect_err(|_| self.failed = true)?;
47        for frame in frames {
48            if let Err(error) = self.normalize_into(frame, &mut out) {
49                self.failed = true;
50                return Err(error);
51            }
52        }
53        Ok(out)
54    }
55
56    pub fn finish(&mut self) -> Result<Vec<u8>, TransformError> {
57        if self.failed {
58            return Err(TransformError::InvalidInput {
59                reason: "cannot finish failed Responses stream normalizer".to_owned(),
60            });
61        }
62        let mut out = Vec::new();
63        if let Some(frame) = self.decoder.finish()? {
64            self.normalize_into(frame, &mut out)?;
65        }
66        for event in self.responses.finish() {
67            encode_responses_event(&event, &mut out)?;
68        }
69        if self.done_seen {
70            out.extend_from_slice(b"data: [DONE]\n\n");
71        }
72        Ok(out)
73    }
74
75    fn normalize_into(&mut self, frame: SseFrame, out: &mut Vec<u8>) -> Result<(), TransformError> {
76        if frame.data.trim() == "[DONE]" {
77            self.done_seen = true;
78            return Ok(());
79        }
80        let event = match serde_json::from_str::<ResponseStreamEvent>(&frame.data) {
81            Ok(event) => event,
82            Err(error) => {
83                // Responses evolves faster than the typed protocol model. An
84                // SSE data frame this crate cannot decode must still reach the
85                // client unchanged; the normalizer simply cannot add lifecycle
86                // events for it.
87                tracing::warn!(
88                    error = %error,
89                    "Responses stream event typed decode failed; forwarding original SSE frame"
90                );
91                out.extend_from_slice(frame.encode().as_bytes());
92                return Ok(());
93            }
94        };
95        for event in self.responses.push(event) {
96            encode_responses_event(&event, out)?;
97        }
98        Ok(())
99    }
100}
101
102#[derive(Default)]
103pub(super) struct ResponsesStreamState {
104    message: ResponsesTextItemState,
105    reasoning: ResponsesTextItemState,
106    tools: BTreeMap<u32, ResponsesToolItemState>,
107    completed: bool,
108}
109
110impl ResponsesStreamState {
111    pub(super) fn push(&mut self, event: ResponseStreamEvent) -> Vec<ResponseStreamEvent> {
112        let mut event = match event {
113            ResponseStreamEvent::Known(known) => known,
114            unknown => return vec![unknown],
115        };
116        let mut out = match &mut event {
117            KnownEvent::ResponseOutputTextDelta {
118                content_index,
119                delta,
120                item_id,
121                output_index,
122                ..
123            } => {
124                let mut out = self.finish_reasoning();
125                out.extend(self.message.ensure(
126                    item_id,
127                    *output_index,
128                    *content_index,
129                    text::message_item_added,
130                ));
131                self.message.text.push_str(delta);
132                out
133            }
134            KnownEvent::ResponseReasoningTextDelta {
135                content_index,
136                delta,
137                item_id,
138                output_index,
139                ..
140            } => {
141                let out = self.reasoning.ensure(
142                    item_id,
143                    *output_index,
144                    *content_index,
145                    text::reasoning_item_added,
146                );
147                self.reasoning.text.push_str(delta);
148                out
149            }
150            KnownEvent::ResponseFunctionCallArgumentsDelta {
151                delta,
152                item_id,
153                output_index,
154                ..
155            } => {
156                self.note_tool_input_delta(
157                    ResponsesToolKind::Function,
158                    *output_index,
159                    item_id,
160                    delta,
161                );
162                Vec::new()
163            }
164            KnownEvent::ResponseCustomToolCallInputDelta {
165                delta,
166                item_id,
167                output_index,
168                ..
169            } => {
170                self.note_tool_input_delta(
171                    ResponsesToolKind::Custom,
172                    *output_index,
173                    item_id,
174                    delta,
175                );
176                Vec::new()
177            }
178            KnownEvent::ResponseFunctionCallArgumentsDone {
179                arguments,
180                item_id,
181                name,
182                output_index,
183                ..
184            } => {
185                self.note_tool_input_done(
186                    ResponsesToolKind::Function,
187                    *output_index,
188                    item_id,
189                    arguments,
190                    Some(name),
191                );
192                Vec::new()
193            }
194            KnownEvent::ResponseCustomToolCallInputDone {
195                input,
196                item_id,
197                output_index,
198                ..
199            } => {
200                self.note_tool_input_done(
201                    ResponsesToolKind::Custom,
202                    *output_index,
203                    item_id,
204                    input,
205                    None,
206                );
207                Vec::new()
208            }
209            KnownEvent::ResponseCompleted { response, .. } => {
210                let mut out = self.finish_reasoning();
211                out.extend(self.finish_message());
212                out.extend(self.finish_tools());
213                self.patch_completed_output(response);
214                self.completed = true;
215                out
216            }
217            KnownEvent::ResponseOutputItemAdded {
218                item, output_index, ..
219            } => {
220                self.note_item_added(item, *output_index);
221                Vec::new()
222            }
223            KnownEvent::ResponseOutputItemDone {
224                item, output_index, ..
225            } => {
226                self.note_item_done(item, *output_index);
227                Vec::new()
228            }
229            KnownEvent::ResponseOutputTextDone { text, .. } => {
230                self.message.note_done_text(text);
231                Vec::new()
232            }
233            KnownEvent::ResponseReasoningTextDone { text, .. } => {
234                self.reasoning.note_done_text(text);
235                Vec::new()
236            }
237            _ => Vec::new(),
238        };
239        out.push(ResponseStreamEvent::Known(event));
240        out
241    }
242
243    pub(super) fn finish(&mut self) -> Vec<ResponseStreamEvent> {
244        if self.completed {
245            return Vec::new();
246        }
247        let mut out = self.finish_reasoning();
248        out.extend(self.finish_message());
249        out.extend(self.finish_tools());
250        if !out.is_empty() {
251            out.push(known(KnownEvent::ResponseCompleted {
252                response: Box::new(fallback_completed_response()),
253                sequence_number: None,
254                extra: Extra::new(),
255            }));
256            self.completed = true;
257        }
258        out
259    }
260
261    fn finish_message(&mut self) -> Vec<ResponseStreamEvent> {
262        self.message.finish(text::message_done_events)
263    }
264
265    fn finish_reasoning(&mut self) -> Vec<ResponseStreamEvent> {
266        self.reasoning.finish(text::reasoning_done_events)
267    }
268
269    fn note_item_added(&mut self, item: &ResponseOutputItem, output_index: u32) {
270        match &item.0 {
271            ResponseItem::Message(message) if message_has_type(message) => {
272                self.message.note_added(message_id(message), output_index);
273            }
274            ResponseItem::Typed(TypedResponseItem::Reasoning { id, .. }) => {
275                self.reasoning.note_added(id.as_deref(), output_index);
276            }
277            ResponseItem::Typed(typed @ TypedResponseItem::FunctionCall { .. }) => {
278                self.note_tool_added(typed, ResponsesToolKind::Function, output_index);
279            }
280            ResponseItem::Typed(typed @ TypedResponseItem::CustomToolCall { .. }) => {
281                self.note_tool_added(typed, ResponsesToolKind::Custom, output_index);
282            }
283            _ => {}
284        }
285    }
286
287    fn note_item_done(&mut self, item: &ResponseOutputItem, output_index: u32) {
288        match &item.0 {
289            ResponseItem::Message(message) if message_has_type(message) => {
290                self.message
291                    .note_item_done(message_id(message), output_index);
292            }
293            ResponseItem::Typed(TypedResponseItem::Reasoning { id, .. }) => {
294                self.reasoning.note_item_done(id.as_deref(), output_index);
295            }
296            ResponseItem::Typed(typed @ TypedResponseItem::FunctionCall { .. }) => {
297                self.note_tool_item_done(typed, ResponsesToolKind::Function, output_index);
298            }
299            ResponseItem::Typed(typed @ TypedResponseItem::CustomToolCall { .. }) => {
300                self.note_tool_item_done(typed, ResponsesToolKind::Custom, output_index);
301            }
302            _ => {}
303        }
304    }
305
306    fn note_tool_added(&mut self, item: &TypedResponseItem, kind: ResponsesToolKind, index: u32) {
307        let state = self.tools.entry(index).or_default();
308        state.note_kind(kind, index);
309        state.note_item(item);
310    }
311
312    fn note_tool_item_done(
313        &mut self,
314        item: &TypedResponseItem,
315        kind: ResponsesToolKind,
316        index: u32,
317    ) {
318        let state = self.tools.entry(index).or_default();
319        state.note_kind(kind, index);
320        state.item_done = true;
321        state.note_item(item);
322    }
323
324    fn note_tool_input_delta(
325        &mut self,
326        kind: ResponsesToolKind,
327        index: u32,
328        item_id: &mut String,
329        delta: &str,
330    ) {
331        let state = self.tools.entry(index).or_default();
332        state.note_kind(kind, index);
333        state.note_event_item_id(item_id);
334        state.rewrite_event_item_id(item_id);
335        state.input.push_str(delta);
336    }
337
338    fn note_tool_input_done(
339        &mut self,
340        kind: ResponsesToolKind,
341        index: u32,
342        item_id: &mut String,
343        input: &str,
344        name: Option<&str>,
345    ) {
346        let state = self.tools.entry(index).or_default();
347        state.note_kind(kind, index);
348        state.note_event_item_id(item_id);
349        state.rewrite_event_item_id(item_id);
350        input.clone_into(&mut state.input);
351        if let Some(name) = name {
352            state.name.get_or_insert_with(|| name.to_owned());
353        }
354        state.input_done = true;
355    }
356
357    fn finish_tools(&mut self) -> Vec<ResponseStreamEvent> {
358        let mut out = Vec::new();
359        for state in self.tools.values_mut() {
360            if !state.can_finish() {
361                continue;
362            }
363            if !state.input_done {
364                out.push(state.input_done_event());
365                state.input_done = true;
366            }
367            if !state.item_done {
368                out.push(state.item_done_event());
369                state.item_done = true;
370            }
371        }
372        out
373    }
374
375    fn patch_completed_output(&self, response: &mut ResponseObject) {
376        if !response.output.is_empty() {
377            return;
378        }
379        let output = self.completed_output_items();
380        if !output.is_empty() {
381            response.output = output;
382        }
383    }
384
385    fn completed_output_items(&self) -> Vec<ResponseOutputItem> {
386        use crate::protocol::openai::ResponseItemLifecycleStatus::Completed;
387        let mut output = Vec::new();
388        if self.reasoning.started {
389            output.push(text::reasoning_item(&self.reasoning, Completed));
390        }
391        if self.message.started {
392            output.push(text::message_item(&self.message, Completed));
393        }
394        output.extend(
395            self.tools
396                .values()
397                .filter(|state| state.can_finish())
398                .map(ResponsesToolItemState::completed_item),
399        );
400        output
401    }
402}
403
404fn known(event: KnownEvent) -> ResponseStreamEvent {
405    ResponseStreamEvent::Known(event)
406}
407
408fn message_has_type(message: &ResponseMessageItem) -> bool {
409    match message {
410        ResponseMessageItem::Output(_) => true,
411        ResponseMessageItem::Input(input) => input.type_.is_some(),
412        ResponseMessageItem::EasyInput(easy) => easy.type_.is_some(),
413        _ => {
414            unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
415        }
416    }
417}
418
419fn message_id(message: &ResponseMessageItem) -> Option<&str> {
420    match message {
421        ResponseMessageItem::Output(output) => Some(&output.id),
422        ResponseMessageItem::Input(input) => input.id.as_deref(),
423        ResponseMessageItem::EasyInput(_) => None,
424        _ => {
425            unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
426        }
427    }
428}
429
430/// Fallback `response.completed` payload for streams that never sent one.
431fn fallback_completed_response() -> ResponseObject {
432    crate::protocol::wire!(ResponseObject {
433        id: "resp_0".to_owned(),
434        created_at: 0,
435        background: None,
436        completed_at: Some(0),
437        conversation: None,
438        error: None,
439        incomplete_details: None,
440        instructions: None,
441        max_output_tokens: None,
442        max_tool_calls: None,
443        metadata: None,
444        model: None,
445        moderation: None,
446        multi_agent: None,
447        object: ResponseObjectType::Response,
448        output: Vec::new(),
449        output_text: None,
450        parallel_tool_calls: None,
451        prompt: None,
452        prompt_cache_key: None,
453        prompt_cache_options: None,
454        prompt_cache_retention: None,
455        previous_response_id: None,
456        reasoning: None,
457        safety_identifier: None,
458        service_tier: None,
459        status: Some(ResponseStatus::Completed),
460        store: None,
461        temperature: None,
462        text: None,
463        tool_choice: None,
464        tools: None,
465        top_logprobs: None,
466        top_p: None,
467        truncation: None,
468        usage: None,
469        user: None,
470        extra: Extra::new(),
471    })
472}