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};
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    /// Feed one upstream chunk; returns encoded inbound bytes (possibly empty).
110    pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
111        Ok(self.push_detailed(chunk)?.value)
112    }
113
114    /// Feed one chunk and return semantic diagnostics produced by its events.
115    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    /// Flush the trailing frame and emit the inbound terminator.
153    pub fn finish(&mut self) -> Result<Vec<u8>, TransformError> {
154        Ok(self.finish_detailed()?.value)
155    }
156
157    /// Flush the stream and return final semantic diagnostics.
158    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                // Preserve the old tolerant-stream behavior for a recognizable
222                // terminal envelope whose full typed body has evolved beyond
223                // what this build understands. This cold error path still uses
224                // small typed envelopes rather than a dynamic Value tree.
225                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
319/// Encode one converted event in the inbound wire format. Claude and Responses
320/// inbound streams carry named SSE events (missing names fall back to
321/// "message", as before the typed path); chat and Gemini are data-only.
322fn 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
336/// Serialize + encode one Responses event as a named SSE frame.
337fn 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;