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};
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 pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<u8>, TransformError> {
110 Ok(self.push_detailed(chunk)?.value)
111 }
112
113 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 pub fn finish(&mut self) -> Result<Vec<u8>, TransformError> {
153 Ok(self.finish_detailed()?.value)
154 }
155
156 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
292fn 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
309fn 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;