potato_provider/providers/
embed.rs1use potato_state::block_on;
2use potato_type::google::EmbeddingConfigTrait;
3use potato_type::Provider;
4use tracing::debug;
5use tracing::instrument;
6
7use crate::error::ProviderError;
8use crate::providers::client::GenAiClient;
9use crate::providers::google::GeminiClient;
10use crate::providers::google::VertexClient;
11use crate::providers::openai::OpenAIClient;
12use crate::providers::types::ServiceType;
13use potato_type::google::v1::embedding::{PredictRequest, PredictResponse};
14use potato_type::google::GeminiEmbeddingConfig;
15use potato_type::google::GeminiEmbeddingResponse;
16use potato_type::openai::v1::embedding::{OpenAIEmbeddingConfig, OpenAIEmbeddingResponse};
17use pyo3::prelude::*;
18use serde::Serialize;
19use std::sync::Arc;
20#[derive(Debug, Clone)]
22pub enum EmbeddingInput {
23 Texts(Vec<String>),
24 PredictRequest(PredictRequest),
25}
26
27impl From<Vec<String>> for EmbeddingInput {
28 fn from(texts: Vec<String>) -> Self {
29 EmbeddingInput::Texts(texts)
30 }
31}
32
33impl From<PredictRequest> for EmbeddingInput {
34 fn from(request: PredictRequest) -> Self {
35 EmbeddingInput::PredictRequest(request)
36 }
37}
38
39#[instrument(skip_all)]
40pub fn modify_predict_request(
41 mut request: PredictRequest,
42 config: &EmbeddingConfig,
43) -> PredictRequest {
44 match config {
45 EmbeddingConfig::OpenAI(_) => request, EmbeddingConfig::Gemini(gemini_config) => {
47 if !gemini_config.is_configured {
51 return request;
52 }
53
54 let mut params = match &request.parameters {
56 serde_json::Value::Object(map) => map.clone(),
57 _ => serde_json::Map::new(),
58 };
59
60 if let Some(dim) = gemini_config.output_dimensionality {
62 params.insert(
63 "outputDimensionality".to_string(),
64 serde_json::Value::Number(serde_json::Number::from(dim)),
65 );
66 }
67
68 if let serde_json::Value::Array(ref mut instances) = request.instances {
69 for instance in instances.iter_mut() {
70 if let Some(task_type) = &gemini_config.task_type {
71 if let serde_json::Value::Object(ref mut map) = instance {
72 map.entry("task_type".to_string())
73 .or_insert_with(|| serde_json::json!(task_type));
74 }
75 }
76 }
77 }
78
79 request.parameters = serde_json::Value::Object(params);
80
81 debug!("Modified PredictRequest: {:?}", request);
82 request
83 }
84 }
85}
86
87#[derive(Debug, Clone, PartialEq, Serialize)]
88#[serde(untagged)]
89pub enum EmbeddingConfig {
90 OpenAI(OpenAIEmbeddingConfig),
91 Gemini(GeminiEmbeddingConfig),
92}
93
94impl EmbeddingConfig {
95 pub fn extract_config(
96 config: Option<&Bound<'_, PyAny>>,
97 provider: &Provider,
98 ) -> Result<Self, ProviderError> {
99 match provider {
100 Provider::OpenAI => {
101 let config = match config {
102 None => OpenAIEmbeddingConfig::default(),
103 Some(cfg) => cfg.extract::<OpenAIEmbeddingConfig>().map_err(|e| {
104 ProviderError::EmbeddingConfigExtractionError(format!(
105 "Failed to extract OpenAIEmbeddingConfig: {}",
106 e
107 ))
108 })?,
109 };
110
111 Ok(EmbeddingConfig::OpenAI(config))
112 }
113 Provider::Gemini | Provider::Vertex => {
114 let config = match config {
115 None => GeminiEmbeddingConfig::default(),
116 Some(cfg) => cfg.extract::<GeminiEmbeddingConfig>().map_err(|e| {
117 ProviderError::EmbeddingConfigExtractionError(format!(
118 "Failed to extract GeminiEmbeddingConfig: {}",
119 e
120 ))
121 })?,
122 };
123
124 Ok(EmbeddingConfig::Gemini(config))
125 }
126 _ => Err(ProviderError::ProviderNotSupportedError(
127 provider.to_string(),
128 )),
129 }
130 }
131
132 pub fn is_configured(&self) -> bool {
133 match self {
134 EmbeddingConfig::OpenAI(_config) => true,
136 EmbeddingConfig::Gemini(config) => config.is_configured,
137 }
138 }
139
140 pub fn get_vertex_config(&self) -> Result<serde_json::Value, ProviderError> {
141 match self {
142 EmbeddingConfig::Gemini(config) => Ok(config.get_parameters_for_predict()),
143 _ => Err(ProviderError::InvalidConfigType(
144 "Only Gemini config can be converted to Vertex config".to_string(),
145 )),
146 }
147 }
148}
149
150impl EmbeddingConfigTrait for EmbeddingConfig {
151 fn get_model(&self) -> &str {
152 match self {
153 EmbeddingConfig::OpenAI(config) => config.model.as_str(),
154 EmbeddingConfig::Gemini(config) => config.get_model(),
155 }
156 }
157}
158
159use tracing::error;
160#[derive(Debug, PartialEq)]
161pub struct Embedder {
162 client: GenAiClient,
163 config: EmbeddingConfig,
164 provider: Provider,
165}
166
167impl Embedder {
168 pub async fn new(provider: Provider, config: EmbeddingConfig) -> Result<Self, ProviderError> {
173 let client = match provider {
174 Provider::OpenAI => GenAiClient::OpenAI(OpenAIClient::new(ServiceType::Embed)?),
175 Provider::Gemini => GenAiClient::Gemini(GeminiClient::new(ServiceType::Embed).await?),
176 Provider::Vertex => GenAiClient::Vertex(VertexClient::new(ServiceType::Embed).await?),
177 _ => {
178 let msg = "No provider specified";
179 error!("{}", msg);
180 return Err(ProviderError::UndefinedError(msg.to_string()));
181 } };
183
184 Ok(Self {
185 client,
186 config,
187 provider,
188 })
189 }
190
191 pub async fn embed(&self, inputs: EmbeddingInput) -> Result<EmbeddingResponse, ProviderError> {
192 self.client.create_embedding(inputs, &self.config).await
194 }
195}
196
197pub enum EmbeddingResponse {
198 OpenAI(OpenAIEmbeddingResponse),
199 Gemini(GeminiEmbeddingResponse),
200 Vertex(PredictResponse),
201}
202
203impl EmbeddingResponse {
204 pub fn to_openai_response(&self) -> Result<&OpenAIEmbeddingResponse, ProviderError> {
205 match self {
206 EmbeddingResponse::OpenAI(response) => Ok(response),
207 _ => Err(ProviderError::InvalidResponseType("OpenAI".to_string())),
208 }
209 }
210
211 pub fn to_gemini_response(&self) -> Result<&GeminiEmbeddingResponse, ProviderError> {
212 match self {
213 EmbeddingResponse::Gemini(response) => Ok(response),
214 _ => Err(ProviderError::InvalidResponseType("Gemini".to_string())),
215 }
216 }
217
218 pub fn to_vertex_response(&self) -> Result<&PredictResponse, ProviderError> {
219 match self {
220 EmbeddingResponse::Vertex(response) => Ok(response),
221 _ => Err(ProviderError::InvalidResponseType("Vertex".to_string())),
222 }
223 }
224
225 pub fn into_py_bound_any<'py>(
226 &self,
227 py: Python<'py>,
228 ) -> Result<Bound<'py, PyAny>, ProviderError> {
229 match self {
230 EmbeddingResponse::OpenAI(response) => Ok(response.into_py_bound_any(py)?),
231 EmbeddingResponse::Gemini(response) => Ok(response.into_py_bound_any(py)?),
232 EmbeddingResponse::Vertex(response) => Ok(response.into_py_bound_any(py)?),
233 }
234 }
235
236 pub fn values(&self) -> Result<&Vec<f32>, ProviderError> {
237 match self {
238 EmbeddingResponse::OpenAI(response) => {
239 let first = response
240 .data
241 .first()
242 .ok_or_else(|| ProviderError::NoEmbeddingsFound)?;
243 Ok(&first.embedding)
244 }
245
246 EmbeddingResponse::Gemini(response) => Ok(&response.embedding.values),
247 _ => Err(ProviderError::InvalidResponseType(
248 "values not available for this response type".to_string(),
249 )),
250 }
251 }
252}
253
254#[pyclass(from_py_object, name = "Embedder")]
255#[derive(Debug, Clone)]
256pub struct PyEmbedder {
257 pub embedder: Arc<Embedder>,
258}
259
260#[pymethods]
261impl PyEmbedder {
262 #[new]
263 #[pyo3(signature = (provider, config=None))]
264 fn new(
265 provider: &Bound<'_, PyAny>,
266 config: Option<&Bound<'_, PyAny>>,
267 ) -> Result<Self, ProviderError> {
268 let provider = Provider::extract_provider(provider)?;
269 let config = EmbeddingConfig::extract_config(config, &provider)?;
270
271 let embedder = block_on(async { Embedder::new(provider, config).await })?;
272
273 Ok(Self {
274 embedder: Arc::new(embedder),
275 })
276 }
277
278 #[pyo3(signature = (input))]
283 pub fn embed<'py>(
284 &self,
285 py: Python<'py>,
286 input: Bound<'py, PyAny>,
287 ) -> Result<Bound<'py, PyAny>, ProviderError> {
288 let embedder = self.embedder.clone();
289
290 let embedding_input = if input.is_instance_of::<pyo3::types::PyList>() {
291 let texts: Vec<String> = input.extract()?;
292 EmbeddingInput::Texts(texts)
293 } else if input.is_instance_of::<pyo3::types::PyString>() {
294 let request: String = input.extract()?;
295 EmbeddingInput::Texts(vec![request])
296 } else if input.is_instance_of::<PredictRequest>() {
297 let request: PredictRequest = input.extract()?;
298 EmbeddingInput::PredictRequest(request)
299 } else {
300 return Err(ProviderError::InvalidInputType(
301 "Input must be a string, list of strings, or PredictRequest dict".to_string(),
302 ));
303 };
304
305 let embeddings = block_on(async { embedder.embed(embedding_input).await })?;
306 embeddings.into_py_bound_any(py)
307 }
308}