gproxy_transform/transform/
stream_adapter.rs1mod 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 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 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 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
105fn 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;