lc_embeddings/
deepseek.rs1use crate::{EmbeddingError, Embeddings};
5use async_trait::async_trait;
6use serde::Deserialize;
7
8pub const DEEPSEEK_BASE_URL: &str = "https://api.deepseek.com/v1";
10
11pub const DEEPSEEK_EMBED_MODEL: &str = "deepseek-embedding";
13
14#[derive(Debug, Clone)]
16pub struct DeepSeekEmbeddingsConfig {
17 pub api_key: String,
18 pub base_url: String,
19 pub model: String,
20}
21
22impl Default for DeepSeekEmbeddingsConfig {
23 fn default() -> Self {
24 Self {
25 api_key: std::env::var("DEEPSEEK_API_KEY").unwrap_or_default(),
26 base_url: DEEPSEEK_BASE_URL.to_string(),
27 model: DEEPSEEK_EMBED_MODEL.to_string(),
28 }
29 }
30}
31
32impl DeepSeekEmbeddingsConfig {
33 pub fn new(api_key: impl Into<String>) -> Self {
35 Self {
36 api_key: api_key.into(),
37 ..Default::default()
38 }
39 }
40
41 #[deprecated(
43 since = "0.7.0",
44 note = "Use from_env_result() which returns Result<Self, String>"
45 )]
46 #[allow(deprecated)]
47 pub fn from_env() -> Self {
48 Self::from_env_result().unwrap_or_else(|_| Self::default())
49 }
50
51 pub fn from_env_result() -> Result<Self, String> {
58 let api_key = std::env::var("DEEPSEEK_API_KEY")
59 .map_err(|_| "DEEPSEEK_API_KEY environment variable not set".to_string())?;
60 let base_url =
61 std::env::var("DEEPSEEK_BASE_URL").unwrap_or_else(|_| DEEPSEEK_BASE_URL.to_string());
62 let model = std::env::var("DEEPSEEK_EMBED_MODEL")
63 .unwrap_or_else(|_| DEEPSEEK_EMBED_MODEL.to_string());
64 Ok(Self {
65 api_key,
66 base_url,
67 model,
68 })
69 }
70
71 pub fn with_model(mut self, model: impl Into<String>) -> Self {
73 self.model = model.into();
74 self
75 }
76}
77
78pub struct DeepSeekEmbeddings {
80 config: DeepSeekEmbeddingsConfig,
81 client: reqwest::Client,
82}
83
84impl DeepSeekEmbeddings {
85 pub fn new(config: DeepSeekEmbeddingsConfig) -> Self {
87 Self {
88 config,
89 client: reqwest::Client::new(),
90 }
91 }
92
93 #[deprecated(
95 since = "0.7.0",
96 note = "Use from_env_result() which returns Result<Self, String>"
97 )]
98 #[allow(deprecated)]
99 pub fn from_env() -> Self {
100 Self::from_env_result().unwrap_or_else(|_| Self::new(DeepSeekEmbeddingsConfig::default()))
101 }
102
103 pub fn from_env_result() -> Result<Self, String> {
105 let config = DeepSeekEmbeddingsConfig::from_env_result()?;
106 Ok(Self::new(config))
107 }
108}
109
110#[async_trait]
111impl Embeddings for DeepSeekEmbeddings {
112 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
113 if text.is_empty() {
114 return Err(EmbeddingError::EmptyInput);
115 }
116
117 let url = format!("{}/embeddings", self.config.base_url);
118
119 let body = serde_json::json!({
120 "model": self.config.model,
121 "input": text,
122 });
123
124 let response = self
125 .client
126 .post(&url)
127 .header("Authorization", format!("Bearer {}", self.config.api_key))
128 .header("Content-Type", "application/json")
129 .json(&body)
130 .send()
131 .await
132 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
133
134 let status = response.status();
135 if !status.is_success() {
136 let error_text = response.text().await.unwrap_or_default();
137 return Err(EmbeddingError::ApiError(format!(
138 "HTTP {}: {}",
139 status, error_text
140 )));
141 }
142
143 let embedding_response: EmbeddingResponse = response
144 .json()
145 .await
146 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
147
148 Ok(embedding_response
149 .data
150 .first()
151 .ok_or_else(|| EmbeddingError::ApiError("No embedding data in response".to_string()))?
152 .embedding
153 .clone())
154 }
155
156 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
157 if texts.is_empty() {
158 return Ok(Vec::new());
159 }
160
161 let url = format!("{}/embeddings", self.config.base_url);
162 let batch_size = 64; let mut all_results = vec![Vec::new(); texts.len()];
164 let mut offset = 0;
165
166 for chunk in texts.chunks(batch_size) {
167 let body = serde_json::json!({
168 "model": self.config.model,
169 "input": chunk,
170 });
171
172 let response = self
173 .client
174 .post(&url)
175 .header("Authorization", format!("Bearer {}", self.config.api_key))
176 .header("Content-Type", "application/json")
177 .json(&body)
178 .send()
179 .await
180 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
181
182 let status = response.status();
183 if !status.is_success() {
184 let error_text = response.text().await.unwrap_or_default();
185 return Err(EmbeddingError::ApiError(format!(
186 "HTTP {}: {}",
187 status, error_text
188 )));
189 }
190
191 let embedding_response: EmbeddingResponse = response
192 .json()
193 .await
194 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
195
196 for item in embedding_response.data {
197 let global_index = offset + item.index as usize;
198 if global_index < all_results.len() {
199 all_results[global_index] = item.embedding;
200 }
201 }
202 offset += chunk.len();
203 }
204
205 Ok(all_results)
206 }
207
208 fn dimension(&self) -> usize {
209 1536
210 }
211
212 fn model_name(&self) -> &str {
213 &self.config.model
214 }
215}
216
217#[derive(Debug, Deserialize)]
218struct EmbeddingResponse {
219 data: Vec<EmbeddingData>,
220}
221
222#[derive(Debug, Deserialize)]
223struct EmbeddingData {
224 embedding: Vec<f32>,
225 index: i32,
226}
227
228#[cfg(test)]
229mod tests {
230 use super::*;
231 use std::env;
232
233 fn save_and_set(key: &str, value: &str) -> Option<String> {
234 let old = env::var(key).ok();
235 env::set_var(key, value);
236 old
237 }
238
239 fn restore(key: &str, old: Option<String>) {
240 match old {
241 Some(v) => env::set_var(key, v),
242 None => env::remove_var(key),
243 }
244 }
245
246 #[test]
247 fn test_from_env_result_ok_when_key_set() {
248 let _lock = crate::ENV_TEST_LOCK
249 .lock()
250 .unwrap_or_else(|e| e.into_inner());
251 let old = save_and_set("DEEPSEEK_API_KEY", "test-key-123");
252 let result = DeepSeekEmbeddingsConfig::from_env_result();
253 assert!(result.is_ok());
254 assert_eq!(result.unwrap().api_key, "test-key-123");
255 restore("DEEPSEEK_API_KEY", old);
256 }
257
258 #[test]
259 fn test_from_env_result_err_when_key_missing() {
260 let _lock = crate::ENV_TEST_LOCK
261 .lock()
262 .unwrap_or_else(|e| e.into_inner());
263 let old = env::var("DEEPSEEK_API_KEY").ok();
264 env::remove_var("DEEPSEEK_API_KEY");
265 let result = DeepSeekEmbeddingsConfig::from_env_result();
266 assert!(result.is_err());
267 assert!(result.unwrap_err().contains("DEEPSEEK_API_KEY"));
268 restore("DEEPSEEK_API_KEY", old);
269 }
270
271 #[test]
272 fn test_from_env_result_uses_optional_vars() {
273 let _lock = crate::ENV_TEST_LOCK
274 .lock()
275 .unwrap_or_else(|e| e.into_inner());
276 let old_key = save_and_set("DEEPSEEK_API_KEY", "key");
277 let old_url = save_and_set("DEEPSEEK_BASE_URL", "https://custom.api.com");
278 let old_model = save_and_set("DEEPSEEK_EMBED_MODEL", "custom-model");
279 let config = DeepSeekEmbeddingsConfig::from_env_result().unwrap();
280 assert_eq!(config.base_url, "https://custom.api.com");
281 assert_eq!(config.model, "custom-model");
282 restore("DEEPSEEK_API_KEY", old_key);
283 restore("DEEPSEEK_BASE_URL", old_url);
284 restore("DEEPSEEK_EMBED_MODEL", old_model);
285 }
286
287 #[test]
288 fn test_from_env_result_uses_defaults_for_optional_vars() {
289 let _lock = crate::ENV_TEST_LOCK
290 .lock()
291 .unwrap_or_else(|e| e.into_inner());
292 let old_key = save_and_set("DEEPSEEK_API_KEY", "key");
293 let old_url = env::var("DEEPSEEK_BASE_URL").ok();
294 env::remove_var("DEEPSEEK_BASE_URL");
295 let old_model = env::var("DEEPSEEK_EMBED_MODEL").ok();
296 env::remove_var("DEEPSEEK_EMBED_MODEL");
297 let config = DeepSeekEmbeddingsConfig::from_env_result().unwrap();
298 assert_eq!(config.base_url, DEEPSEEK_BASE_URL.to_string());
299 assert_eq!(config.model, DEEPSEEK_EMBED_MODEL);
300 restore("DEEPSEEK_API_KEY", old_key);
301 restore("DEEPSEEK_BASE_URL", old_url);
302 restore("DEEPSEEK_EMBED_MODEL", old_model);
303 }
304
305 #[test]
306 fn test_embeddings_from_env_result_ok() {
307 let _lock = crate::ENV_TEST_LOCK
308 .lock()
309 .unwrap_or_else(|e| e.into_inner());
310 let old = save_and_set("DEEPSEEK_API_KEY", "test-key");
311 assert!(DeepSeekEmbeddings::from_env_result().is_ok());
312 restore("DEEPSEEK_API_KEY", old);
313 }
314
315 #[test]
316 fn test_embeddings_from_env_result_err_when_key_missing() {
317 let _lock = crate::ENV_TEST_LOCK
318 .lock()
319 .unwrap_or_else(|e| e.into_inner());
320 let old = env::var("DEEPSEEK_API_KEY").ok();
321 env::remove_var("DEEPSEEK_API_KEY");
322 assert!(DeepSeekEmbeddings::from_env_result().is_err());
323 restore("DEEPSEEK_API_KEY", old);
324 }
325}