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", deny_unknown_fields)]
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 #[serde(default)]
68 embeddings: Vec<EmbeddingValues>,
69}
70
71#[derive(Debug, Deserialize)]
72struct EmbeddingValues {
73 #[serde(default)]
74 values: Embedding,
75}
76
77#[derive(Debug, Clone)]
79pub struct GoogleEmbeddingModel {
80 config: SharedConfig,
81 provider: ProviderId,
82 model_id: ModelId,
83}
84
85#[derive(Debug, Clone, PartialEq)]
87pub struct PreparedEmbeddingRequest {
88 pub action: &'static str,
90 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 #[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 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}