Skip to main content

gemini_rs/
stream.rs

1use crate::client::{Formatter, Request};
2use std::{
3    ops::{Deref, DerefMut},
4    pin::Pin,
5    task::Poll,
6};
7
8use bytes::Bytes;
9use futures::Stream;
10use reqwest::Method;
11use serde::ser::Error as _;
12
13use crate::{
14    Error, Result,
15    client::{BASE_URI, GenerateContent, Route},
16    types,
17};
18
19pub struct StreamGenerateContent(pub(crate) GenerateContent);
20
21impl StreamGenerateContent {
22    pub fn new(model: &str) -> Self {
23        Self(GenerateContent::new(model.into()))
24    }
25}
26
27impl Deref for Route<StreamGenerateContent> {
28    type Target = GenerateContent;
29
30    fn deref(&self) -> &Self::Target {
31        &self.kind.0
32    }
33}
34
35impl DerefMut for Route<StreamGenerateContent> {
36    fn deref_mut(&mut self) -> &mut Self::Target {
37        &mut self.kind.0
38    }
39}
40
41impl Route<StreamGenerateContent> {
42    pub async fn stream(self) -> std::result::Result<RouteStream<StreamGenerateContent>, String> {
43        let url = format!("{BASE_URI}/{}", self);
44        let body = self.kind.body().clone();
45        let mut request = self
46            .client
47            .reqwest
48            .request(StreamGenerateContent::METHOD, url);
49
50        if let Some(body) = body {
51            request = request.json(&body);
52        }
53
54        let response = request.send().await.map_err(|e| e.to_string())?;
55        let stream = response.bytes_stream();
56
57        Ok(RouteStream {
58            phantom: std::marker::PhantomData,
59            stream: Box::pin(stream),
60            buffer: Vec::new(),
61            pos: 0,
62            state: ParseState::CannotAdvance,
63        })
64    }
65}
66
67pub struct RouteStream<T> {
68    phantom: std::marker::PhantomData<T>,
69    stream: Pin<Box<dyn Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Send>>,
70    buffer: Vec<u8>,
71    pos: usize, // A cursor into the buffer.
72    state: ParseState,
73}
74
75#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76enum ParseState {
77    CannotAdvance,
78    ReadingChars,
79    ReadingValue,
80    Finished,
81}
82
83#[derive(Debug)]
84enum ParseOutcome {
85    Ok(Option<types::Response>),
86    Err(serde_json::Error),
87    Eof,
88}
89
90impl RouteStream<StreamGenerateContent> {
91    fn next_char_pos(&self) -> Option<usize> {
92        self.buffer[self.pos..]
93            .iter()
94            .position(|&b| !b.is_ascii_whitespace())
95            .map(|p| self.pos + p)
96    }
97
98    fn advance_next_char(&mut self) -> Option<u8> {
99        self.pos = self.next_char_pos().unwrap_or(self.buffer.len());
100        self.buffer.get(self.pos).copied()
101    }
102
103    fn current_char(&self) -> Option<u8> {
104        self.buffer.get(self.pos).copied()
105    }
106
107    fn is_bridge_char(&self) -> bool {
108        matches!(self.current_char(), Some(b'[') | Some(b','))
109    }
110
111    fn parse_chunk(&mut self) -> ParseOutcome {
112        let mut de = serde_json::Deserializer::from_slice(&self.buffer[self.pos..])
113            .into_iter::<types::Response>();
114        match de.next() {
115            Some(Ok(value)) => {
116                self.pos += de.byte_offset();
117                ParseOutcome::Ok(Some(value))
118            }
119            Some(Err(e)) if e.is_eof() => ParseOutcome::Eof,
120            Some(Err(e)) => ParseOutcome::Err(e),
121            None => ParseOutcome::Ok(None), // No more objects to read.
122        }
123    }
124
125    fn try_parse_next(&mut self) -> Option<ParseOutcome> {
126        match self.state {
127            ParseState::CannotAdvance => None, // nothing to read
128            ParseState::ReadingChars => {
129                self.advance_next_char();
130                if self.is_bridge_char() {
131                    self.pos += 1; // Move past this '[' or ','
132                    self.state = ParseState::ReadingValue;
133                    None
134                } else if let Some(b']') = self.current_char() {
135                    // If we hit a ']', we can finish reading.
136                    self.state = ParseState::Finished;
137                    Some(ParseOutcome::Ok(None))
138                } else {
139                    None
140                }
141            }
142            ParseState::ReadingValue => {
143                self.advance_next_char();
144                // Deserialize one object from our current position.
145                let outcome = self.parse_chunk();
146                match &outcome {
147                    ParseOutcome::Ok(Some(_)) => {
148                        self.state = ParseState::ReadingChars;
149                    }
150                    ParseOutcome::Ok(None) | ParseOutcome::Err(_) => {
151                        self.state = ParseState::Finished;
152                    }
153                    ParseOutcome::Eof => {}
154                };
155                Some(outcome)
156            }
157            ParseState::Finished => None,
158        }
159    }
160}
161
162impl Stream for RouteStream<StreamGenerateContent> {
163    type Item = Result<types::Response>;
164
165    fn poll_next(
166        mut self: Pin<&mut Self>,
167        cx: &mut std::task::Context<'_>,
168    ) -> Poll<Option<Self::Item>> {
169        loop {
170            // Housekeeping: drain the buffer if we've processed a lot.
171            if self.pos > 2048 {
172                let this_pos = self.pos;
173                self.buffer.drain(..this_pos);
174                self.pos = 0;
175            }
176
177            if let Some(outcome) = self.try_parse_next() {
178                match outcome {
179                    ParseOutcome::Ok(Some(response)) => return Poll::Ready(Some(Ok(response))),
180                    ParseOutcome::Ok(None) if self.state == ParseState::Finished => {
181                        return Poll::Ready(None);
182                    }
183                    ParseOutcome::Err(error) => return Poll::Ready(Some(Err(Error::Serde(error)))),
184                    ParseOutcome::Eof => {} // Continue to read more data.
185                    _ => {}
186                }
187            };
188
189            // If we fell through, we need more data. Poll the underlying stream.
190            match self.stream.as_mut().poll_next(cx) {
191                Poll::Ready(Some(Ok(bytes))) => {
192                    if self.buffer.is_empty() && !bytes.is_empty() {
193                        self.state = ParseState::ReadingChars;
194                    }
195                    self.buffer.extend_from_slice(&bytes);
196                    continue; // Loop again to process new data.
197                }
198                Poll::Pending => return Poll::Pending,
199                Poll::Ready(Some(Err(e))) => {
200                    self.state = ParseState::Finished;
201                    return Poll::Ready(Some(Err(Error::Http(e))));
202                }
203                Poll::Ready(None) => {
204                    // Underlying stream ended. Check if we're in a clean state.
205                    if self.state != ParseState::Finished && self.pos < self.buffer.len() {
206                        let msg =
207                            format!("stream ended with unparsed data in state {:?}", self.state);
208                        return Poll::Ready(Some(Err(serde_json::Error::custom(msg).into())));
209                    }
210                    self.state = ParseState::Finished;
211                    return Poll::Ready(None);
212                }
213            }
214        }
215    }
216}
217
218impl Request for StreamGenerateContent {
219    type Model = types::Response;
220    type Body = types::GenerateContent;
221
222    const METHOD: Method = Method::POST;
223
224    fn format_uri(&self, fmt: &mut Formatter<'_, '_>) -> std::fmt::Result {
225        fmt.write_str("v1beta/")?;
226        fmt.write_str("models/")?;
227        fmt.write_str(&self.0.model)?;
228        fmt.write_str(":streamGenerateContent")
229    }
230
231    fn body(&self) -> Option<Self::Body> {
232        Some(self.0.body.clone())
233    }
234}
235
236impl std::fmt::Display for StreamGenerateContent {
237    fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
238        let mut fmt = Formatter::new(fmt);
239        self.format_uri(&mut fmt)?;
240        fmt.write_query_param("key", &self.0.model)
241    }
242}