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, 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), }
123 }
124
125 fn try_parse_next(&mut self) -> Option<ParseOutcome> {
126 match self.state {
127 ParseState::CannotAdvance => None, ParseState::ReadingChars => {
129 self.advance_next_char();
130 if self.is_bridge_char() {
131 self.pos += 1; self.state = ParseState::ReadingValue;
133 None
134 } else if let Some(b']') = self.current_char() {
135 self.state = ParseState::Finished;
137 Some(ParseOutcome::Ok(None))
138 } else {
139 None
140 }
141 }
142 ParseState::ReadingValue => {
143 self.advance_next_char();
144 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 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 => {} _ => {}
186 }
187 };
188
189 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; }
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 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}