Skip to main content

potato_provider/providers/
embed.rs

1use 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/// Input types for embedding creation
21#[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, // OpenAI config does not apply to PredictRequest
46        EmbeddingConfig::Gemini(gemini_config) => {
47            // If configured, we need to modify the request
48            // If not configured, return early
49
50            if !gemini_config.is_configured {
51                return request;
52            }
53
54            // Handle parameters modification
55            let mut params = match &request.parameters {
56                serde_json::Value::Object(map) => map.clone(),
57                _ => serde_json::Map::new(),
58            };
59
60            // :predict endpoints expect dimensionality in parameters
61            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            // is configured only applies to Gemini and Vertex at the moment
135            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    /// Create a new Embedder instance that can be used to generate embeddings.
169    /// # Arguments
170    /// * `provider`: The provider to use for generating embeddings.
171    /// * `config`: The configuration for the embedding.
172    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            } // Add other providers here as needed
182        };
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        // Implementation for creating an embedding
193        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    /// Create a new embedding from a single input string
279    /// # Arguments
280    /// * `inputs`: The input string to embed.
281    /// * `config`: The configuration for the embedding.
282    #[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}