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", deny_unknown_fields)]
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    #[serde(default)]
68    embeddings: Vec<EmbeddingValues>,
69}
70
71#[derive(Debug, Deserialize)]
72struct EmbeddingValues {
73    #[serde(default)]
74    values: Embedding,
75}
76
77/// Embedding model backed by `embedContent`.
78#[derive(Debug, Clone)]
79pub struct GoogleEmbeddingModel {
80    config: SharedConfig,
81    provider: ProviderId,
82    model_id: ModelId,
83}
84
85/// A prepared embedding request.
86#[derive(Debug, Clone, PartialEq)]
87pub struct PreparedEmbeddingRequest {
88    /// Action (`embedContent` or `batchEmbedContents`).
89    pub action: &'static str,
90    /// Request body.
91    pub body: JsonValue,
92}
93
94fn parts(value: &str, extra: Option<&Vec<JsonValue>>) -> Vec<JsonValue> {
95    let mut parts = Vec::new();
96    match extra {
97        Some(extra) => {
98            if !value.is_empty() {
99                parts.push(json!({"text": value}));
100            }
101            parts.extend(extra.iter().cloned());
102        }
103        None => parts.push(json!({"text": value})),
104    }
105    parts
106}
107
108impl GoogleEmbeddingModel {
109    /// Creates the model.
110    #[must_use]
111    pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
112        Self {
113            provider: ProviderId::new(config.name.clone()),
114            config,
115            model_id: model_id.into(),
116        }
117    }
118
119    /// Builds the request for `options`.
120    ///
121    /// # Errors
122    ///
123    /// Returns [`ProviderError::TooManyEmbeddingValues`] above
124    /// [`MAX_EMBEDDINGS_PER_CALL`] and [`ProviderError::InvalidArgument`] for
125    /// invalid options or a `content` list whose length differs from the
126    /// values.
127    pub fn prepare_request(
128        &self,
129        options: &EmbedOptions,
130    ) -> Result<PreparedEmbeddingRequest, ProviderError> {
131        if options.values.len() > MAX_EMBEDDINGS_PER_CALL {
132            return Err(TooManyEmbeddingValuesForCallError {
133                provider: self.provider.clone(),
134                model_id: self.model_id.clone(),
135                max_embeddings_per_call: MAX_EMBEDDINGS_PER_CALL,
136                value_count: options.values.len(),
137            }
138            .into());
139        }
140        let google = parse_merged::<GoogleEmbeddingOptions>(
141            &self.config,
142            &options.provider_options,
143            GoogleEmbeddingOptions::merge,
144        )?;
145        if let Some(content) = &google.content
146            && content.len() != options.values.len()
147        {
148            return Err(InvalidArgumentError::new(
149                "content",
150                format!(
151                    "the number of multimodal content entries ({}) must match the number of values ({})",
152                    content.len(),
153                    options.values.len()
154                ),
155            )
156            .into());
157        }
158        let model = GoogleConfig::model_path(self.model_id.as_str());
159        let extra = |index: usize| {
160            google
161                .content
162                .as_ref()
163                .and_then(|content| content.get(index))
164                .and_then(Option::as_ref)
165        };
166        let common = |request: &mut JsonObject| {
167            if let Some(dimensionality) = google.output_dimensionality {
168                request.insert(
169                    "outputDimensionality".to_owned(),
170                    JsonValue::from(dimensionality),
171                );
172            }
173            if let Some(task_type) = &google.task_type {
174                request.insert("taskType".to_owned(), JsonValue::from(task_type.as_str()));
175            }
176        };
177        if let [value] = options.values.as_slice() {
178            let mut request = JsonObject::new();
179            request.insert("model".to_owned(), JsonValue::from(model.as_str()));
180            request.insert(
181                "content".to_owned(),
182                json!({"parts": parts(value, extra(0))}),
183            );
184            common(&mut request);
185            return Ok(PreparedEmbeddingRequest {
186                action: "embedContent",
187                body: JsonValue::Object(request),
188            });
189        }
190        let requests: Vec<JsonValue> = options
191            .values
192            .iter()
193            .enumerate()
194            .map(|(index, value)| {
195                let mut request = JsonObject::new();
196                request.insert("model".to_owned(), JsonValue::from(model.as_str()));
197                request.insert(
198                    "content".to_owned(),
199                    json!({"role": "user", "parts": parts(value, extra(index))}),
200                );
201                common(&mut request);
202                JsonValue::Object(request)
203            })
204            .collect();
205        Ok(PreparedEmbeddingRequest {
206            action: "batchEmbedContents",
207            body: json!({"requests": requests}),
208        })
209    }
210}
211
212impl EmbeddingModel for GoogleEmbeddingModel {
213    fn provider(&self) -> &ProviderId {
214        &self.provider
215    }
216
217    fn model_id(&self) -> &ModelId {
218        &self.model_id
219    }
220
221    fn max_embeddings_per_call(&self) -> Option<usize> {
222        Some(MAX_EMBEDDINGS_PER_CALL)
223    }
224
225    fn supports_parallel_calls(&self) -> bool {
226        true
227    }
228
229    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
230    async fn do_embed(&self, options: EmbedOptions) -> Result<EmbedResult, ProviderError> {
231        let prepared = self.prepare_request(&options)?;
232        let url = self
233            .config
234            .model_url(self.model_id.as_str(), prepared.action);
235        let headers = self.config.headers(&options.headers)?;
236        let (embeddings, response_headers, raw) = if prepared.action == "embedContent" {
237            let handlers = ResponseHandlers::new(
238                json_response_handler::<SingleEmbeddingResponse>(),
239                failed_response_handler(),
240            );
241            let response = post_json(
242                self.config.transport.as_ref(),
243                url,
244                headers,
245                &prepared.body,
246                &handlers,
247                options.cancellation.clone(),
248            )
249            .await?;
250            (
251                vec![response.value.embedding.values],
252                response.response_headers,
253                response.raw,
254            )
255        } else {
256            let handlers = ResponseHandlers::new(
257                json_response_handler::<BatchEmbeddingResponse>(),
258                failed_response_handler(),
259            );
260            let response = post_json(
261                self.config.transport.as_ref(),
262                url,
263                headers,
264                &prepared.body,
265                &handlers,
266                options.cancellation.clone(),
267            )
268            .await?;
269            (
270                response
271                    .value
272                    .embeddings
273                    .into_iter()
274                    .map(|embedding| embedding.values)
275                    .collect(),
276                response.response_headers,
277                response.raw,
278            )
279        };
280        Ok(EmbedResult {
281            embeddings,
282            usage: None,
283            provider_metadata: None,
284            response: ResponseMetadata {
285                id: None,
286                timestamp: Some(chrono::Utc::now()),
287                model_id: Some(self.model_id.clone()),
288                headers: Some(response_headers),
289                body: raw,
290            },
291            warnings: Vec::new(),
292        })
293    }
294}