Skip to main content

gproxy_transform/transform/stream_adapter/
responses.rs

1mod text;
2mod tool;
3
4use std::collections::BTreeMap;
5
6use serde_json::{Value, json};
7
8use super::{ContentGenerationKind, SseDecoder, SseFrame, encode_frame};
9use text::{
10    ResponsesTextItemState, message_item, message_item_added, reasoning_item, reasoning_item_added,
11};
12use tool::{ResponsesToolItemState, ResponsesToolKind};
13
14/// Stateful normalizer for an upstream that already speaks Responses SSE.
15#[derive(Default)]
16pub struct ResponsesStreamNormalizer {
17    decoder: SseDecoder,
18    responses: ResponsesStreamState,
19}
20
21impl ResponsesStreamNormalizer {
22    pub fn new() -> Self {
23        Self::default()
24    }
25
26    pub fn push(&mut self, chunk: &[u8]) -> Vec<u8> {
27        let mut out = Vec::new();
28        for frame in self.decoder.push(chunk) {
29            self.normalize_into(frame, &mut out);
30        }
31        out
32    }
33
34    pub fn finish(&mut self) -> Vec<u8> {
35        let mut out = Vec::new();
36        if let Some(frame) = self.decoder.finish() {
37            self.normalize_into(frame, &mut out);
38        }
39        out
40    }
41
42    fn normalize_into(&mut self, frame: SseFrame, out: &mut Vec<u8>) {
43        if frame.data.trim() == "[DONE]" {
44            out.extend_from_slice(frame.encode().as_bytes());
45            return;
46        }
47        let Ok(event) = serde_json::from_str::<Value>(&frame.data) else {
48            out.extend_from_slice(frame.encode().as_bytes());
49            return;
50        };
51        for event in self.responses.push(event) {
52            out.extend_from_slice(
53                encode_frame(ContentGenerationKind::OpenAiResponses, &event).as_bytes(),
54            );
55        }
56    }
57}
58
59#[derive(Default)]
60pub(super) struct ResponsesStreamState {
61    message: ResponsesTextItemState,
62    reasoning: ResponsesTextItemState,
63    tools: BTreeMap<u32, ResponsesToolItemState>,
64    completed: bool,
65}
66
67impl ResponsesStreamState {
68    pub(super) fn push(&mut self, mut event: Value) -> Vec<Value> {
69        match event.get("type").and_then(Value::as_str) {
70            Some("response.output_text.delta") => {
71                let mut out = self.finish_reasoning();
72                out.extend(self.message.ensure(&event, "msg_0", message_item_added));
73                self.message.push_delta(&event);
74                out.push(event);
75                out
76            }
77            Some("response.reasoning_text.delta") => {
78                let mut out = self
79                    .reasoning
80                    .ensure(&event, "reasoning_0", reasoning_item_added);
81                self.reasoning.push_delta(&event);
82                out.push(event);
83                out
84            }
85            Some("response.function_call_arguments.delta") => {
86                self.note_tool_input_delta(&mut event, ResponsesToolKind::Function);
87                vec![event]
88            }
89            Some("response.custom_tool_call_input.delta") => {
90                self.note_tool_input_delta(&mut event, ResponsesToolKind::Custom);
91                vec![event]
92            }
93            Some("response.function_call_arguments.done") => {
94                self.note_tool_input_done(&mut event, ResponsesToolKind::Function);
95                vec![event]
96            }
97            Some("response.custom_tool_call_input.done") => {
98                self.note_tool_input_done(&mut event, ResponsesToolKind::Custom);
99                vec![event]
100            }
101            Some("response.completed") => {
102                let mut out = self.finish_reasoning();
103                out.extend(self.finish_message());
104                out.extend(self.finish_tools());
105                self.patch_completed_output(&mut event);
106                self.completed = true;
107                out.push(event);
108                out
109            }
110            Some("response.output_item.added") => {
111                self.note_item_added(&event);
112                vec![event]
113            }
114            Some("response.output_item.done") => {
115                self.note_item_done(&event);
116                vec![event]
117            }
118            Some("response.output_text.done") => {
119                self.message.note_done_text(&event);
120                vec![event]
121            }
122            Some("response.reasoning_text.done") => {
123                self.reasoning.note_done_text(&event);
124                vec![event]
125            }
126            _ => vec![event],
127        }
128    }
129
130    pub(super) fn finish(&mut self) -> Vec<Value> {
131        if self.completed {
132            return Vec::new();
133        }
134        let mut out = self.finish_reasoning();
135        out.extend(self.finish_message());
136        if !out.is_empty() {
137            out.extend(self.finish_tools());
138            out.push(json!({
139                "type": "response.completed",
140                "response": {"id":"resp_0","object":"response","created_at":0,
141                    "completed_at":0,"status":"completed","output":[]},
142            }));
143            self.completed = true;
144        }
145        out
146    }
147
148    fn finish_message(&mut self) -> Vec<Value> {
149        self.message.finish(|state| {
150            vec![
151                json!({"type":"response.output_text.done","output_index":state.output_index(),
152                "item_id":state.id(),"content_index":state.content_index(),"text":state.text}),
153                json!({"type":"response.content_part.done","output_index":state.output_index(),
154                "item_id":state.id(),"content_index":state.content_index(),
155                "part":{"type":"output_text","text":state.text,"annotations":[]}}),
156                json!({"type":"response.output_item.done","output_index":state.output_index(),
157                "item":message_item(state,"completed")}),
158            ]
159        })
160    }
161
162    fn finish_reasoning(&mut self) -> Vec<Value> {
163        self.reasoning.finish(|state| {
164            vec![
165                json!({"type":"response.reasoning_text.done","output_index":state.output_index(),
166                "item_id":state.id(),"content_index":state.content_index(),"text":state.text}),
167                json!({"type":"response.output_item.done","output_index":state.output_index(),
168                "item":reasoning_item(state,"completed")}),
169            ]
170        })
171    }
172
173    fn note_item_added(&mut self, event: &Value) {
174        match item_type(event) {
175            Some("message") => self.message.note_added(event),
176            Some("reasoning") => self.reasoning.note_added(event),
177            Some("function_call") => self.note_tool_added(event, ResponsesToolKind::Function),
178            Some("custom_tool_call") => self.note_tool_added(event, ResponsesToolKind::Custom),
179            _ => {}
180        }
181    }
182
183    fn note_item_done(&mut self, event: &Value) {
184        match item_type(event) {
185            Some("message") => self.message.note_item_done(event),
186            Some("reasoning") => self.reasoning.note_item_done(event),
187            Some("function_call") => self.note_tool_item_done(event, ResponsesToolKind::Function),
188            Some("custom_tool_call") => self.note_tool_item_done(event, ResponsesToolKind::Custom),
189            _ => {}
190        }
191    }
192
193    fn note_tool_added(&mut self, event: &Value, kind: ResponsesToolKind) {
194        let Some(index) = event_output_index(event) else {
195            return;
196        };
197        let state = self.tools.entry(index).or_default();
198        state.note_kind(kind, index);
199        if let Some(item) = event.get("item") {
200            state.note_item(item);
201        }
202    }
203
204    fn note_tool_item_done(&mut self, event: &Value, kind: ResponsesToolKind) {
205        let Some(index) = event_output_index(event) else {
206            return;
207        };
208        let state = self.tools.entry(index).or_default();
209        state.note_kind(kind, index);
210        state.item_done = true;
211        if let Some(item) = event.get("item") {
212            state.note_item(item);
213        }
214    }
215
216    fn note_tool_input_delta(&mut self, event: &mut Value, kind: ResponsesToolKind) {
217        let Some(index) = event_output_index(event) else {
218            return;
219        };
220        let state = self.tools.entry(index).or_default();
221        state.note_kind(kind, index);
222        state.note_event_item_id(event);
223        if let Some(id) = state.item_id.as_deref() {
224            event["item_id"] = Value::String(id.into());
225        }
226        if let Some(delta) = event.get("delta").and_then(Value::as_str) {
227            state.input.push_str(delta);
228        }
229    }
230
231    fn note_tool_input_done(&mut self, event: &mut Value, kind: ResponsesToolKind) {
232        let Some(index) = event_output_index(event) else {
233            return;
234        };
235        let state = self.tools.entry(index).or_default();
236        state.note_kind(kind, index);
237        state.note_event_item_id(event);
238        if let Some(id) = state.item_id.as_deref() {
239            event["item_id"] = Value::String(id.into());
240        }
241        let field = match kind {
242            ResponsesToolKind::Function => "arguments",
243            ResponsesToolKind::Custom => "input",
244        };
245        if let Some(done) = event.get(field).and_then(Value::as_str) {
246            state.input = done.into();
247        }
248        if matches!(kind, ResponsesToolKind::Function)
249            && let Some(name) = event.get("name").and_then(Value::as_str)
250        {
251            state.name.get_or_insert_with(|| name.into());
252        }
253        state.input_done = true;
254    }
255
256    fn finish_tools(&mut self) -> Vec<Value> {
257        let mut out = Vec::new();
258        for state in self.tools.values_mut() {
259            if !state.can_finish() {
260                continue;
261            }
262            if !state.input_done {
263                out.push(state.input_done_event());
264                state.input_done = true;
265            }
266            if !state.item_done {
267                out.push(state.item_done_event());
268                state.item_done = true;
269            }
270        }
271        out
272    }
273
274    fn patch_completed_output(&self, event: &mut Value) {
275        let Some(response) = event.get_mut("response").and_then(Value::as_object_mut) else {
276            return;
277        };
278        if !response
279            .get("output")
280            .and_then(Value::as_array)
281            .is_none_or(Vec::is_empty)
282        {
283            return;
284        }
285        let output = self.completed_output_items();
286        if !output.is_empty() {
287            response.insert("output".into(), Value::Array(output));
288        }
289    }
290
291    fn completed_output_items(&self) -> Vec<Value> {
292        let mut output = Vec::new();
293        if self.reasoning.started {
294            output.push(reasoning_item(&self.reasoning, "completed"));
295        }
296        if self.message.started {
297            output.push(message_item(&self.message, "completed"));
298        }
299        output.extend(
300            self.tools
301                .values()
302                .filter(|state| state.can_finish())
303                .map(|state| state.item("completed")),
304        );
305        output
306    }
307}
308
309fn item_type(event: &Value) -> Option<&str> {
310    event.get("item")?.get("type")?.as_str()
311}
312
313fn event_output_index(event: &Value) -> Option<u32> {
314    event.get("output_index")?.as_u64()?.try_into().ok()
315}