Skip to main content

ferrin_google/
language_model.rs

1//! Gemini language model (`generateContent` / `streamGenerateContent`).
2
3use ferrin_provider_util::http::ResponseHandlers;
4use ferrin_provider_util::http::event_source_response_handler;
5use ferrin_provider_util::http::json_response_handler;
6use ferrin_provider_util::http::post_json;
7use ferrin_provider_util::stream_driver::drive_stream;
8use ferrin_provider_util::tool_name_mapping::ToolNameMapping;
9use ferrin_spec::Content;
10use ferrin_spec::FinishReason;
11use ferrin_spec::FinishReasonKind;
12use ferrin_spec::JsonValue;
13use ferrin_spec::ModelId;
14use ferrin_spec::ProviderId;
15use ferrin_spec::ResponseMetadata;
16use ferrin_spec::error::ProviderError;
17use ferrin_spec::language_model::CallOptions;
18use ferrin_spec::language_model::GenerateResult;
19use ferrin_spec::language_model::LanguageModel;
20use ferrin_spec::language_model::RequestMetadata;
21use ferrin_spec::language_model::StreamPart;
22use ferrin_spec::language_model::StreamResult;
23use ferrin_spec::language_model::SupportedUrls;
24use ferrin_spec::shared::Warning;
25use regex::Regex;
26use url::Url;
27
28use crate::api_types::GenerateContentResponse;
29use crate::capabilities::capabilities;
30use crate::config::SharedConfig;
31use crate::error::failed_response_handler;
32use crate::output::OutputMapper;
33use crate::output::convert_usage;
34use crate::output::has_client_tool_calls;
35use crate::output::map_finish_reason;
36use crate::output::raw_usage;
37use crate::request::PreparedRequest;
38use crate::request::prepare_request;
39use crate::stream::GoogleStreamState;
40
41/// Provider id family of the language model.
42pub const FAMILY: &str = "generative-ai";
43
44/// Media types the Gemini API fetches from arbitrary HTTPS URLs.
45pub const EXTERNAL_URL_MEDIA_TYPES: [&str; 22] = [
46    "text/html",
47    "text/css",
48    "text/plain",
49    "text/xml",
50    "text/csv",
51    "text/rtf",
52    "text/javascript",
53    "application/json",
54    "application/pdf",
55    "image/bmp",
56    "image/jpeg",
57    "image/png",
58    "image/webp",
59    "video/mp4",
60    "video/mpeg",
61    "video/quicktime",
62    "video/avi",
63    "video/x-flv",
64    "video/mpg",
65    "video/webm",
66    "video/wmv",
67    "video/3gpp",
68];
69
70fn regex(pattern: &str) -> Option<Regex> {
71    Regex::new(pattern).ok()
72}
73
74/// URLs every Gemini request accepts: Files API URIs (public endpoint and
75/// the configured base URL) and YouTube videos.
76#[must_use]
77pub fn base_supported_urls(base_url: &Url) -> SupportedUrls {
78    let escaped_base = regex::escape(base_url.as_str().trim_end_matches('/'));
79    let patterns = [
80        r"^https://generativelanguage\.googleapis\.com/v1beta/files/.*$".to_owned(),
81        format!("^{escaped_base}/files/.*$"),
82        r"^https://(?:www\.)?youtube\.com/watch\?v=[\w-]+(?:&[\w=&.-]*)?$".to_owned(),
83        r"^https://youtu\.be/[\w-]+(?:\?[\w=&.-]*)?$".to_owned(),
84    ];
85    SupportedUrls::none().with("*", patterns.iter().filter_map(|pattern| regex(pattern)))
86}
87
88/// Supported URLs of `model_id`: the base set plus, for Gemini models other
89/// than Gemini 2.0, any HTTPS URL for [`EXTERNAL_URL_MEDIA_TYPES`].
90#[must_use]
91pub fn supported_urls(base_url: &Url, model_id: &str) -> SupportedUrls {
92    let mut urls = base_supported_urls(base_url);
93    let lower = model_id.to_ascii_lowercase();
94    let is_gemini = lower
95        .split('/')
96        .any(|segment| segment.starts_with("gemini-"));
97    let is_gemini_2_0 = lower
98        .split('/')
99        .any(|segment| segment.starts_with("gemini-2.0"));
100    if is_gemini && !is_gemini_2_0 {
101        for media_type in EXTERNAL_URL_MEDIA_TYPES {
102            urls.insert(media_type, regex("^https://.*$"));
103        }
104    }
105    urls
106}
107
108/// Language model backed by `generateContent`.
109#[derive(Debug, Clone)]
110pub struct GoogleLanguageModel {
111    config: SharedConfig,
112    provider: ProviderId,
113    model_id: ModelId,
114}
115
116impl GoogleLanguageModel {
117    /// Creates the model.
118    #[must_use]
119    pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
120        Self {
121            provider: config.provider_id(FAMILY),
122            config,
123            model_id: model_id.into(),
124        }
125    }
126
127    /// Shared configuration.
128    #[must_use]
129    pub fn config(&self) -> &SharedConfig {
130        &self.config
131    }
132
133    /// Builds the request body and warnings for `options` without sending.
134    ///
135    /// # Errors
136    ///
137    /// See [`prepare_request`].
138    pub fn prepare_request(&self, options: &CallOptions) -> Result<PreparedRequest, ProviderError> {
139        prepare_request(&self.config, self.model_id.as_str(), options)
140    }
141
142    /// Converts a parsed `generateContent` response (`raw` is the original
143    /// JSON) into a result. Used for direct calls and batch results.
144    ///
145    /// # Errors
146    ///
147    /// See [`convert_generate_content_response`].
148    pub fn convert_response(
149        &self,
150        prepared: &PreparedRequest,
151        body: &GenerateContentResponse,
152        raw: Option<&JsonValue>,
153    ) -> Result<GenerateResult, ProviderError> {
154        convert_generate_content_response(
155            &self.config,
156            prepared.tool_name_mapping.clone(),
157            prepared.warnings.clone(),
158            body,
159            raw,
160        )
161    }
162}
163
164/// Converts a parsed `generateContent` response (`raw` is the original JSON)
165/// into a result using `mapping` to restore custom tool names; `warnings`
166/// are attached to the result.
167///
168/// # Errors
169///
170/// Returns [`ProviderError::InvalidResponseData`] for undecodable inline
171/// data.
172pub fn convert_generate_content_response(
173    config: &SharedConfig,
174    mapping: ToolNameMapping,
175    warnings: Vec<Warning>,
176    body: &GenerateContentResponse,
177    raw: Option<&JsonValue>,
178) -> Result<GenerateResult, ProviderError> {
179    let mut mapper = OutputMapper::new(config.clone(), mapping);
180    let candidate = body.candidate();
181    let mut content = Vec::new();
182    if let Some(candidate) = candidate {
183        content = mapper.map_parts(candidate.parts())?;
184        content.extend(
185            mapper
186                .sources(&candidate.grounding_chunks())
187                .into_iter()
188                .map(Content::Source),
189        );
190    }
191    let block_reason = body.block_reason();
192    let candidate_reason = candidate.and_then(|candidate| candidate.finish_reason.as_deref());
193    let finish_reason = match (candidate_reason, block_reason) {
194        (None, Some(reason)) => FinishReason::with_raw(FinishReasonKind::ContentFilter, reason),
195        (reason, _) => map_finish_reason(reason, has_client_tool_calls(&content)),
196    };
197    let raw_usage_object = raw_usage(raw);
198    let usage_value = raw_usage_object.clone().map(JsonValue::Object);
199    let metadata = mapper.response_metadata(body, candidate, usage_value.as_ref());
200    let mut result = GenerateResult::new(content, finish_reason);
201    result.usage = convert_usage(body.usage_metadata.as_ref(), raw_usage_object);
202    result.provider_metadata = Some(metadata);
203    result.warnings = warnings;
204    result.response = ResponseMetadata {
205        id: body.response_id.clone(),
206        timestamp: body
207            .create_time
208            .as_deref()
209            .and_then(|time| chrono::DateTime::parse_from_rfc3339(time).ok())
210            .map(|time| time.with_timezone(&chrono::Utc)),
211        model_id: body.model_version.clone().map(Into::into),
212        headers: None,
213        body: raw.cloned(),
214    };
215    Ok(result)
216}
217
218impl LanguageModel for GoogleLanguageModel {
219    fn provider(&self) -> &ProviderId {
220        &self.provider
221    }
222
223    fn model_id(&self) -> &ModelId {
224        &self.model_id
225    }
226
227    async fn supported_urls(&self) -> SupportedUrls {
228        supported_urls(&self.config.base_url, self.model_id.as_str())
229    }
230
231    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
232    async fn do_generate(&self, options: CallOptions) -> Result<GenerateResult, ProviderError> {
233        let prepared = self.prepare_request(&options)?;
234        let url = self
235            .config
236            .model_url(self.model_id.as_str(), "generateContent");
237        let headers = self.config.headers(&options.headers)?;
238        let handlers = ResponseHandlers::new(
239            json_response_handler::<GenerateContentResponse>(),
240            failed_response_handler(),
241        );
242        let request_body = JsonValue::Object(prepared.body.clone());
243        let response = post_json(
244            self.config.transport.as_ref(),
245            url,
246            headers,
247            &request_body,
248            &handlers,
249            options.cancellation.clone(),
250        )
251        .await?;
252        let mut result =
253            self.convert_response(&prepared, &response.value, response.raw.as_ref())?;
254        result.request = RequestMetadata::with_body(request_body);
255        result.response.headers = Some(response.response_headers);
256        Ok(result)
257    }
258
259    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
260    async fn do_stream(&self, options: CallOptions) -> Result<StreamResult, ProviderError> {
261        let prepared = self.prepare_request(&options)?;
262        let mut url = self
263            .config
264            .model_url(self.model_id.as_str(), "streamGenerateContent");
265        url.set_query(Some("alt=sse"));
266        let headers = self.config.headers(&options.headers)?;
267        let handlers = ResponseHandlers::new(
268            event_source_response_handler::<GenerateContentResponse>(),
269            failed_response_handler(),
270        );
271        let request_body = JsonValue::Object(prepared.body.clone());
272        let response = post_json(
273            self.config.transport.as_ref(),
274            url,
275            headers,
276            &request_body,
277            &handlers,
278            options.cancellation.clone(),
279        )
280        .await?;
281        let mapper = OutputMapper::new(self.config.clone(), prepared.tool_name_mapping.clone());
282        let state = GoogleStreamState::new(mapper);
283        let stream = drive_stream(
284            StreamPart::StreamStart {
285                warnings: prepared.warnings,
286            },
287            response.value,
288            state,
289            options.include_raw_chunks,
290        );
291        let mut result = StreamResult::new(stream);
292        result.request = RequestMetadata::with_body(request_body);
293        result.response = ResponseMetadata::with_headers(response.response_headers);
294        Ok(result)
295    }
296}
297
298/// Whether `model_id` is a Gemini model (used by the image model).
299#[must_use]
300pub fn is_gemini_model(model_id: &str) -> bool {
301    capabilities(model_id).is_gemini
302}