Skip to main content

gproxy_transform/envelope/collector/
mod.rs

1use bytes::Bytes;
2use gproxy_protocol::{ContentGenerationKind, claude as claude_wire, openai};
3
4use self::chat::ChatCollector;
5use self::claude::ClaudeCollector;
6use self::gemini::GeminiCollector;
7use self::responses::ResponsesCollector;
8
9mod chat;
10mod claude;
11mod gemini;
12mod responses;
13
14use super::{SseDecoder, SseFrame};
15use crate::TransformError;
16
17pub enum BufferedResponse {
18    OpenAiChat(Box<openai::ChatCompletionResponse>),
19    OpenAiResponses(Box<openai::ResponseObject>),
20    Claude(Box<claude_wire::CreateMessageResponseBody>),
21    Gemini(Box<gproxy_protocol::gemini::GenerateContentResponse>),
22}
23
24impl BufferedResponse {
25    pub fn into_bytes(self) -> Result<Bytes, TransformError> {
26        Ok(Bytes::from(match self {
27            Self::OpenAiChat(response) => serde_json::to_vec(&response)?,
28            Self::OpenAiResponses(response) => serde_json::to_vec(&response)?,
29            Self::Claude(response) => serde_json::to_vec(&response)?,
30            Self::Gemini(response) => serde_json::to_vec(&response)?,
31        }))
32    }
33}
34
35pub struct ResponseCollector {
36    decoder: SseDecoder,
37    state: Collector,
38}
39
40enum Collector {
41    Chat(Box<ChatCollector>),
42    Responses(Box<ResponsesCollector>),
43    Claude(Box<ClaudeCollector>),
44    Gemini(Box<GeminiCollector>),
45}
46
47impl ResponseCollector {
48    pub fn new(kind: ContentGenerationKind) -> Result<Self, TransformError> {
49        let state = match kind {
50            ContentGenerationKind::OpenAiChat => Collector::Chat(Box::default()),
51            ContentGenerationKind::OpenAiResponses
52            | ContentGenerationKind::OpenAiResponsesWebSocket => {
53                Collector::Responses(Box::default())
54            }
55            ContentGenerationKind::ClaudeMessages => Collector::Claude(Box::default()),
56            ContentGenerationKind::GeminiGenerateContent => Collector::Gemini(Box::default()),
57            #[cfg(not(feature = "exhaustive"))]
58            _ => {
59                return Err(crate::TransformError::unsupported(
60                    "protocol enum",
61                    "unrecognized external variant",
62                ));
63            }
64        };
65        Ok(Self {
66            decoder: SseDecoder::default(),
67            state,
68        })
69    }
70
71    pub fn push(&mut self, chunk: Bytes) -> Result<(), TransformError> {
72        for frame in self.decoder.push(&chunk)? {
73            self.state.frame(frame)?;
74        }
75        Ok(())
76    }
77
78    pub fn is_complete(&self) -> bool {
79        self.state.is_complete()
80    }
81
82    pub fn claude_has_output(&self) -> bool {
83        matches!(&self.state, Collector::Claude(state) if state.has_output())
84    }
85
86    pub fn claude_has_open_tool(&self) -> bool {
87        matches!(&self.state, Collector::Claude(state) if !state.open_tools.is_empty())
88    }
89
90    pub fn finish(mut self) -> Result<BufferedResponse, TransformError> {
91        if let Some(frame) = self.decoder.finish()? {
92            self.state.frame(frame)?;
93        }
94        self.state.finish()
95    }
96}
97
98impl Collector {
99    fn frame(&mut self, frame: SseFrame) -> Result<(), TransformError> {
100        match self {
101            Self::Chat(state) => state.frame(frame),
102            Self::Responses(state) => state.frame(frame),
103            Self::Claude(state) => state.frame(frame),
104            Self::Gemini(state) => state.frame(frame),
105        }
106    }
107
108    fn is_complete(&self) -> bool {
109        match self {
110            Self::Chat(state) => state.is_complete(),
111            Self::Responses(state) => state.response.is_some(),
112            Self::Claude(state) => state.complete,
113            Self::Gemini(state) => state.is_complete(),
114        }
115    }
116
117    fn finish(self) -> Result<BufferedResponse, TransformError> {
118        match self {
119            Self::Chat(state) => state
120                .finish()
121                .map(Box::new)
122                .map(BufferedResponse::OpenAiChat),
123            Self::Responses(state) => state
124                .finish()
125                .map(Box::new)
126                .map(BufferedResponse::OpenAiResponses),
127            Self::Claude(state) => state.finish().map(Box::new).map(BufferedResponse::Claude),
128            Self::Gemini(state) => state.finish().map(Box::new).map(BufferedResponse::Gemini),
129        }
130    }
131}