lc_embeddings/
openai_compat.rs1use crate::{EmbeddingError, Embeddings};
9use async_trait::async_trait;
10use serde::Deserialize;
11
12pub trait CompatConfigAccess {
14 fn api_key(&self) -> &str;
15 fn base_url(&self) -> &str;
16 fn model(&self) -> &str;
17}
18
19pub trait CompatSpec: CompatConfigAccess + Sized + Default {
24 fn api_key_env() -> &'static str;
26 fn batch_size() -> usize;
28 fn dimension_for(model: &str) -> Result<usize, EmbeddingError>;
30 fn from_env_result() -> Result<Self, String>;
32}
33
34pub 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 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 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 #[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 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 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 crate::l2_normalize(&mut embedding);
139 Ok(embedding)
140 }
141
142 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
143 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 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 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 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 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 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#[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}