1use crate::{EmbeddingError, Embeddings};
5use async_trait::async_trait;
6use serde::Deserialize;
7
8pub const QWEN_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1";
10
11pub const QWEN_EMBED_MODEL: &str = "text-embedding-v1";
13
14#[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 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("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 pub fn with_model(mut self, model: impl Into<String>) -> Self {
72 self.model = model.into();
73 self
74 }
75}
76
77pub struct QwenEmbeddings {
79 config: QwenEmbeddingsConfig,
80 client: reqwest::Client,
81}
82
83impl QwenEmbeddings {
84 pub fn new(config: QwenEmbeddingsConfig) -> Self {
86 Self {
87 config,
88 client: reqwest::Client::new(),
89 }
90 }
91
92 #[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 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}