Skip to main content

lc_embeddings/
openai_compat.rs

1// lc-embeddings/src/openai_compat.rs
2//! OpenAI 兼容协议 embedding 客户端公共基类(P1-5)。
3//!
4//! DeepSeek 与 Qwen 走同一套 OpenAI `/embeddings` 协议(同样的请求体、同样的
5//! `data[index]` 对齐、同样的 Bearer 认证),二者源码几乎逐行重复。本模块抽出
6//! 通用实现,DeepSeek/Qwen 仅通过 [`CompatSpec`] 配置 URL / 模型 / 维度 / 批量大小。
7
8use crate::{EmbeddingError, Embeddings};
9use async_trait::async_trait;
10use serde::Deserialize;
11
12/// 访问 provider 配置字段的抽象——DeepSeek/Qwen 的 config 结构体字段名相同。
13pub trait CompatConfigAccess {
14    fn api_key(&self) -> &str;
15    fn base_url(&self) -> &str;
16    fn model(&self) -> &str;
17}
18
19/// OpenAI 兼容协议 provider 的静态规格。
20///
21/// 实现该 trait 即可获得 `OpenAICompatEmbeddings` 提供的完整 embedding 能力,
22/// 是接入新增 OpenAI 兼容 provider 的扩展点。
23pub trait CompatSpec: CompatConfigAccess + Sized + Default {
24    /// 环境变量名:API key(用于构造期错误信息,P1-3)。
25    fn api_key_env() -> &'static str;
26    /// 单次请求的批量上限。
27    fn batch_size() -> usize;
28    /// 给定模型的向量维度;未知模型必须报错(P1-2),不得回落默认值撒谎。
29    fn dimension_for(model: &str) -> Result<usize, EmbeddingError>;
30    /// 从环境变量构造 config(复用各 config 已实现的 from_env_result)。
31    fn from_env_result() -> Result<Self, String>;
32}
33
34/// 通用 OpenAI 兼容 embedding 客户端(P1-5)。
35///
36/// DeepSeek/Qwen 等走 OpenAI `/embeddings` 协议的 provider 通过
37/// [`CompatSpec`] 配置规格复用本实现;`C` 即各自的 config 类型。
38pub struct OpenAICompatEmbeddings<C: CompatConfigAccess + CompatSpec> {
39    config: C,
40    client: reqwest::Client,
41    dimension: usize,
42}
43
44impl<C: CompatConfigAccess + CompatSpec> std::fmt::Debug for OpenAICompatEmbeddings<C> {
45    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46        f.debug_struct("OpenAICompatEmbeddings")
47            .field("model", &self.config.model())
48            .field("dimension", &self.dimension)
49            .finish()
50    }
51}
52
53impl<C: CompatConfigAccess + CompatSpec> OpenAICompatEmbeddings<C> {
54    /// 构造时 fail fast(P1-3):API key 为空立即报错,而不是拖到发请求才 401;
55    /// 同时校验模型维度已知(P1-2)。
56    pub fn new(config: C) -> Result<Self, EmbeddingError> {
57        if config.api_key().trim().is_empty() {
58            return Err(EmbeddingError::Config(format!(
59                "{} is empty",
60                C::api_key_env()
61            )));
62        }
63        let dimension = C::dimension_for(config.model())?;
64        Ok(Self {
65            config,
66            client: reqwest::Client::new(),
67            dimension,
68        })
69    }
70
71    /// Creates from environment variables, returning a Result.
72    pub fn from_env_result() -> Result<Self, String> {
73        let config = C::from_env_result()?;
74        Self::new(config).map_err(|e| e.to_string())
75    }
76
77    /// Creates from environment variables.
78    #[deprecated(
79        since = "0.7.0",
80        note = "Use from_env_result() which returns Result<Self, String>"
81    )]
82    #[allow(deprecated)]
83    pub fn from_env() -> Self {
84        Self::from_env_result()
85            .unwrap_or_else(|_| Self::new(C::default()).expect("from_env(): missing API key"))
86    }
87}
88
89#[async_trait]
90impl<C: CompatConfigAccess + CompatSpec + Send + Sync> Embeddings for OpenAICompatEmbeddings<C> {
91    async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
92        if text.trim().is_empty() {
93            return Err(EmbeddingError::EmptyInput);
94        }
95
96        let url = format!("{}/embeddings", self.config.base_url());
97
98        let body = serde_json::json!({
99            "model": self.config.model(),
100            "input": text,
101        });
102
103        // P2-5: 429/5xx 指数退避重试。
104        let response = crate::retry::post_json_with_retry(
105            &self.client,
106            &url,
107            self.config.api_key(),
108            &body,
109            &crate::retry::DEFAULT_RETRY,
110        )
111        .await
112        .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
113
114        let status = response.status();
115        if !status.is_success() {
116            // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
117            let error_text = response.text().await.map_err(|e| {
118                EmbeddingError::HttpError(format!("failed to read error response body: {e}"))
119            })?;
120            return Err(EmbeddingError::ApiError(format!(
121                "HTTP {}: {}",
122                status, error_text
123            )));
124        }
125
126        let embedding_response: EmbeddingResponse = response
127            .json()
128            .await
129            .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
130
131        let mut embedding = embedding_response
132            .data
133            .first()
134            .ok_or_else(|| EmbeddingError::ApiError("No embedding data in response".to_string()))?
135            .embedding
136            .clone();
137        // P2-8: 统一 L2 归一化,保证单位长度。
138        crate::l2_normalize(&mut embedding);
139        Ok(embedding)
140    }
141
142    async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
143        // P1-1: 空切片不是错误(无事可做),含空/全空白文本才报错。
144        if texts.is_empty() {
145            return Ok(Vec::new());
146        }
147        if texts.iter().any(|t| t.trim().is_empty()) {
148            return Err(EmbeddingError::EmptyInput);
149        }
150
151        let url = format!("{}/embeddings", self.config.base_url());
152        let batch_size = C::batch_size().max(1);
153        // P0-1: 用 Option 槽位逐项收集,拒绝静默空向量。某 chunk 少返回/错位会
154        // 留下 None 槽位并在收尾时报错,而不是产出零向量被下游当成"不相似"。
155        let mut all_results: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
156        let mut offset = 0;
157
158        for chunk in texts.chunks(batch_size) {
159            let body = serde_json::json!({
160                "model": self.config.model(),
161                "input": chunk,
162            });
163
164            // P2-5: 429/5xx 指数退避重试。
165            let response = crate::retry::post_json_with_retry(
166                &self.client,
167                &url,
168                self.config.api_key(),
169                &body,
170                &crate::retry::DEFAULT_RETRY,
171            )
172            .await
173            .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
174
175            let status = response.status();
176            if !status.is_success() {
177                // P1-4: 读失败的错误体也要报错,不能 unwrap_or_default() 吞掉。
178                let error_text = response.text().await.map_err(|e| {
179                    EmbeddingError::HttpError(format!("failed to read error response body: {e}"))
180                })?;
181                return Err(EmbeddingError::ApiError(format!(
182                    "HTTP {}: {}",
183                    status, error_text
184                )));
185            }
186
187            let embedding_response: EmbeddingResponse = response
188                .json()
189                .await
190                .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
191
192            for item in embedding_response.data {
193                let global_index = offset + item.index as usize;
194                if global_index >= all_results.len() {
195                    // 服务端 index 超出请求范围 = 批次错位,直接报错。
196                    return Err(EmbeddingError::BatchMismatch {
197                        expected: all_results.len(),
198                        actual: global_index + 1,
199                    });
200                }
201                all_results[global_index] = Some(item.embedding);
202            }
203            offset += chunk.len();
204        }
205
206        // 展开为 Result:任一槽位空缺即显式报错,而非留下零向量;并统一 L2 归一化(P2-8)。
207        all_results
208            .into_iter()
209            .map(|opt| {
210                let mut v = opt.ok_or(EmbeddingError::EmptyVectorInBatch)?;
211                crate::l2_normalize(&mut v);
212                Ok(v)
213            })
214            .collect()
215    }
216
217    fn dimension(&self) -> usize {
218        self.dimension
219    }
220
221    fn model_name(&self) -> &str {
222        self.config.model()
223    }
224}
225
226/// OpenAI 兼容协议的 embedding 响应体(DeepSeek/Qwen 共用)。
227#[derive(Debug, Deserialize)]
228struct EmbeddingResponse {
229    data: Vec<EmbeddingData>,
230}
231
232#[derive(Debug, Deserialize)]
233struct EmbeddingData {
234    embedding: Vec<f32>,
235    index: i32,
236}