Skip to main content

lc_embeddings/
deepseek.rs

1// lc-embeddings/src/deepseek.rs
2//! DeepSeek embeddings implementation.
3
4use crate::{EmbeddingError, Embeddings};
5use async_trait::async_trait;
6use serde::Deserialize;
7
8/// Default base URL for the DeepSeek API.
9pub const DEEPSEEK_BASE_URL: &str = "https://api.deepseek.com/v1";
10
11/// Default embedding model for DeepSeek.
12pub const DEEPSEEK_EMBED_MODEL: &str = "deepseek-embedding";
13
14/// Configuration for DeepSeek embeddings API.
15#[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    /// Creates a new DeepSeekEmbeddingsConfig with the given API key.
34    pub fn new(api_key: impl Into<String>) -> Self {
35        Self {
36            api_key: api_key.into(),
37            ..Default::default()
38        }
39    }
40
41    /// Creates a DeepSeekEmbeddingsConfig from environment variables.
42    #[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    /// Creates a DeepSeekEmbeddingsConfig from environment variables, returning a Result.
52    ///
53    /// Environment variables:
54    /// - `DEEPSEEK_API_KEY`: API key (required)
55    /// - `DEEPSEEK_BASE_URL`: API endpoint (optional)
56    /// - `DEEPSEEK_EMBED_MODEL`: Model name (optional)
57    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    /// Sets the embedding model.
72    pub fn with_model(mut self, model: impl Into<String>) -> Self {
73        self.model = model.into();
74        self
75    }
76}
77
78/// DeepSeek embeddings client for generating vector embeddings.
79pub struct DeepSeekEmbeddings {
80    config: DeepSeekEmbeddingsConfig,
81    client: reqwest::Client,
82}
83
84impl DeepSeekEmbeddings {
85    /// Creates a DeepSeekEmbeddings with the given configuration.
86    pub fn new(config: DeepSeekEmbeddingsConfig) -> Self {
87        Self {
88            config,
89            client: reqwest::Client::new(),
90        }
91    }
92
93    /// Creates a DeepSeekEmbeddings from environment variables.
94    #[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    /// Creates a DeepSeekEmbeddings from environment variables, returning a Result.
104    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; // DeepSeek API batch limit
163        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}