Skip to main content

gproxy_transform/transform/stream_adapter/
buffered.rs

1use super::{SseDecoder, SseTransformer};
2use crate::protocol::ContentGenerationKind;
3use crate::transform::TransformError;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
6pub struct BufferedDiagnostics {
7    pub decoded_frames: usize,
8}
9
10#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct BufferedAggregation {
12    pub body: Vec<u8>,
13    pub diagnostics: BufferedDiagnostics,
14}
15
16/// Convert a fully-buffered SSE body.
17pub fn convert_buffered(
18    mut transformer: SseTransformer,
19    body: &[u8],
20) -> Result<Vec<u8>, TransformError> {
21    let mut out = transformer.push(body)?;
22    out.extend(transformer.finish()?);
23    Ok(out)
24}
25
26/// Collapse a provider SSE stream into one response JSON of the same wire kind.
27pub fn aggregate_buffered(
28    kind: ContentGenerationKind,
29    sse_body: &[u8],
30) -> Result<BufferedAggregation, TransformError> {
31    use crate::transform::generate_content::stream_to_response as s2r;
32    use ContentGenerationKind as K;
33
34    let mut decoder = SseDecoder::new();
35    let mut frames = decoder.push(sse_body)?;
36    if let Some(tail) = decoder.finish()? {
37        frames.push(tail);
38    }
39    let decoded_frames = frames.len();
40    let datas: Vec<String> = frames
41        .into_iter()
42        .map(|frame| frame.data)
43        .filter(|data| data.trim() != "[DONE]")
44        .collect();
45
46    macro_rules! collapse {
47        ($ty:ty, $aggregate:path) => {{
48            let events = datas
49                .iter()
50                .enumerate()
51                .map(|(index, data)| {
52                    serde_json::from_str::<$ty>(data).map_err(|error| {
53                        TransformError::InvalidInput {
54                            reason: format!("decode buffered stream frame {index}: {error}"),
55                        }
56                    })
57                })
58                .collect::<Result<Vec<_>, _>>()?;
59            serde_json::to_vec(&$aggregate(events.into_iter())).map_err(|error| {
60                TransformError::Serialization {
61                    reason: error.to_string(),
62                }
63            })
64        }};
65    }
66
67    let out = match kind {
68        K::OpenAiResponses | K::OpenAiResponsesWebSocket => collapse!(
69            crate::protocol::openai::ResponseStreamEvent,
70            s2r::openai_responses::response
71        ),
72        K::OpenAiChatCompletions => collapse!(
73            crate::protocol::openai::ChatCompletionChunk,
74            s2r::openai_chat::response
75        ),
76        K::ClaudeMessages => collapse!(
77            crate::protocol::claude::StreamEvent,
78            s2r::claude_messages::response
79        ),
80        K::GeminiGenerateContent => collapse!(
81            crate::protocol::gemini::StreamGenerateContentChunk,
82            s2r::gemini_generate_content::response
83        ),
84        _ => {
85            unreachable!("new non-exhaustive protocol variant requires a lockstep transform update")
86        }
87    }?;
88    Ok(BufferedAggregation {
89        body: out,
90        diagnostics: BufferedDiagnostics { decoded_frames },
91    })
92}