code_repo_wiki/generate/
embed.rs1use std::sync::atomic::{AtomicUsize, Ordering};
2
3use anyhow::{Context, Result};
4
5use crate::analysis::feature::Embedder;
6use crate::config::schema::EmbedSection;
7
8pub struct EmbeddingEngine {
12 client: reqwest::Client,
13 config: EmbedSection,
14 call_count: AtomicUsize,
15 rt: tokio::runtime::Handle,
17}
18
19impl EmbeddingEngine {
20 pub fn new(config: &EmbedSection, rt: tokio::runtime::Handle) -> Result<Self> {
25 let client = reqwest::Client::builder()
26 .timeout(std::time::Duration::from_secs(60))
27 .build()
28 .context("创建 Embedding HTTP 客户端失败")?;
29 Ok(Self {
30 client,
31 config: config.clone(),
32 call_count: AtomicUsize::new(0),
33 rt,
34 })
35 }
36
37 fn resolve_api_key(&self) -> Result<String> {
39 self.config
40 .api_key
41 .clone()
42 .or_else(|| std::env::var(&self.config.api_key_env).ok())
43 .context(format!(
44 "Embedding API Key 未设置(api_key 为空且环境变量 {} 未定义)",
45 self.config.api_key_env
46 ))
47 }
48
49 fn resolve_base_url(&self) -> String {
51 self.config
52 .base_url
53 .clone()
54 .unwrap_or_else(|| "https://api.openai.com/v1".to_string())
55 }
56
57 pub async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
61 let api_key = self.resolve_api_key()?;
62 let url = format!("{}/embeddings", self.resolve_base_url());
63
64 let mut all_embeddings = Vec::with_capacity(texts.len());
65
66 for chunk in texts.chunks(crate::config::schema::EMBED_BATCH_SIZE) {
67 let body = serde_json::json!({
68 "model": self.config.model,
69 "input": chunk,
70 });
71
72 let resp = crate::generate::llm::retry_with_backoff(
76 crate::generate::llm::MAX_RETRIES,
77 || {
78 let body = &body;
79 let url = &url;
80 let api_key = &api_key;
81 let client = &self.client;
82 async move {
83 client
84 .post(url)
85 .bearer_auth(api_key)
86 .json(body)
87 .send()
88 .await
89 }
90 },
91 )
92 .await
93 .with_context(|| "Embedding API 请求失败")?;
94
95 if !resp.status().is_success() {
96 let status = resp.status();
97 let text = resp.text().await.unwrap_or_default();
98 anyhow::bail!("Embedding API 返回错误 ({}): {}", status, text);
99 }
100
101 self.call_count.fetch_add(1, Ordering::Relaxed);
102
103 let data: serde_json::Value = resp
104 .json()
105 .await
106 .context("解析 Embedding API 响应 JSON 失败")?;
107
108 let embeddings = data["data"]
109 .as_array()
110 .context("Embedding 响应缺少 data 字段")?
111 .iter()
112 .map(|item| {
113 let arr = item["embedding"]
114 .as_array()
115 .context("嵌入向量缺失")?;
116 arr.iter()
120 .map(|v| {
121 v.as_f64()
122 .map(|f| f as f32)
123 .with_context(|| "嵌入向量包含非数字元素(模型输出异常,拒绝静默丢弃)")
124 })
125 .collect::<Result<Vec<f32>>>()
126 })
127 .collect::<Result<Vec<_>>>()?;
128
129 if embeddings.len() != chunk.len() {
134 anyhow::bail!(
135 "Embedding 响应数量不匹配:请求 {} 条,返回 {} 条",
136 chunk.len(),
137 embeddings.len()
138 );
139 }
140 if let Some(first) = embeddings.first() {
141 let dim = first.len();
142 if let Some(bad) = embeddings.iter().find(|v| v.len() != dim) {
143 anyhow::bail!(
144 "Embedding 响应维度不一致:{} 维与 {} 维并存",
145 dim,
146 bad.len()
147 );
148 }
149 }
150
151 all_embeddings.extend(embeddings);
152 }
153
154 Ok(all_embeddings)
155 }
156
157 pub async fn embed(&self, text: &str) -> Result<Vec<f32>> {
159 let results = self.embed_batch(&[text.to_string()]).await?;
160 results.into_iter().next().context("Embedding 返回空结果")
161 }
162
163 pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
165 if a.len() != b.len() || a.is_empty() {
166 return 0.0;
167 }
168 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
169 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
170 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
171 if norm_a == 0.0 || norm_b == 0.0 {
172 return 0.0;
173 }
174 (dot / (norm_a * norm_b)).clamp(-1.0, 1.0)
175 }
176
177 pub fn call_count(&self) -> usize {
179 self.call_count.load(Ordering::Relaxed)
180 }
181}
182
183impl Embedder for EmbeddingEngine {
187 fn embed(&self, text: &str) -> Result<Vec<f32>> {
188 self.rt.block_on(EmbeddingEngine::embed(self, text))
189 }
190
191 fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
192 self.rt
193 .block_on(EmbeddingEngine::embed_batch(self, texts))
194 }
195
196 fn cosine_similarity(&self, a: &[f32], b: &[f32]) -> f64 {
197 EmbeddingEngine::cosine_similarity(a, b) as f64
198 }
199}
200
201#[cfg(test)]
202mod tests {
203 use super::*;
204
205 #[test]
206 fn test_cosine_similarity_identical() {
207 let a = vec![1.0, 0.0, 0.0];
208 let sim = EmbeddingEngine::cosine_similarity(&a, &a);
209 assert!((sim - 1.0).abs() < 1e-6);
210 }
211
212 #[test]
213 fn test_cosine_similarity_orthogonal() {
214 let a = vec![1.0, 0.0];
215 let b = vec![0.0, 1.0];
216 let sim = EmbeddingEngine::cosine_similarity(&a, &b);
217 assert!((sim - 0.0).abs() < 1e-6);
218 }
219
220 #[test]
221 fn test_cosine_similarity_opposite() {
222 let a = vec![1.0, 0.0];
223 let b = vec![-1.0, 0.0];
224 let sim = EmbeddingEngine::cosine_similarity(&a, &b);
225 assert!((sim + 1.0).abs() < 1e-6);
226 }
227
228 #[test]
229 fn test_cosine_similarity_zero_vector() {
230 let a = vec![0.0, 0.0];
231 let b = vec![1.0, 0.0];
232 let sim = EmbeddingEngine::cosine_similarity(&a, &b);
233 assert!((sim - 0.0).abs() < 1e-6);
234 }
235
236 #[test]
237 fn test_cosine_similarity_empty() {
238 let sim = EmbeddingEngine::cosine_similarity(&[], &[]);
239 assert!((sim - 0.0).abs() < 1e-6);
240 }
241
242 #[test]
243 fn test_cosine_similarity_mismatched_length() {
244 let a = vec![1.0, 0.0];
245 let b = vec![1.0];
246 let sim = EmbeddingEngine::cosine_similarity(&a, &b);
247 assert!((sim - 0.0).abs() < 1e-6);
248 }
249}