Skip to main content

gproxy_transform/transform/
stream_adapter.rs

1//! Runtime SSE adaptation for cross-protocol content-generation streams.
2
3mod 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};
13
14use responses::ResponsesStreamState;
15
16pub use buffered::{
17    BufferedAggregation, BufferedDiagnostics, aggregate_buffered, convert_buffered,
18};
19pub use responses::ResponsesStreamNormalizer;
20pub use synthesize::synthesize_sse;
21
22pub struct SseTransformer {
23    decoder: SseDecoder,
24    converter: dispatch::StreamConverter,
25    source: ContentGenerationKind,
26    inbound: ContentGenerationKind,
27    responses: Option<ResponsesStreamState>,
28    error_mode: StreamErrorMode,
29    require_terminal: bool,
30    terminal_seen: bool,
31    failed: bool,
32    finished: bool,
33    skipped: u64,
34    semantic_diagnostics: Vec<TransformDiagnostic>,
35}
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
38pub enum StreamErrorMode {
39    #[default]
40    Strict,
41    SkipInvalid,
42}
43
44#[derive(Debug, Clone, Copy, PartialEq, Eq)]
45pub struct StreamOptions {
46    pub limits: SseLimits,
47    pub error_mode: StreamErrorMode,
48    pub require_terminal: bool,
49}
50
51impl Default for StreamOptions {
52    fn default() -> Self {
53        Self {
54            limits: SseLimits::default(),
55            error_mode: StreamErrorMode::Strict,
56            require_terminal: true,
57        }
58    }
59}
60
61#[derive(Debug, Clone, PartialEq, Eq, Default)]
62pub struct StreamDiagnostics {
63    pub skipped_frames: u64,
64    pub semantic_diagnostics: Vec<TransformDiagnostic>,
65}
66
67impl SseTransformer {
68    pub fn new(pair: TransformPair, ctx: TransformContext) -> Result<Self, TransformError> {
69        Self::with_options(pair, ctx, StreamOptions::default())
70    }
71
72    pub fn with_options(
73        pair: TransformPair,
74        ctx: TransformContext,
75        options: StreamOptions,
76    ) -> Result<Self, TransformError> {
77        let OperationKind::ContentGeneration(source) = ctx.source.kind() else {
78            return Err(TransformError::InvalidInput {
79                reason: "stream source is not a content-generation operation".to_owned(),
80            });
81        };
82        let OperationKind::ContentGeneration(inbound) = ctx.target.kind() else {
83            return Err(TransformError::InvalidInput {
84                reason: "stream target is not a content-generation operation".to_owned(),
85            });
86        };
87        Ok(Self {
88            decoder: SseDecoder::with_limits(options.limits),
89            converter: dispatch::StreamConverter::new(pair, ctx)?,
90            source,
91            inbound,
92            responses: matches!(
93                inbound,
94                ContentGenerationKind::OpenAiResponses
95                    | ContentGenerationKind::OpenAiResponsesWebSocket
96            )
97            .then(ResponsesStreamState::default),
98            error_mode: options.error_mode,
99            require_terminal: options.require_terminal,
100            terminal_seen: false,
101            failed: false,
102            finished: false,
103            skipped: 0,
104            semantic_diagnostics: Vec::new(),
105        })
106    }
107
108    /// Feed one upstream chunk; returns encoded inbound bytes (possibly empty).
109    pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
110        Ok(self.push_detailed(chunk)?.value)
111    }
112
113    /// Feed one chunk and return semantic diagnostics produced by its events.
114    pub fn push_detailed(
115        &mut self,
116        chunk: &[u8],
117    ) -> Result<TransformOutput<Vec<u8>>, TransformError> {
118        let diagnostic_start = self.semantic_diagnostics.len();
119        let value = self.push_value(chunk)?;
120        Ok(TransformOutput::new(
121            value,
122            self.semantic_diagnostics[diagnostic_start..].to_vec(),
123        ))
124    }
125
126    fn push_value(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
127        if self.finished {
128            return Err(TransformError::InvalidInput {
129                reason: "cannot push after stream finish".to_owned(),
130            });
131        }
132        if self.failed {
133            return Err(TransformError::InvalidInput {
134                reason: "stream is failed after an earlier conversion error".to_owned(),
135            });
136        }
137        let mut out = Vec::new();
138        let frames = self
139            .decoder
140            .push(chunk)
141            .inspect_err(|_| self.failed = true)?;
142        for frame in frames {
143            if let Err(error) = self.convert_into(frame, &mut out) {
144                self.failed = true;
145                return Err(error);
146            }
147        }
148        Ok(out)
149    }
150
151    /// Flush the trailing frame and emit the inbound terminator.
152    pub fn finish(&mut self) -> Result<Vec<u8>, TransformError> {
153        Ok(self.finish_detailed()?.value)
154    }
155
156    /// Flush the stream and return final semantic diagnostics.
157    pub fn finish_detailed(&mut self) -> Result<TransformOutput<Vec<u8>>, TransformError> {
158        let diagnostic_start = self.semantic_diagnostics.len();
159        let value = self.finish_value()?;
160        Ok(TransformOutput::new(
161            value,
162            self.semantic_diagnostics[diagnostic_start..].to_vec(),
163        ))
164    }
165
166    fn finish_value(&mut self) -> Result<Vec<u8>, TransformError> {
167        if self.finished {
168            return Ok(Vec::new());
169        }
170        if self.failed {
171            return Err(TransformError::InvalidInput {
172                reason: "cannot finish a stream after a conversion error".to_owned(),
173            });
174        }
175        let mut out = Vec::new();
176        if let Some(frame) = self.decoder.finish().inspect_err(|_| self.failed = true)?
177            && let Err(error) = self.convert_into(frame, &mut out)
178        {
179            self.failed = true;
180            return Err(error);
181        }
182        if self.require_terminal && !self.terminal_seen {
183            self.failed = true;
184            return Err(TransformError::UnexpectedEof {
185                reason: "upstream ended before a protocol terminal event",
186            });
187        }
188        let converted = self.converter.finish_detailed()?;
189        self.semantic_diagnostics.extend(converted.diagnostics);
190        for event in converted.value {
191            self.encode_converted(event, &mut out)?;
192        }
193        if let Some(responses) = self.responses.as_mut() {
194            for event in responses.finish() {
195                encode_responses_event(&event, &mut out)?;
196            }
197        }
198        if self.inbound == ContentGenerationKind::OpenAiChatCompletions {
199            out.extend_from_slice(b"data: [DONE]\n\n");
200        }
201        self.finished = true;
202        Ok(out)
203    }
204
205    pub fn diagnostics(&self) -> StreamDiagnostics {
206        StreamDiagnostics {
207            skipped_frames: self.skipped,
208            semantic_diagnostics: self.semantic_diagnostics.clone(),
209        }
210    }
211
212    fn convert_into(&mut self, frame: SseFrame, out: &mut Vec<u8>) -> Result<(), TransformError> {
213        if frame.data.trim() == "[DONE]" {
214            self.terminal_seen = true;
215            return Ok(());
216        }
217        self.terminal_seen |= is_terminal_event(self.source, &frame.data);
218        let events = match self.converter.push_detailed(&frame.data) {
219            Ok(events) => events,
220            Err(_) if self.error_mode == StreamErrorMode::SkipInvalid => {
221                self.skipped += 1;
222                return Ok(());
223            }
224            Err(error) => return Err(error),
225        };
226        self.semantic_diagnostics.extend(events.diagnostics);
227        for event in events.value {
228            self.encode_converted(event, out)?;
229        }
230        Ok(())
231    }
232
233    fn encode_converted(
234        &mut self,
235        event: dispatch::StreamEventOut,
236        out: &mut Vec<u8>,
237    ) -> Result<(), TransformError> {
238        match event {
239            dispatch::StreamEventOut::Encoded { event, data } => {
240                encode_frame(self.inbound, event.as_deref(), &data, out);
241            }
242            dispatch::StreamEventOut::Responses(event) => {
243                if let Some(responses) = self.responses.as_mut() {
244                    for event in responses.push(*event) {
245                        encode_responses_event(&event, out)?;
246                    }
247                } else {
248                    encode_responses_event(&event, out)?;
249                }
250            }
251        }
252        Ok(())
253    }
254}
255
256fn is_terminal_event(kind: ContentGenerationKind, data: &str) -> bool {
257    let Ok(value) = serde_json::from_str::<serde_json::Value>(data) else {
258        return false;
259    };
260    match kind {
261        ContentGenerationKind::OpenAiChatCompletions => false,
262        ContentGenerationKind::ClaudeMessages => matches!(
263            value.get("type").and_then(serde_json::Value::as_str),
264            Some("message_stop" | "error")
265        ),
266        ContentGenerationKind::OpenAiResponses
267        | ContentGenerationKind::OpenAiResponsesWebSocket => matches!(
268            value.get("type").and_then(serde_json::Value::as_str),
269            Some("response.completed" | "response.incomplete" | "response.failed" | "error")
270        ),
271        ContentGenerationKind::GeminiGenerateContent => {
272            value
273                .get("candidates")
274                .and_then(serde_json::Value::as_array)
275                .is_some_and(|candidates| {
276                    candidates.iter().any(|candidate| {
277                        candidate
278                            .get("finishReason")
279                            .is_some_and(|reason| !reason.is_null())
280                    })
281                })
282                || value
283                    .pointer("/promptFeedback/blockReason")
284                    .is_some_and(|reason| !reason.is_null())
285        }
286        _ => {
287            unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
288        }
289    }
290}
291
292/// Encode one converted event in the inbound wire format. Claude and Responses
293/// inbound streams carry named SSE events (missing names fall back to
294/// "message", as before the typed path); chat and Gemini are data-only.
295fn encode_frame(kind: ContentGenerationKind, event: Option<&str>, data: &str, out: &mut Vec<u8>) {
296    use ContentGenerationKind as K;
297    let frame = match kind {
298        K::ClaudeMessages | K::OpenAiResponses | K::OpenAiResponsesWebSocket => {
299            SseFrame::event(event.unwrap_or("message"), data)
300        }
301        K::OpenAiChatCompletions | K::GeminiGenerateContent => SseFrame::data(data),
302        _ => {
303            unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
304        }
305    };
306    out.extend_from_slice(frame.encode().as_bytes());
307}
308
309/// Serialize + encode one Responses event as a named SSE frame.
310fn encode_responses_event(
311    event: &ResponseStreamEvent,
312    out: &mut Vec<u8>,
313) -> Result<(), TransformError> {
314    let data = serde_json::to_string(event).map_err(|error| TransformError::Serialization {
315        reason: error.to_string(),
316    })?;
317    let frame = SseFrame::event(event.event_name().unwrap_or("message"), data);
318    out.extend_from_slice(frame.encode().as_bytes());
319    Ok(())
320}
321
322#[cfg(test)]
323mod tests;