ferrin_google/interactions/
mod.rs1mod 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#[derive(Debug, Clone, Default)]
43pub struct GoogleInteractionOptions {
44 pub headers: Headers,
46 pub cancellation: CancellationToken,
48}
49
50#[derive(Debug, Clone)]
65pub struct GoogleInteractionsLanguageModel {
66 config: SharedConfig,
67 provider: ProviderId,
68 model_id: ModelId,
69}
70
71impl GoogleInteractionsLanguageModel {
72 #[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 #[must_use]
84 pub fn config(&self) -> &SharedConfig {
85 &self.config
86 }
87
88 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 #[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 #[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 #[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}