gproxy_transform/transform/
stream_adapter.rs1mod buffered;
4mod responses;
5mod synthesize;
6
7use super::common::sse::{SseDecoder, SseFrame, SseLimits};
8use super::{
9 TransformContext, TransformDiagnostic, TransformError, TransformOutput, TransformPair, dispatch,
10};
11use crate::protocol::openai::ResponseStreamEvent;
12use crate::protocol::{ContentGenerationKind, OperationKind};
13use serde::Deserialize;
14
15use responses::ResponsesStreamState;
16
17pub use buffered::{
18 BufferedAggregation, BufferedDiagnostics, aggregate_buffered, convert_buffered,
19};
20pub use responses::ResponsesStreamNormalizer;
21pub use synthesize::synthesize_sse;
22
23pub struct SseTransformer {
24 decoder: SseDecoder,
25 converter: dispatch::StreamConverter,
26 source: ContentGenerationKind,
27 inbound: ContentGenerationKind,
28 responses: Option<ResponsesStreamState>,
29 error_mode: StreamErrorMode,
30 require_terminal: bool,
31 terminal_seen: bool,
32 failed: bool,
33 finished: bool,
34 skipped: u64,
35 semantic_diagnostics: Vec<TransformDiagnostic>,
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
39pub enum StreamErrorMode {
40 #[default]
41 Strict,
42 SkipInvalid,
43}
44
45#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46pub struct StreamOptions {
47 pub limits: SseLimits,
48 pub error_mode: StreamErrorMode,
49 pub require_terminal: bool,
50}
51
52impl Default for StreamOptions {
53 fn default() -> Self {
54 Self {
55 limits: SseLimits::default(),
56 error_mode: StreamErrorMode::Strict,
57 require_terminal: true,
58 }
59 }
60}
61
62#[derive(Debug, Clone, PartialEq, Eq, Default)]
63pub struct StreamDiagnostics {
64 pub skipped_frames: u64,
65 pub semantic_diagnostics: Vec<TransformDiagnostic>,
66}
67
68impl SseTransformer {
69 pub fn new(pair: TransformPair, ctx: TransformContext) -> Result<Self, TransformError> {
70 Self::with_options(pair, ctx, StreamOptions::default())
71 }
72
73 pub fn with_options(
74 pair: TransformPair,
75 ctx: TransformContext,
76 options: StreamOptions,
77 ) -> Result<Self, TransformError> {
78 let OperationKind::ContentGeneration(source) = ctx.source.kind() else {
79 return Err(TransformError::InvalidInput {
80 reason: "stream source is not a content-generation operation".to_owned(),
81 });
82 };
83 let OperationKind::ContentGeneration(inbound) = ctx.target.kind() else {
84 return Err(TransformError::InvalidInput {
85 reason: "stream target is not a content-generation operation".to_owned(),
86 });
87 };
88 Ok(Self {
89 decoder: SseDecoder::with_limits(options.limits),
90 converter: dispatch::StreamConverter::new(pair, ctx)?,
91 source,
92 inbound,
93 responses: matches!(
94 inbound,
95 ContentGenerationKind::OpenAiResponses
96 | ContentGenerationKind::OpenAiResponsesWebSocket
97 )
98 .then(ResponsesStreamState::default),
99 error_mode: options.error_mode,
100 require_terminal: options.require_terminal,
101 terminal_seen: false,
102 failed: false,
103 finished: false,
104 skipped: 0,
105 semantic_diagnostics: Vec::new(),
106 })
107 }
108
109 pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
111 Ok(self.push_detailed(chunk)?.value)
112 }
113
114 pub fn push_detailed(
116 &mut self,
117 chunk: &[u8],
118 ) -> Result<TransformOutput<Vec<u8>>, TransformError> {
119 let diagnostic_start = self.semantic_diagnostics.len();
120 let value = self.push_value(chunk)?;
121 Ok(TransformOutput::new(
122 value,
123 self.semantic_diagnostics[diagnostic_start..].to_vec(),
124 ))
125 }
126
127 fn push_value(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
128 if self.finished {
129 return Err(TransformError::InvalidInput {
130 reason: "cannot push after stream finish".to_owned(),
131 });
132 }
133 if self.failed {
134 return Err(TransformError::InvalidInput {
135 reason: "stream is failed after an earlier conversion error".to_owned(),
136 });
137 }
138 let mut out = Vec::new();
139 let frames = self
140 .decoder
141 .push(chunk)
142 .inspect_err(|_| self.failed = true)?;
143 for frame in frames {
144 if let Err(error) = self.convert_into(frame, &mut out) {
145 self.failed = true;
146 return Err(error);
147 }
148 }
149 Ok(out)
150 }
151
152 pub fn finish(&mut self) -> Result<Vec<u8>, TransformError> {
154 Ok(self.finish_detailed()?.value)
155 }
156
157 pub fn finish_detailed(&mut self) -> Result<TransformOutput<Vec<u8>>, TransformError> {
159 let diagnostic_start = self.semantic_diagnostics.len();
160 let value = self.finish_value()?;
161 Ok(TransformOutput::new(
162 value,
163 self.semantic_diagnostics[diagnostic_start..].to_vec(),
164 ))
165 }
166
167 fn finish_value(&mut self) -> Result<Vec<u8>, TransformError> {
168 if self.finished {
169 return Ok(Vec::new());
170 }
171 if self.failed {
172 return Err(TransformError::InvalidInput {
173 reason: "cannot finish a stream after a conversion error".to_owned(),
174 });
175 }
176 let mut out = Vec::new();
177 if let Some(frame) = self.decoder.finish().inspect_err(|_| self.failed = true)?
178 && let Err(error) = self.convert_into(frame, &mut out)
179 {
180 self.failed = true;
181 return Err(error);
182 }
183 if self.require_terminal && !self.terminal_seen {
184 self.failed = true;
185 return Err(TransformError::UnexpectedEof {
186 reason: "upstream ended before a protocol terminal event",
187 });
188 }
189 let converted = self.converter.finish_detailed()?;
190 self.semantic_diagnostics.extend(converted.diagnostics);
191 for event in converted.value {
192 self.encode_converted(event, &mut out)?;
193 }
194 if let Some(responses) = self.responses.as_mut() {
195 for event in responses.finish() {
196 encode_responses_event(&event, &mut out)?;
197 }
198 }
199 if self.inbound == ContentGenerationKind::OpenAiChatCompletions {
200 out.extend_from_slice(b"data: [DONE]\n\n");
201 }
202 self.finished = true;
203 Ok(out)
204 }
205
206 pub fn diagnostics(&self) -> StreamDiagnostics {
207 StreamDiagnostics {
208 skipped_frames: self.skipped,
209 semantic_diagnostics: self.semantic_diagnostics.clone(),
210 }
211 }
212
213 fn convert_into(&mut self, frame: SseFrame, out: &mut Vec<u8>) -> Result<(), TransformError> {
214 if frame.data.trim() == "[DONE]" {
215 self.terminal_seen = true;
216 return Ok(());
217 }
218 let events = match self.converter.push_detailed_with_status(&frame.data) {
219 Ok(events) => events,
220 Err(_) if self.error_mode == StreamErrorMode::SkipInvalid => {
221 self.terminal_seen |= terminal_hint(self.source, &frame.data);
226 self.skipped += 1;
227 return Ok(());
228 }
229 Err(error) => return Err(error),
230 };
231 self.terminal_seen |= events.value.terminal;
232 self.semantic_diagnostics.extend(events.diagnostics);
233 for event in events.value.events {
234 self.encode_converted(event, out)?;
235 }
236 Ok(())
237 }
238
239 fn encode_converted(
240 &mut self,
241 event: dispatch::StreamEventOut,
242 out: &mut Vec<u8>,
243 ) -> Result<(), TransformError> {
244 match event {
245 dispatch::StreamEventOut::Encoded { event, data } => {
246 encode_frame(self.inbound, event.as_deref(), &data, out);
247 }
248 dispatch::StreamEventOut::Responses(event) => {
249 if let Some(responses) = self.responses.as_mut() {
250 for event in responses.push(*event) {
251 encode_responses_event(&event, out)?;
252 }
253 } else {
254 encode_responses_event(&event, out)?;
255 }
256 }
257 }
258 Ok(())
259 }
260}
261
262fn terminal_hint(kind: ContentGenerationKind, data: &str) -> bool {
263 #[derive(Deserialize)]
264 struct TaggedEnvelope {
265 #[serde(rename = "type")]
266 type_: String,
267 }
268
269 #[derive(Deserialize)]
270 #[serde(rename_all = "camelCase")]
271 struct GeminiEnvelope {
272 #[serde(default)]
273 candidates: Vec<GeminiCandidateEnvelope>,
274 prompt_feedback: Option<GeminiPromptFeedbackEnvelope>,
275 }
276
277 #[derive(Deserialize)]
278 #[serde(rename_all = "camelCase")]
279 struct GeminiCandidateEnvelope {
280 finish_reason: Option<serde::de::IgnoredAny>,
281 }
282
283 #[derive(Deserialize)]
284 #[serde(rename_all = "camelCase")]
285 struct GeminiPromptFeedbackEnvelope {
286 block_reason: Option<serde::de::IgnoredAny>,
287 }
288
289 match kind {
290 ContentGenerationKind::OpenAiChatCompletions => false,
291 ContentGenerationKind::ClaudeMessages => serde_json::from_str::<TaggedEnvelope>(data)
292 .is_ok_and(|event| matches!(event.type_.as_str(), "message_stop" | "error")),
293 ContentGenerationKind::OpenAiResponses
294 | ContentGenerationKind::OpenAiResponsesWebSocket => {
295 serde_json::from_str::<TaggedEnvelope>(data).is_ok_and(|event| {
296 matches!(
297 event.type_.as_str(),
298 "response.completed" | "response.incomplete" | "response.failed" | "error"
299 )
300 })
301 }
302 ContentGenerationKind::GeminiGenerateContent => {
303 serde_json::from_str::<GeminiEnvelope>(data).is_ok_and(|event| {
304 event
305 .candidates
306 .iter()
307 .any(|candidate| candidate.finish_reason.is_some())
308 || event
309 .prompt_feedback
310 .is_some_and(|feedback| feedback.block_reason.is_some())
311 })
312 }
313 _ => {
314 unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
315 }
316 }
317}
318
319fn encode_frame(kind: ContentGenerationKind, event: Option<&str>, data: &str, out: &mut Vec<u8>) {
323 use ContentGenerationKind as K;
324 let frame = match kind {
325 K::ClaudeMessages | K::OpenAiResponses | K::OpenAiResponsesWebSocket => {
326 SseFrame::event(event.unwrap_or("message"), data)
327 }
328 K::OpenAiChatCompletions | K::GeminiGenerateContent => SseFrame::data(data),
329 _ => {
330 unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
331 }
332 };
333 out.extend_from_slice(frame.encode().as_bytes());
334}
335
336fn encode_responses_event(
338 event: &ResponseStreamEvent,
339 out: &mut Vec<u8>,
340) -> Result<(), TransformError> {
341 let data = serde_json::to_string(event).map_err(|error| TransformError::Serialization {
342 reason: error.to_string(),
343 })?;
344 let frame = SseFrame::event(event.event_name().unwrap_or("message"), data);
345 out.extend_from_slice(frame.encode().as_bytes());
346 Ok(())
347}
348
349#[cfg(test)]
350mod tests;