Skip to main content

ferrin_google/
embedding.rs

1//! Gemini embedding model (`embedContent` / `batchEmbedContents`).
2
3use ferrin_provider_util::http::ResponseHandlers;
4use ferrin_provider_util::http::json_response_handler;
5use ferrin_provider_util::http::post_json;
6use ferrin_spec::JsonObject;
7use ferrin_spec::JsonValue;
8use ferrin_spec::ModelId;
9use ferrin_spec::ProviderId;
10use ferrin_spec::ResponseMetadata;
11use ferrin_spec::embedding_model::EmbedOptions;
12use ferrin_spec::embedding_model::EmbedResult;
13use ferrin_spec::embedding_model::Embedding;
14use ferrin_spec::embedding_model::EmbeddingModel;
15use ferrin_spec::error::InvalidArgumentError;
16use ferrin_spec::error::ProviderError;
17use ferrin_spec::error::TooManyEmbeddingValuesForCallError;
18use serde::Deserialize;
19use serde_json::json;
20
21use crate::config::GoogleConfig;
22use crate::config::SharedConfig;
23use crate::error::failed_response_handler;
24use crate::options::parse_merged;
25
26/// Maximum values per call.
27pub const MAX_EMBEDDINGS_PER_CALL: usize = 100;
28
29/// Embedding options (`provider_options["google"]`).
30#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
31#[serde(rename_all = "camelCase")]
32pub struct GoogleEmbeddingOptions {
33    /// Output dimensionality.
34    #[serde(default)]
35    pub output_dimensionality: Option<u32>,
36    /// Task type (`SEMANTIC_SIMILARITY`, `RETRIEVAL_DOCUMENT`, ...).
37    #[serde(default)]
38    pub task_type: Option<String>,
39    /// Extra multimodal parts per value (`[{text} | {inlineData} | {fileData}]`
40    /// or `null`); must have one entry per value.
41    #[serde(default)]
42    pub content: Option<Vec<Option<Vec<JsonValue>>>>,
43}
44
45impl GoogleEmbeddingOptions {
46    fn merge(mut self, other: Self) -> Self {
47        if other.output_dimensionality.is_some() {
48            self.output_dimensionality = other.output_dimensionality;
49        }
50        if other.task_type.is_some() {
51            self.task_type = other.task_type;
52        }
53        if other.content.is_some() {
54            self.content = other.content;
55        }
56        self
57    }
58}
59
60#[derive(Debug, Deserialize)]
61struct SingleEmbeddingResponse {
62    embedding: EmbeddingValues,
63}
64
65#[derive(Debug, Deserialize)]
66struct BatchEmbeddingResponse {
67    embeddings: Vec<EmbeddingValues>,
68}
69
70#[derive(Debug, Deserialize)]
71struct EmbeddingValues {
72    values: Embedding,
73}
74
75/// Embedding model backed by `embedContent`.
76#[derive(Debug, Clone)]
77pub struct GoogleEmbeddingModel {
78    config: SharedConfig,
79    provider: ProviderId,
80    model_id: ModelId,
81}
82
83/// A prepared embedding request.
84#[derive(Debug, Clone, PartialEq)]
85pub struct PreparedEmbeddingRequest {
86    /// Action (`embedContent` or `batchEmbedContents`).
87    pub action: &'static str,
88    /// Request body.
89    pub body: JsonValue,
90}
91
92fn parts(value: &str, extra: Option<&Vec<JsonValue>>) -> Vec<JsonValue> {
93    let mut parts = Vec::new();
94    match extra {
95        Some(extra) => {
96            if !value.is_empty() {
97                parts.push(json!({"text": value}));
98            }
99            parts.extend(extra.iter().cloned());
100        }
101        None => parts.push(json!({"text": value})),
102    }
103    parts
104}
105
106fn valid_part(part: &JsonValue) -> bool {
107    part.get("text").is_some_and(JsonValue::is_string)
108        || part.get("inlineData").is_some_and(|data| {
109            data.get("mimeType").is_some_and(JsonValue::is_string)
110                && data.get("data").is_some_and(JsonValue::is_string)
111        })
112        || part.get("fileData").is_some_and(|data| {
113            data.get("mimeType").is_some_and(JsonValue::is_string)
114                && data.get("fileUri").is_some_and(JsonValue::is_string)
115        })
116}
117
118impl GoogleEmbeddingModel {
119    /// Creates the model.
120    #[must_use]
121    pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
122        Self {
123            provider: ProviderId::new(config.name.clone()),
124            config,
125            model_id: model_id.into(),
126        }
127    }
128
129    /// Builds the request for `options`.
130    ///
131    /// # Errors
132    ///
133    /// Returns [`ProviderError::TooManyEmbeddingValues`] above
134    /// [`MAX_EMBEDDINGS_PER_CALL`] and [`ProviderError::InvalidArgument`] for
135    /// invalid options or a `content` list whose length differs from the
136    /// values.
137    pub fn prepare_request(
138        &self,
139        options: &EmbedOptions,
140    ) -> Result<PreparedEmbeddingRequest, ProviderError> {
141        if options.values.len() > MAX_EMBEDDINGS_PER_CALL {
142            return Err(TooManyEmbeddingValuesForCallError {
143                provider: self.provider.clone(),
144                model_id: self.model_id.clone(),
145                max_embeddings_per_call: MAX_EMBEDDINGS_PER_CALL,
146                value_count: options.values.len(),
147            }
148            .into());
149        }
150        let google = parse_merged::<GoogleEmbeddingOptions>(
151            &self.config,
152            &options.provider_options,
153            GoogleEmbeddingOptions::merge,
154        )?;
155        if google.task_type.as_deref().is_some_and(|task| {
156            !matches!(
157                task,
158                "SEMANTIC_SIMILARITY"
159                    | "CLASSIFICATION"
160                    | "CLUSTERING"
161                    | "RETRIEVAL_DOCUMENT"
162                    | "RETRIEVAL_QUERY"
163                    | "QUESTION_ANSWERING"
164                    | "FACT_VERIFICATION"
165                    | "CODE_RETRIEVAL_QUERY"
166            )
167        }) {
168            return Err(
169                InvalidArgumentError::new("taskType", "unsupported embedding task type").into(),
170            );
171        }
172        for content in google.content.iter().flatten().flatten() {
173            if content.is_empty() || content.iter().any(|part| !valid_part(part)) {
174                return Err(InvalidArgumentError::new("content", "embedding content must contain nonempty lists of text, inlineData or fileData parts").into());
175            }
176        }
177        if let Some(content) = &google.content
178            && content.len() != options.values.len()
179        {
180            return Err(InvalidArgumentError::new(
181                "content",
182                format!(
183                    "the number of multimodal content entries ({}) must match the number of values ({})",
184                    content.len(),
185                    options.values.len()
186                ),
187            )
188            .into());
189        }
190        let model = GoogleConfig::model_path(self.model_id.as_str());
191        let extra = |index: usize| {
192            google
193                .content
194                .as_ref()
195                .and_then(|content| content.get(index))
196                .and_then(Option::as_ref)
197        };
198        let common = |request: &mut JsonObject| {
199            if let Some(dimensionality) = google.output_dimensionality {
200                request.insert(
201                    "outputDimensionality".to_owned(),
202                    JsonValue::from(dimensionality),
203                );
204            }
205            if let Some(task_type) = &google.task_type {
206                request.insert("taskType".to_owned(), JsonValue::from(task_type.as_str()));
207            }
208        };
209        if let [value] = options.values.as_slice() {
210            let mut request = JsonObject::new();
211            request.insert("model".to_owned(), JsonValue::from(model.as_str()));
212            request.insert(
213                "content".to_owned(),
214                json!({"parts": parts(value, extra(0))}),
215            );
216            common(&mut request);
217            return Ok(PreparedEmbeddingRequest {
218                action: "embedContent",
219                body: JsonValue::Object(request),
220            });
221        }
222        let requests: Vec<JsonValue> = options
223            .values
224            .iter()
225            .enumerate()
226            .map(|(index, value)| {
227                let mut request = JsonObject::new();
228                request.insert("model".to_owned(), JsonValue::from(model.as_str()));
229                request.insert(
230                    "content".to_owned(),
231                    json!({"role": "user", "parts": parts(value, extra(index))}),
232                );
233                common(&mut request);
234                JsonValue::Object(request)
235            })
236            .collect();
237        Ok(PreparedEmbeddingRequest {
238            action: "batchEmbedContents",
239            body: json!({"requests": requests}),
240        })
241    }
242}
243
244impl EmbeddingModel for GoogleEmbeddingModel {
245    fn provider(&self) -> &ProviderId {
246        &self.provider
247    }
248
249    fn model_id(&self) -> &ModelId {
250        &self.model_id
251    }
252
253    fn max_embeddings_per_call(&self) -> Option<usize> {
254        Some(MAX_EMBEDDINGS_PER_CALL)
255    }
256
257    fn supports_parallel_calls(&self) -> bool {
258        true
259    }
260
261    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
262    async fn do_embed(&self, options: EmbedOptions) -> Result<EmbedResult, ProviderError> {
263        let prepared = self.prepare_request(&options)?;
264        let url = self
265            .config
266            .model_url(self.model_id.as_str(), prepared.action);
267        let headers = self.config.headers(&options.headers)?;
268        let (embeddings, response_headers, raw) = if prepared.action == "embedContent" {
269            let handlers = ResponseHandlers::new(
270                json_response_handler::<SingleEmbeddingResponse>(),
271                failed_response_handler(),
272            );
273            let response = post_json(
274                self.config.transport.as_ref(),
275                url,
276                headers,
277                &prepared.body,
278                &handlers,
279                options.cancellation.clone(),
280            )
281            .await?;
282            (
283                vec![response.value.embedding.values],
284                response.response_headers,
285                response.raw,
286            )
287        } else {
288            let handlers = ResponseHandlers::new(
289                json_response_handler::<BatchEmbeddingResponse>(),
290                failed_response_handler(),
291            );
292            let response = post_json(
293                self.config.transport.as_ref(),
294                url,
295                headers,
296                &prepared.body,
297                &handlers,
298                options.cancellation.clone(),
299            )
300            .await?;
301            (
302                response
303                    .value
304                    .embeddings
305                    .into_iter()
306                    .map(|embedding| embedding.values)
307                    .collect(),
308                response.response_headers,
309                response.raw,
310            )
311        };
312        Ok(EmbedResult {
313            embeddings,
314            usage: None,
315            provider_metadata: None,
316            response: ResponseMetadata {
317                id: None,
318                timestamp: Some(chrono::Utc::now()),
319                model_id: Some(self.model_id.clone()),
320                headers: Some(response_headers),
321                body: raw,
322            },
323            warnings: Vec::new(),
324        })
325    }
326}