(
&self,
request: EmbeddingRequest,
)
| 237 | #[async_trait] |
| 238 | impl EmbeddingProviderTrait for OpenAIEmbeddingProvider { |
| 239 | async fn generate_embeddings( |
| 240 | &self, |
| 241 | request: EmbeddingRequest, |
| 242 | ) -> GraphBitResult<EmbeddingResponse> { |
| 243 | let url = format!("{}/embeddings", self.base_url()); |
| 244 | |
| 245 | let input = match &request.input { |
| 246 | EmbeddingInput::Single(text) => serde_json::Value::String(text.clone()), |
| 247 | EmbeddingInput::Multiple(texts) => serde_json::Value::Array( |
| 248 | texts |
| 249 | .iter() |
| 250 | .map(|t| serde_json::Value::String(t.clone())) |
| 251 | .collect(), |
| 252 | ), |
| 253 | }; |
| 254 | |
| 255 | let mut body = serde_json::json!({ |
| 256 | "model": self.config.model, |
| 257 | "input": input, |
| 258 | }); |
| 259 | |
| 260 | // Add user if provided |
| 261 | if let Some(user) = &request.user { |
| 262 | body["user"] = serde_json::Value::String(user.clone()); |
| 263 | } |
| 264 | |
| 265 | // Add extra parameters |
| 266 | for (key, value) in &request.params { |
| 267 | body[key] = value.clone(); |
| 268 | } |
| 269 | |
| 270 | let response = self |
| 271 | .client |
| 272 | .post(&url) |
| 273 | .header("Authorization", format!("Bearer {}", self.config.api_key)) |
| 274 | .header("Content-Type", "application/json") |
| 275 | .json(&body) |
| 276 | .send() |
| 277 | .await |
| 278 | .map_err(|e| GraphBitError::llm(format!("Failed to send request to OpenAI: {e}")))?; |
| 279 | |
| 280 | if !response.status().is_success() { |
| 281 | let error_text = response |
| 282 | .text() |
| 283 | .await |
| 284 | .unwrap_or_else(|_| "Unknown error".to_string()); |
| 285 | return Err(GraphBitError::llm(format!( |
| 286 | "OpenAI API error: {error_text}" |
| 287 | ))); |
| 288 | } |
| 289 | |
| 290 | let response_json: serde_json::Value = response |
| 291 | .json() |
| 292 | .await |
| 293 | .map_err(|e| GraphBitError::llm(format!("Failed to parse OpenAI response: {e}")))?; |
| 294 | |
| 295 | // Parse embeddings |
| 296 | let embeddings_data = response_json["data"] |
no test coverage detected