Skip to main content

lc_embeddings/
qwen.rs

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