1use 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
26pub const MAX_EMBEDDINGS_PER_CALL: usize = 100;
28
29#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
31#[serde(rename_all = "camelCase")]
32pub struct GoogleEmbeddingOptions {
33 #[serde(default)]
35 pub output_dimensionality: Option<u32>,
36 #[serde(default)]
38 pub task_type: Option<String>,
39 #[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#[derive(Debug, Clone)]
77pub struct GoogleEmbeddingModel {
78 config: SharedConfig,
79 provider: ProviderId,
80 model_id: ModelId,
81}
82
83#[derive(Debug, Clone, PartialEq)]
85pub struct PreparedEmbeddingRequest {
86 pub action: &'static str,
88 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 #[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 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}