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]) -> Result<Vec<String>, TransformError> {
132 Ok(self
133 .decoder
134 .push(chunk)?
135 .into_iter()
136 .filter_map(frame_data)
137 .collect())
138 }
139
140 pub fn finish(&mut self) -> Result<Vec<String>, TransformError> {
141 Ok(self
142 .decoder
143 .finish()?
144 .and_then(frame_data)
145 .into_iter()
146 .collect())
147 }
148}
149
150fn frame_data(frame: crate::transform::common::sse::SseFrame) -> Option<String> {
151 (frame.data.trim() != "[DONE]").then_some(frame.data)
152}
153
154#[cfg(test)]
155mod tests {
156 use serde_json::{Value, json};
157
158 use crate::protocol::{ContentGenerationKind, Operation, OperationKey};
159
160 use super::*;
161
162 fn ctx() -> TransformContext {
163 TransformContext::new(
164 OperationKey::content_generation(
165 Operation::GenerateContent,
166 ContentGenerationKind::OpenAiResponses,
167 ),
168 OperationKey::content_generation(
169 Operation::GenerateContent,
170 ContentGenerationKind::OpenAiResponsesWebSocket,
171 ),
172 )
173 }
174
175 #[test]
176 fn response_create_frame_becomes_streaming_responses_body() {
177 let body = response_create_frame_to_response_body(
178 br#"{"type":"response.create","model":"gpt-test","input":"hi","stream":false,"generate":false,"client_metadata":{"k":"v"},"future_field":{"x":1}}"#,
179 )
180 .unwrap();
181 let value: Value = serde_json::from_slice(&body).unwrap();
182
183 assert_eq!(value.get("type"), None);
184 assert_eq!(value["model"], "gpt-test");
185 assert_eq!(value["input"], "hi");
186 assert_eq!(value["stream"], true);
187 assert_eq!(value["generate"], false);
188 assert_eq!(value["client_metadata"]["k"], "v");
189 assert_eq!(value["future_field"], json!({ "x": 1 }));
190 }
191
192 #[test]
193 fn http_request_becomes_response_create_frame() {
194 let value = http_request_to_ws_request(
195 json!({"model":"gpt-test","input":"hi","stream":true}),
196 &ctx(),
197 )
198 .unwrap();
199
200 assert_eq!(value["type"], "response.create");
201 assert_eq!(value["model"], "gpt-test");
202 assert_eq!(value["stream"], true);
203 }
204
205 #[test]
206 fn rejects_non_create_frame() {
207 let err = response_create_frame_to_response_body(br#"{"type":"session.update"}"#)
208 .expect_err("unsupported frame should fail");
209 assert!(err.to_string().contains("unsupported websocket frame type"));
210 }
211
212 #[test]
213 fn sse_decoder_returns_plain_json_messages() {
214 let mut decoder = ResponseWebSocketSseDecoder::new();
215 let messages = decoder.push(
216 b"event: response.created\ndata: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n",
217 ).unwrap();
218 assert_eq!(messages, vec![r#"{"type":"response.created"}"#]);
219 assert!(decoder.finish().unwrap().is_empty());
220 }
221}