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