rig_core/providers/gemini/
embedding.rs1use 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
15pub const EMBEDDING_001: &str = "gemini-embedding-001";
17pub const EMBEDDING_004: &str = "text-embedding-004";
19
20fn 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 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 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 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
155mod 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 let model =
198 <EmbeddingModel as embeddings::EmbeddingModel>::make(&client, EMBEDDING_001, None);
199 assert_eq!(embeddings::EmbeddingModel::ndims(&model), 3072);
200
201 let model =
203 <EmbeddingModel as embeddings::EmbeddingModel>::make(&client, EMBEDDING_004, None);
204 assert_eq!(embeddings::EmbeddingModel::ndims(&model), 768);
205
206 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 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 let body = r#"{"error":{"code":503,"message":"boom","status":"UNAVAILABLE"}}"#;
272 let http_client = RecordingHttpClient::new(body); 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}