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 response_create_frame_to_response_body(frame: &[u8]) -> Result<Vec<u8>, TransformError> {
68 let mut value: Value =
69 serde_json::from_slice(frame).map_err(|error| TransformError::InvalidInput {
70 reason: format!("decode websocket frame: {error}"),
71 })?;
72 let object = value
73 .as_object_mut()
74 .ok_or_else(|| TransformError::InvalidInput {
75 reason: "websocket frame must be a JSON object".to_owned(),
76 })?;
77 let frame_type = object
78 .remove("type")
79 .and_then(|value| value.as_str().map(str::to_owned))
80 .ok_or_else(|| TransformError::InvalidInput {
81 reason: "websocket frame missing type".to_owned(),
82 })?;
83 if frame_type != "response.create" {
84 return Err(TransformError::InvalidInput {
85 reason: format!("unsupported websocket frame type: {frame_type}"),
86 });
87 }
88
89 object.insert("stream".to_owned(), Value::Bool(true));
90 serde_json::to_vec(&value).map_err(|error| TransformError::Serialization {
91 reason: error.to_string(),
92 })
93}
94
95pub fn validate_response_create_frame(frame: &[u8]) -> Result<(), TransformError> {
97 let value: Value =
98 serde_json::from_slice(frame).map_err(|error| TransformError::InvalidInput {
99 reason: format!("decode websocket frame: {error}"),
100 })?;
101 let object = value
102 .as_object()
103 .ok_or_else(|| TransformError::InvalidInput {
104 reason: "websocket frame must be a JSON object".to_owned(),
105 })?;
106 let frame_type = object
107 .get("type")
108 .and_then(|value| value.as_str())
109 .ok_or_else(|| TransformError::InvalidInput {
110 reason: "websocket frame missing type".to_owned(),
111 })?;
112 if frame_type != "response.create" {
113 return Err(TransformError::InvalidInput {
114 reason: format!("unsupported websocket frame type: {frame_type}"),
115 });
116 }
117 Ok(())
118}
119
120#[derive(Debug, Default)]
122pub struct ResponseWebSocketSseDecoder {
123 decoder: SseDecoder,
124}
125
126impl ResponseWebSocketSseDecoder {
127 pub fn new() -> Self {
128 Self::default()
129 }
130
131 pub fn push(&mut self, chunk: &[u8]) -> Vec<String> {
132 self.decoder
133 .push(chunk)
134 .into_iter()
135 .filter_map(frame_data)
136 .collect()
137 }
138
139 pub fn finish(&mut self) -> Vec<String> {
140 self.decoder
141 .finish()
142 .and_then(frame_data)
143 .into_iter()
144 .collect()
145 }
146}
147
148fn frame_data(frame: crate::transform::common::sse::SseFrame) -> Option<String> {
149 (frame.data.trim() != "[DONE]").then_some(frame.data)
150}
151
152#[cfg(test)]
153mod tests {
154 use serde_json::{Value, json};
155
156 use crate::protocol::{ContentGenerationKind, Operation, OperationKey};
157
158 use super::*;
159
160 fn ctx() -> TransformContext {
161 TransformContext::new(
162 OperationKey::content_generation(
163 Operation::GenerateContent,
164 ContentGenerationKind::OpenAiResponses,
165 ),
166 OperationKey::content_generation(
167 Operation::GenerateContent,
168 ContentGenerationKind::OpenAiResponsesWebSocket,
169 ),
170 )
171 }
172
173 #[test]
174 fn response_create_frame_becomes_streaming_responses_body() {
175 let body = response_create_frame_to_response_body(
176 br#"{"type":"response.create","model":"gpt-test","input":"hi","stream":false,"generate":false,"client_metadata":{"k":"v"},"future_field":{"x":1}}"#,
177 )
178 .unwrap();
179 let value: Value = serde_json::from_slice(&body).unwrap();
180
181 assert_eq!(value.get("type"), None);
182 assert_eq!(value["model"], "gpt-test");
183 assert_eq!(value["input"], "hi");
184 assert_eq!(value["stream"], true);
185 assert_eq!(value["generate"], false);
186 assert_eq!(value["client_metadata"]["k"], "v");
187 assert_eq!(value["future_field"], json!({ "x": 1 }));
188 }
189
190 #[test]
191 fn http_request_becomes_response_create_frame() {
192 let value = http_request_to_ws_request(
193 json!({"model":"gpt-test","input":"hi","stream":true}),
194 &ctx(),
195 )
196 .unwrap();
197
198 assert_eq!(value["type"], "response.create");
199 assert_eq!(value["model"], "gpt-test");
200 assert_eq!(value["stream"], true);
201 }
202
203 #[test]
204 fn rejects_non_create_frame() {
205 let err = response_create_frame_to_response_body(br#"{"type":"session.update"}"#)
206 .expect_err("unsupported frame should fail");
207 assert!(err.to_string().contains("unsupported websocket frame type"));
208 }
209
210 #[test]
211 fn sse_decoder_returns_plain_json_messages() {
212 let mut decoder = ResponseWebSocketSseDecoder::new();
213 let messages = decoder.push(
214 b"event: response.created\ndata: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n",
215 );
216 assert_eq!(messages, vec![r#"{"type":"response.created"}"#]);
217 assert!(decoder.finish().is_empty());
218 }
219}