gproxy_transform/transform/stream_adapter/
buffered.rs1use 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
16pub 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
26pub 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}