Skip to main content

rig_core/providers/gemini/
embedding.rs

1// ================================================================
2//! Google Gemini Embeddings Integration
3//! From [Gemini API Reference](https://ai.google.dev/api/embeddings)
4// ================================================================
5
6use serde_json::json;
7
8use super::{Client, client::ApiResponse};
9use crate::{
10    embeddings::{self, EmbeddingError},
11    http_client::HttpClientExt,
12    wasm_compat::WasmCompatSend,
13};
14
15/// `gemini-embedding-001` embedding model (3072 dimensions by default)
16pub const EMBEDDING_001: &str = "gemini-embedding-001";
17/// `text-embedding-004` embedding model (768 dimensions by default)
18pub const EMBEDDING_004: &str = "text-embedding-004";
19
20/// Returns the default output dimensionality for known Gemini embedding models.
21///
22/// See <https://ai.google.dev/gemini-api/docs/models#gemini-embedding>
23fn model_default_ndims(model: &str) -> Option<usize> {
24    match model {
25        EMBEDDING_001 => Some(3072),
26        EMBEDDING_004 => Some(768),
27        _ => None,
28    }
29}
30
31#[derive(Clone)]
32pub struct EmbeddingModel<T = reqwest::Client> {
33    client: Client<T>,
34    model: String,
35    ndims: usize,
36}
37
38impl<T> EmbeddingModel<T> {
39    pub fn new(client: Client<T>, model: impl Into<String>, ndims: usize) -> Self {
40        Self {
41            client,
42            model: model.into(),
43            ndims,
44        }
45    }
46
47    pub fn with_model(client: Client<T>, model: &str, ndims: usize) -> Self {
48        Self {
49            client,
50            model: model.to_string(),
51            ndims,
52        }
53    }
54}
55
56impl<T> embeddings::EmbeddingModel for EmbeddingModel<T>
57where
58    T: Clone + HttpClientExt + 'static,
59{
60    type Client = Client<T>;
61
62    const MAX_DOCUMENTS: usize = 1024;
63
64    fn make(client: &Self::Client, model: impl Into<String>, dims: Option<usize>) -> Self {
65        let model = model.into();
66        let ndims = dims.or_else(|| model_default_ndims(&model)).unwrap_or(768);
67        Self::new(client.clone(), model, ndims)
68    }
69
70    fn ndims(&self) -> usize {
71        self.ndims
72    }
73
74    /// <https://ai.google.dev/api/embeddings#batch_embed_contents-SHELL>
75    async fn embed_texts(
76        &self,
77        documents: impl IntoIterator<Item = String> + WasmCompatSend,
78    ) -> Result<Vec<embeddings::Embedding>, EmbeddingError> {
79        let documents: Vec<String> = documents.into_iter().collect();
80
81        // Google batch embed requests. See docstrings for API ref link.
82        let requests: Vec<_> = documents
83            .iter()
84            .map(|doc| {
85                json!({
86                    "model": format!("models/{}", self.model),
87                    "content": json!({
88                        "parts": [json!({
89                            "text": doc.to_string()
90                        })]
91                    }),
92                    "output_dimensionality": self.ndims,
93                })
94            })
95            .collect();
96
97        let request_body = json!({ "requests": requests  });
98
99        if let Ok(pretty_body) = serde_json::to_string_pretty(&request_body) {
100            tracing::trace!(
101                target: "rig::embedding",
102                "Sending embedding request to Gemini API {pretty_body}"
103            );
104        }
105
106        let request_body = serde_json::to_vec(&request_body)?;
107        let path = format!("/v1beta/models/{}:batchEmbedContents", self.model);
108        let req = self
109            .client
110            .post(path.as_str())?
111            .body(request_body)
112            .map_err(|e| EmbeddingError::HttpError(e.into()))?;
113        let response = self.client.send::<_, Vec<u8>>(req).await?;
114
115        let status = response.status();
116        let body = response.into_body().await?;
117
118        // Preserve non-success bodies before deserialization because providers
119        // may return empty, non-JSON, or otherwise unexpected error payloads.
120        if !status.is_success() {
121            return Err(EmbeddingError::from_http_response(
122                status,
123                String::from_utf8_lossy(&body),
124            ));
125        }
126
127        match serde_json::from_slice::<ApiResponse<gemini_api_types::EmbeddingResponse>>(&body)? {
128            ApiResponse::Ok(response) => {
129                let docs = documents
130                    .into_iter()
131                    .zip(response.embeddings)
132                    .map(|(document, embedding)| embeddings::Embedding {
133                        document,
134                        vec: embedding
135                            .values
136                            .into_iter()
137                            .filter_map(|n| n.as_f64())
138                            .collect(),
139                    })
140                    .collect();
141
142                Ok(docs)
143            }
144            ApiResponse::Err(err) => {
145                tracing::warn!(message = %err.error.message, "provider returned an error response");
146                Err(EmbeddingError::from_http_response(
147                    status,
148                    String::from_utf8_lossy(&body),
149                ))
150            }
151        }
152    }
153}
154
155// =================================================================
156// Gemini API Types
157// =================================================================
158/// Rust Implementation of the Gemini Types from [Gemini API Reference](https://ai.google.dev/api/embeddings)
159mod gemini_api_types {
160    use serde::Deserialize;
161
162    #[derive(Debug, Deserialize)]
163    pub struct EmbeddingResponse {
164        pub embeddings: Vec<EmbeddingValues>,
165    }
166
167    #[derive(Debug, Deserialize)]
168    pub struct EmbeddingValues {
169        #[serde(default)]
170        pub values: Vec<serde_json::Number>,
171    }
172}
173
174#[cfg(test)]
175mod tests {
176    use super::*;
177
178    #[test]
179    fn test_embedding_values_deserializes_without_empty_values_field() {
180        let values: gemini_api_types::EmbeddingValues =
181            serde_json::from_str("{}").expect("empty embedding values should deserialize");
182        assert!(values.values.is_empty());
183    }
184
185    #[test]
186    fn test_model_default_ndims_lookup() {
187        assert_eq!(model_default_ndims(EMBEDDING_001), Some(3072));
188        assert_eq!(model_default_ndims(EMBEDDING_004), Some(768));
189        assert_eq!(model_default_ndims("unknown-model"), None);
190    }
191
192    #[test]
193    fn test_make_resolves_default_dims() {
194        let client = Client::new("test_key").unwrap();
195
196        // EMBEDDING_001 defaults to 3072
197        let model =
198            <EmbeddingModel as embeddings::EmbeddingModel>::make(&client, EMBEDDING_001, None);
199        assert_eq!(embeddings::EmbeddingModel::ndims(&model), 3072);
200
201        // EMBEDDING_004 defaults to 768
202        let model =
203            <EmbeddingModel as embeddings::EmbeddingModel>::make(&client, EMBEDDING_004, None);
204        assert_eq!(embeddings::EmbeddingModel::ndims(&model), 768);
205
206        // Unknown model falls back to 768
207        let model = <EmbeddingModel as embeddings::EmbeddingModel>::make(
208            &client,
209            "some-future-model",
210            None,
211        );
212        assert_eq!(embeddings::EmbeddingModel::ndims(&model), 768);
213    }
214
215    #[test]
216    fn test_make_respects_explicit_dims() {
217        let client = Client::new("test_key").unwrap();
218
219        let model =
220            <EmbeddingModel as embeddings::EmbeddingModel>::make(&client, EMBEDDING_001, Some(256));
221        assert_eq!(embeddings::EmbeddingModel::ndims(&model), 256);
222    }
223
224    #[test]
225    fn test_new_uses_provided_ndims() {
226        let client = Client::new("test_key").unwrap();
227
228        let model = EmbeddingModel::new(client, EMBEDDING_001, 512);
229        assert_eq!(embeddings::EmbeddingModel::ndims(&model), 512);
230    }
231
232    #[tokio::test]
233    async fn embedding_non_success_preserves_status_and_body() {
234        use crate::client::embeddings::EmbeddingsClient;
235        use crate::embeddings::EmbeddingModel as _;
236        use crate::test_utils::RecordingHttpClient;
237
238        // The non-success status guard preserves the raw provider body without
239        // depending on its envelope shape.
240        let body =
241            r#"{"error":{"code":503,"message":"service unavailable","status":"UNAVAILABLE"}}"#;
242        let http_client =
243            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
244        let client = Client::builder()
245            .api_key("test-key")
246            .http_client(http_client)
247            .build()
248            .expect("build client");
249        let model = client.embedding_model(EMBEDDING_001);
250
251        let error = model
252            .embed_texts(vec!["hello".to_string()])
253            .await
254            .expect_err("should fail with non-success status");
255
256        assert!(matches!(error, EmbeddingError::HttpError(_)));
257        assert_eq!(
258            error.provider_response_status(),
259            Some(http::StatusCode::SERVICE_UNAVAILABLE)
260        );
261        assert_eq!(error.provider_response_body(), Some(body));
262    }
263
264    #[tokio::test]
265    async fn embedding_2xx_error_envelope_preserves_status_and_body() {
266        use crate::client::embeddings::EmbeddingsClient;
267        use crate::embeddings::EmbeddingModel as _;
268        use crate::test_utils::RecordingHttpClient;
269
270        // 200 OK carrying Gemini's standard nested error envelope.
271        let body = r#"{"error":{"code":503,"message":"boom","status":"UNAVAILABLE"}}"#;
272        let http_client = RecordingHttpClient::new(body); // 200 OK
273        let client = Client::builder()
274            .api_key("test-key")
275            .http_client(http_client)
276            .build()
277            .expect("build client");
278        let model = client.embedding_model(EMBEDDING_001);
279
280        let error = model
281            .embed_texts(vec!["hello".to_string()])
282            .await
283            .expect_err("should fail with provider error envelope");
284
285        match &error {
286            EmbeddingError::ProviderResponse(stored) => {
287                assert_eq!(stored.body, body);
288                assert_eq!(stored.status, Some(http::StatusCode::OK));
289            }
290            other => panic!("expected ProviderResponse, got {other:?}"),
291        }
292    }
293}