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 serde_json::Value;
8
9use super::common::sse::{SseDecoder, SseFrame};
10use super::{TransformContext, TransformPair, dispatch};
11use crate::protocol::ContentGenerationKind;
12
13use responses::ResponsesStreamState;
14
15pub use buffered::{aggregate_buffered, convert_buffered};
16pub use responses::ResponsesStreamNormalizer;
17pub use synthesize::synthesize_sse;
18
19pub struct SseTransformer {
20    decoder: SseDecoder,
21    /// Reverse pair: upstream kind to inbound kind.
22    pair: TransformPair,
23    ctx: TransformContext,
24    inbound: ContentGenerationKind,
25    responses: Option<ResponsesStreamState>,
26    skipped: u64,
27}
28
29impl SseTransformer {
30    pub fn new(pair: TransformPair, ctx: TransformContext, inbound: ContentGenerationKind) -> Self {
31        Self {
32            decoder: SseDecoder::new(),
33            pair,
34            ctx,
35            inbound,
36            responses: matches!(
37                inbound,
38                ContentGenerationKind::OpenAiResponses
39                    | ContentGenerationKind::OpenAiResponsesWebSocket
40            )
41            .then(ResponsesStreamState::default),
42            skipped: 0,
43        }
44    }
45
46    /// Feed one upstream chunk; returns encoded inbound bytes (possibly empty).
47    pub fn push(&mut self, chunk: &[u8]) -> Vec<u8> {
48        let mut out = Vec::new();
49        for frame in self.decoder.push(chunk) {
50            self.convert_into(frame, &mut out);
51        }
52        out
53    }
54
55    /// Flush the trailing frame and emit the inbound terminator.
56    pub fn finish(&mut self) -> Vec<u8> {
57        let mut out = Vec::new();
58        if let Some(frame) = self.decoder.finish() {
59            self.convert_into(frame, &mut out);
60        }
61        if let Some(responses) = self.responses.as_mut() {
62            for event in responses.finish() {
63                out.extend_from_slice(encode_frame(self.inbound, &event).as_bytes());
64            }
65        }
66        if self.inbound == ContentGenerationKind::OpenAiChatCompletions {
67            out.extend_from_slice(b"data: [DONE]\n\n");
68        }
69        if self.skipped > 0 {
70            tracing::warn!(
71                skipped = self.skipped,
72                "stream transform skipped unconvertible frames"
73            );
74        }
75        out
76    }
77
78    fn convert_into(&mut self, frame: SseFrame, out: &mut Vec<u8>) {
79        if frame.data.trim() == "[DONE]" {
80            return;
81        }
82        let event: Value = match serde_json::from_str(&frame.data) {
83            Ok(value) => value,
84            Err(_) => {
85                self.skipped += 1;
86                return;
87            }
88        };
89        match dispatch::stream_event_value(self.pair, &self.ctx, event) {
90            Ok(converted) => {
91                let events = if let Some(responses) = self.responses.as_mut() {
92                    responses.push(converted)
93                } else {
94                    vec![converted]
95                };
96                for event in events {
97                    out.extend_from_slice(encode_frame(self.inbound, &event).as_bytes());
98                }
99            }
100            Err(_) => self.skipped += 1,
101        }
102    }
103}
104
105/// Encode one converted event in the inbound wire format.
106fn encode_frame(kind: ContentGenerationKind, value: &Value) -> String {
107    use ContentGenerationKind as K;
108    let data = value.to_string();
109    match kind {
110        K::ClaudeMessages | K::OpenAiResponses | K::OpenAiResponsesWebSocket => {
111            let name = value
112                .get("type")
113                .and_then(Value::as_str)
114                .unwrap_or("message");
115            SseFrame::event(name, data).encode()
116        }
117        K::OpenAiChatCompletions | K::GeminiGenerateContent => SseFrame::data(data).encode(),
118    }
119}
120
121#[cfg(test)]
122mod tests;