gproxy_transform/transform/generate_content/
openai_responses_websocket.rs1use serde_json::Value;
9
10use crate::transform::common::sse::SseDecoder;
11use crate::transform::{TransformContext, TransformError};
12
13pub fn http_request_to_ws_request(
16 mut value: Value,
17 _ctx: &TransformContext,
18) -> Result<Value, TransformError> {
19 let object = value
20 .as_object_mut()
21 .ok_or_else(|| TransformError::InvalidInput {
22 reason: "responses websocket request must be a JSON object".to_owned(),
23 })?;
24 object.insert(
25 "type".to_owned(),
26 Value::String("response.create".to_owned()),
27 );
28 Ok(value)
29}
30
31pub fn ws_request_to_http_request(
34 mut value: Value,
35 _ctx: &TransformContext,
36) -> Result<Value, TransformError> {
37 let object = value
38 .as_object_mut()
39 .ok_or_else(|| TransformError::InvalidInput {
40 reason: "websocket frame must be a JSON object".to_owned(),
41 })?;
42 let frame_type = object
43 .remove("type")
44 .and_then(|value| value.as_str().map(str::to_owned))
45 .ok_or_else(|| TransformError::InvalidInput {
46 reason: "websocket frame missing type".to_owned(),
47 })?;
48 if frame_type != "response.create" {
49 return Err(TransformError::InvalidInput {
50 reason: format!("unsupported websocket frame type: {frame_type}"),
51 });
52 }
53 object.insert("stream".to_owned(), Value::Bool(true));
54 Ok(value)
55}
56
57pub fn identity(value: Value, _ctx: &TransformContext) -> Value {
58 value
59}
60
61pub fn identity_result(value: Value, _ctx: &TransformContext) -> Result<Value, TransformError> {
62 Ok(value)
63}
64
65pub fn response_create_frame_to_response_body(frame: &[u8]) -> Result<Vec<u8>, TransformError> {
72 let mut value: Value =
73 serde_json::from_slice(frame).map_err(|error| TransformError::InvalidInput {
74 reason: format!("decode websocket frame: {error}"),
75 })?;
76 let object = value
77 .as_object_mut()
78 .ok_or_else(|| TransformError::InvalidInput {
79 reason: "websocket frame must be a JSON object".to_owned(),
80 })?;
81 let frame_type = object
82 .remove("type")
83 .and_then(|value| value.as_str().map(str::to_owned))
84 .ok_or_else(|| TransformError::InvalidInput {
85 reason: "websocket frame missing type".to_owned(),
86 })?;
87 if frame_type != "response.create" {
88 return Err(TransformError::InvalidInput {
89 reason: format!("unsupported websocket frame type: {frame_type}"),
90 });
91 }
92
93 object.insert("stream".to_owned(), Value::Bool(true));
94 serde_json::to_vec(&value).map_err(|error| TransformError::Serialization {
95 reason: error.to_string(),
96 })
97}
98
99pub fn validate_response_create_frame(frame: &[u8]) -> Result<(), TransformError> {
101 let value: Value =
102 serde_json::from_slice(frame).map_err(|error| TransformError::InvalidInput {
103 reason: format!("decode websocket frame: {error}"),
104 })?;
105 let object = value
106 .as_object()
107 .ok_or_else(|| TransformError::InvalidInput {
108 reason: "websocket frame must be a JSON object".to_owned(),
109 })?;
110 let frame_type = object
111 .get("type")
112 .and_then(|value| value.as_str())
113 .ok_or_else(|| TransformError::InvalidInput {
114 reason: "websocket frame missing type".to_owned(),
115 })?;
116 if frame_type != "response.create" {
117 return Err(TransformError::InvalidInput {
118 reason: format!("unsupported websocket frame type: {frame_type}"),
119 });
120 }
121 Ok(())
122}
123
124#[derive(Debug, Default)]
126pub struct ResponseWebSocketSseDecoder {
127 decoder: SseDecoder,
128}
129
130impl ResponseWebSocketSseDecoder {
131 pub fn new() -> Self {
132 Self::default()
133 }
134
135 pub fn push(&mut self, chunk: &[u8]) -> Vec<String> {
136 self.decoder
137 .push(chunk)
138 .into_iter()
139 .filter_map(frame_data)
140 .collect()
141 }
142
143 pub fn finish(&mut self) -> Vec<String> {
144 self.decoder
145 .finish()
146 .and_then(frame_data)
147 .into_iter()
148 .collect()
149 }
150}
151
152fn frame_data(frame: crate::transform::common::sse::SseFrame) -> Option<String> {
153 (frame.data.trim() != "[DONE]").then_some(frame.data)
154}
155
156#[cfg(test)]
157mod tests {
158 use serde_json::{Value, json};
159
160 use crate::protocol::{ContentGenerationKind, Operation, OperationKey};
161
162 use super::*;
163
164 fn ctx() -> TransformContext {
165 TransformContext::new(
166 OperationKey::content_generation(
167 Operation::GenerateContent,
168 ContentGenerationKind::OpenAiResponses,
169 ),
170 OperationKey::content_generation(
171 Operation::GenerateContent,
172 ContentGenerationKind::OpenAiResponsesWebSocket,
173 ),
174 )
175 }
176
177 #[test]
178 fn response_create_frame_becomes_streaming_responses_body() {
179 let body = response_create_frame_to_response_body(
180 br#"{"type":"response.create","model":"gpt-test","input":"hi","stream":false,"generate":false,"client_metadata":{"k":"v"},"future_field":{"x":1}}"#,
181 )
182 .unwrap();
183 let value: Value = serde_json::from_slice(&body).unwrap();
184
185 assert_eq!(value.get("type"), None);
186 assert_eq!(value["model"], "gpt-test");
187 assert_eq!(value["input"], "hi");
188 assert_eq!(value["stream"], true);
189 assert_eq!(value["generate"], false);
190 assert_eq!(value["client_metadata"]["k"], "v");
191 assert_eq!(value["future_field"], json!({ "x": 1 }));
192 }
193
194 #[test]
195 fn http_request_becomes_response_create_frame() {
196 let value = http_request_to_ws_request(
197 json!({"model":"gpt-test","input":"hi","stream":true}),
198 &ctx(),
199 )
200 .unwrap();
201
202 assert_eq!(value["type"], "response.create");
203 assert_eq!(value["model"], "gpt-test");
204 assert_eq!(value["stream"], true);
205 }
206
207 #[test]
208 fn rejects_non_create_frame() {
209 let err = response_create_frame_to_response_body(br#"{"type":"session.update"}"#)
210 .expect_err("unsupported frame should fail");
211 assert!(err.to_string().contains("unsupported websocket frame type"));
212 }
213
214 #[test]
215 fn sse_decoder_returns_plain_json_messages() {
216 let mut decoder = ResponseWebSocketSseDecoder::new();
217 let messages = decoder.push(
218 b"event: response.created\ndata: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n",
219 );
220 assert_eq!(messages, vec![r#"{"type":"response.created"}"#]);
221 assert!(decoder.finish().is_empty());
222 }
223}