Skip to main content

ferrin_google/interactions/
mod.rs

1//! Gemini Interactions language models, derived from the Vercel AI SDK
2//! (Apache-2.0, Copyright 2023 Vercel, Inc.), translated and modified; see `NOTICE`.
3
4mod background;
5mod output;
6mod prompt;
7mod request;
8mod sources;
9mod stream;
10mod stream_content;
11mod tools;
12
13use std::time::Duration;
14
15use ferrin_provider_util::http::ResponseHandlers;
16use ferrin_provider_util::http::event_source_response_handler;
17use ferrin_provider_util::http::get;
18use ferrin_provider_util::http::json_response_handler;
19use ferrin_provider_util::http::post_json;
20use ferrin_provider_util::stream_driver::drive_stream;
21use ferrin_spec::Headers;
22use ferrin_spec::JsonValue;
23use ferrin_spec::ModelId;
24use ferrin_spec::ProviderId;
25use ferrin_spec::ResponseMetadata;
26use ferrin_spec::error::InvalidArgumentError;
27use ferrin_spec::error::InvalidResponseDataError;
28use ferrin_spec::error::ProviderError;
29use ferrin_spec::language_model::CallOptions;
30use ferrin_spec::language_model::GenerateResult;
31use ferrin_spec::language_model::LanguageModel;
32use ferrin_spec::language_model::RequestMetadata;
33use ferrin_spec::language_model::StreamPart;
34use ferrin_spec::language_model::StreamResult;
35use ferrin_spec::language_model::SupportedUrls;
36use tokio_util::sync::CancellationToken;
37
38use crate::config::SharedConfig;
39use crate::error::failed_response_handler;
40
41/// Headers and cancellation for an existing Interactions resource.
42#[derive(Debug, Clone, Default)]
43pub struct GoogleInteractionOptions {
44    /// Extra request headers.
45    pub headers: Headers,
46    /// Cancellation of this resource operation.
47    pub cancellation: CancellationToken,
48}
49
50/// A language model backed by Google's Interactions API.
51///
52/// # Examples
53///
54/// ```no_run
55/// # async fn run() -> Result<(), ferrin_spec::error::ProviderError> {
56/// use ferrin_google::{create_google, GoogleSettings};
57/// use ferrin_spec::LanguageModel;
58/// use ferrin_spec::language_model::{CallOptions, PromptMessage};
59/// let google = create_google(GoogleSettings::default())?;
60/// let model = google.interactions("gemini-2.5-flash");
61/// let result = model.do_generate(CallOptions::new(vec![PromptMessage::user_text("Hello")])).await?;
62/// # Ok(()) }
63/// ```
64#[derive(Debug, Clone)]
65pub struct GoogleInteractionsLanguageModel {
66    config: SharedConfig,
67    provider: ProviderId,
68    model_id: ModelId,
69}
70
71impl GoogleInteractionsLanguageModel {
72    /// Creates an Interactions model using the shared provider configuration.
73    #[must_use]
74    pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
75        Self {
76            provider: config.provider_id("interactions"),
77            config,
78            model_id: model_id.into(),
79        }
80    }
81
82    /// Returns the shared configuration.
83    #[must_use]
84    pub fn config(&self) -> &SharedConfig {
85        &self.config
86    }
87
88    /// Converts a call to the Interactions request body without sending it.
89    ///
90    /// # Errors
91    ///
92    /// Returns an invalid argument error for unsupported options or file references.
93    pub fn prepare_request(&self, options: &CallOptions) -> Result<JsonValue, ProviderError> {
94        Ok(request::prepare(&self.config, self.model_id.as_str(), options)?.body)
95    }
96
97    /// Creates an interaction and returns its initial resource without polling.
98    ///
99    /// # Errors
100    ///
101    /// Returns option conversion, transport, provider or cancellation errors.
102    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
103    pub async fn start_interaction(
104        &self,
105        options: CallOptions,
106    ) -> Result<JsonValue, ProviderError> {
107        let body = self.prepare_request(&options)?;
108        let handlers = ResponseHandlers::new(
109            json_response_handler::<JsonValue>(),
110            failed_response_handler(),
111        );
112        Ok(post_json(
113            self.config.transport.as_ref(),
114            self.config.url("interactions"),
115            self.config.headers(&options.headers)?,
116            &body,
117            &handlers,
118            options.cancellation,
119        )
120        .await?
121        .value)
122    }
123
124    /// Retrieves an interaction resource by its server-assigned identifier.
125    ///
126    /// # Errors
127    ///
128    /// Returns an invalid argument error for an empty ID, or an HTTP/provider error.
129    #[tracing::instrument(skip_all)]
130    pub async fn get_interaction(
131        &self,
132        id: &str,
133        options: GoogleInteractionOptions,
134    ) -> Result<JsonValue, ProviderError> {
135        let handlers = ResponseHandlers::new(
136            json_response_handler::<JsonValue>(),
137            failed_response_handler(),
138        );
139        Ok(get(
140            self.config.transport.as_ref(),
141            background::url(&self.config, id)?,
142            self.config.headers(&options.headers)?,
143            &handlers,
144            options.cancellation,
145        )
146        .await?
147        .value)
148    }
149
150    /// Stops an interaction on the server, including a background run.
151    ///
152    /// # Errors
153    ///
154    /// Returns an invalid argument error for an empty ID, or an HTTP/provider error.
155    #[tracing::instrument(skip_all)]
156    pub async fn cancel_interaction(
157        &self,
158        id: &str,
159        options: GoogleInteractionOptions,
160    ) -> Result<JsonValue, ProviderError> {
161        background::cancel(
162            &self.config,
163            id,
164            self.config.headers(&options.headers)?,
165            options.cancellation,
166        )
167        .await
168    }
169
170    async fn generate(
171        &self,
172        options: &CallOptions,
173        prepared: &request::Prepared,
174    ) -> Result<GenerateResult, ProviderError> {
175        let handlers = ResponseHandlers::new(
176            json_response_handler::<JsonValue>(),
177            failed_response_handler(),
178        );
179        let headers = self.config.headers(&options.headers)?;
180        let mut response = post_json(
181            self.config.transport.as_ref(),
182            self.config.url("interactions"),
183            headers.clone(),
184            &prepared.body,
185            &handlers,
186            options.cancellation.clone(),
187        )
188        .await?;
189        let deadline = tokio::time::Instant::now() + Duration::from_millis(prepared.timeout_ms);
190        while matches!(response.value["status"].as_str(), Some("in_progress")) {
191            if prepared.body["background"] != true && prepared.body.get("agent").is_none() {
192                return Err(bad_response(
193                    "nonterminal interactions response without background mode",
194                ));
195            }
196            let id = response.value["id"]
197                .as_str()
198                .filter(|id| !id.is_empty())
199                .ok_or_else(|| bad_response("background interactions response omitted id"))?;
200            let url = background::url(&self.config, id)?;
201            let poll = async {
202                tokio::time::sleep(Duration::from_secs(1)).await;
203                get(
204                    self.config.transport.as_ref(),
205                    url,
206                    headers.clone(),
207                    &handlers,
208                    options.cancellation.clone(),
209                )
210                .await
211            };
212            response = tokio::select! {
213                biased;
214                () = options.cancellation.cancelled() => {
215                    let _ = tokio::time::timeout(Duration::from_secs(2), background::cancel(&self.config, id, headers.clone(), CancellationToken::new())).await;
216                    return Err(ProviderError::Cancelled);
217                },
218                response = tokio::time::timeout_at(deadline, poll) => response.map_err(|_| bad_response("background interactions polling timed out"))??,
219            };
220        }
221        let mut result = output::convert(&self.config, &response.value, &prepared.aliases)?;
222        result.warnings = prepared.warnings.clone();
223        result.request = RequestMetadata::with_body(prepared.body.clone());
224        result.response.headers = Some(response.response_headers);
225        Ok(result)
226    }
227}
228
229impl LanguageModel for GoogleInteractionsLanguageModel {
230    fn provider(&self) -> &ProviderId {
231        &self.provider
232    }
233    fn model_id(&self) -> &ModelId {
234        &self.model_id
235    }
236    async fn supported_urls(&self) -> SupportedUrls {
237        crate::language_model::supported_urls(&self.config.base_url, self.model_id.as_str())
238    }
239
240    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
241    async fn do_generate(&self, options: CallOptions) -> Result<GenerateResult, ProviderError> {
242        let prepared = request::prepare(&self.config, self.model_id.as_str(), &options)?;
243        self.generate(&options, &prepared).await
244    }
245
246    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
247    async fn do_stream(&self, options: CallOptions) -> Result<StreamResult, ProviderError> {
248        let mut prepared = request::prepare(&self.config, self.model_id.as_str(), &options)?;
249        if prepared.body["background"] == true {
250            let handlers = ResponseHandlers::new(
251                json_response_handler::<JsonValue>(),
252                failed_response_handler(),
253            );
254            let headers = self.config.headers(&options.headers)?;
255            let response = post_json(
256                self.config.transport.as_ref(),
257                self.config.url("interactions"),
258                headers.clone(),
259                &prepared.body,
260                &handlers,
261                options.cancellation.clone(),
262            )
263            .await?;
264            if response.value["status"] == "in_progress" {
265                let id = response.value["id"]
266                    .as_str()
267                    .filter(|id| !id.is_empty())
268                    .ok_or_else(|| bad_response("background interaction omitted id"))?;
269                let chunks = background::stream(
270                    self.config.clone(),
271                    id.to_owned(),
272                    headers,
273                    options.cancellation,
274                    prepared.timeout_ms,
275                );
276                let stream = drive_stream(
277                    StreamPart::StreamStart {
278                        warnings: prepared.warnings,
279                    },
280                    chunks,
281                    stream::State::new(self.config.clone(), prepared.aliases)
282                        .with_interaction(response.value),
283                    options.include_raw_chunks,
284                );
285                let mut result = StreamResult::new(stream);
286                result.request = RequestMetadata::with_body(prepared.body);
287                result.response = ResponseMetadata::with_headers(response.response_headers);
288                return Ok(result);
289            }
290            let mut result = output::convert(&self.config, &response.value, &prepared.aliases)?;
291            result.warnings = prepared.warnings;
292            result.request = RequestMetadata::with_body(prepared.body);
293            result.response.headers = Some(response.response_headers);
294            let mut parts = vec![StreamPart::StreamStart {
295                warnings: result.warnings,
296            }];
297            if options.include_raw_chunks
298                && let Some(raw) = result.response.body.clone()
299            {
300                parts.push(StreamPart::Raw { raw_value: raw });
301            }
302            parts.push(StreamPart::ResponseMetadata {
303                id: result.response.id.clone(),
304                timestamp: result.response.timestamp,
305                model_id: result.response.model_id.clone(),
306            });
307            for (index, content) in result.content.into_iter().enumerate() {
308                parts.extend(stream_content::content_parts(
309                    content,
310                    &format!("step-{index}").into(),
311                ));
312            }
313            parts.push(StreamPart::Finish {
314                finish_reason: result.finish_reason,
315                usage: result.usage,
316                provider_metadata: result.provider_metadata,
317            });
318            let mut stream = StreamResult::new(Box::pin(futures_util::stream::iter(parts)));
319            stream.request = result.request;
320            stream.response = result.response;
321            return Ok(stream);
322        }
323        prepared.body["stream"] = JsonValue::Bool(true);
324        let handlers = ResponseHandlers::new(
325            event_source_response_handler::<JsonValue>(),
326            failed_response_handler(),
327        );
328        let response = post_json(
329            self.config.transport.as_ref(),
330            self.config.url("interactions"),
331            self.config.headers(&options.headers)?,
332            &prepared.body,
333            &handlers,
334            options.cancellation.clone(),
335        )
336        .await?;
337        let stream = drive_stream(
338            StreamPart::StreamStart {
339                warnings: prepared.warnings,
340            },
341            stream::terminal_chunks(response.value),
342            stream::State::new(self.config.clone(), prepared.aliases),
343            options.include_raw_chunks,
344        );
345        let mut result = StreamResult::new(stream);
346        result.request = RequestMetadata::with_body(prepared.body);
347        result.response = ResponseMetadata::with_headers(response.response_headers);
348        Ok(result)
349    }
350}
351
352fn invalid(argument: &str, message: &str) -> ProviderError {
353    InvalidArgumentError::new(argument, message).into()
354}
355
356fn bad_response(message: &str) -> ProviderError {
357    InvalidResponseDataError::new(message, JsonValue::Null).into()
358}