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;
16 fn base_url(&self) -> &str;
18 fn model(&self) -> &str;
20}
21
22pub trait CompatSpec: CompatConfigAccess + Sized + Default {
27 fn api_key_env() -> &'static str;
29 fn batch_size() -> usize;
31 fn dimension_for(model: &str) -> Result<usize, EmbeddingError>;
33 fn from_env_result() -> Result<Self, EmbeddingError>;
35}
36
37pub struct OpenAICompatEmbeddings<C: CompatConfigAccess + CompatSpec> {
42 config: C,
43 client: reqwest::Client,
44 dimension: usize,
45}
46
47impl<C: CompatConfigAccess + CompatSpec> std::fmt::Debug for OpenAICompatEmbeddings<C> {
48 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
49 f.debug_struct("OpenAICompatEmbeddings")
50 .field("model", &self.config.model())
51 .field("dimension", &self.dimension)
52 .finish()
53 }
54}
55
56impl<C: CompatConfigAccess + CompatSpec> OpenAICompatEmbeddings<C> {
57 pub fn new(config: C) -> Result<Self, EmbeddingError> {
60 if config.api_key().trim().is_empty() {
61 return Err(EmbeddingError::Config(format!(
62 "{} is empty",
63 C::api_key_env()
64 )));
65 }
66 let dimension = C::dimension_for(config.model())?;
67 Ok(Self {
68 config,
69 client: reqwest::Client::new(),
70 dimension,
71 })
72 }
73
74 pub fn from_env_result() -> Result<Self, EmbeddingError> {
76 let config = C::from_env_result()?;
77 Self::new(config)
78 }
79}
80
81#[async_trait]
82impl<C: CompatConfigAccess + CompatSpec + Send + Sync> Embeddings for OpenAICompatEmbeddings<C> {
83 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
84 if text.trim().is_empty() {
85 return Err(EmbeddingError::EmptyInput);
86 }
87
88 let url = format!("{}/embeddings", self.config.base_url());
89
90 let body = serde_json::json!({
91 "model": self.config.model(),
92 "input": text,
93 });
94
95 let response = crate::retry::post_json_with_retry(
97 &self.client,
98 &url,
99 self.config.api_key(),
100 &body,
101 &crate::retry::DEFAULT_RETRY,
102 )
103 .await
104 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
105
106 let status = response.status();
107 if !status.is_success() {
108 let error_text = response.text().await.map_err(|e| {
110 EmbeddingError::HttpError(format!("failed to read error response body: {e}"))
111 })?;
112 return Err(EmbeddingError::ApiError(format!(
113 "HTTP {}: {}",
114 status, error_text
115 )));
116 }
117
118 let embedding_response: EmbeddingResponse = response
119 .json()
120 .await
121 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
122
123 let mut embedding = embedding_response
124 .data
125 .first()
126 .ok_or_else(|| EmbeddingError::ApiError("No embedding data in response".to_string()))?
127 .embedding
128 .clone();
129 crate::l2_normalize(&mut embedding);
131 Ok(embedding)
132 }
133
134 async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
135 if texts.is_empty() {
137 return Ok(Vec::new());
138 }
139 if texts.iter().any(|t| t.trim().is_empty()) {
140 return Err(EmbeddingError::EmptyInput);
141 }
142
143 let url = format!("{}/embeddings", self.config.base_url());
144 let batch_size = C::batch_size().max(1);
145 let mut all_results: Vec<Option<Vec<f32>>> = vec![None; texts.len()];
148 let mut offset = 0;
149
150 for chunk in texts.chunks(batch_size) {
151 let body = serde_json::json!({
152 "model": self.config.model(),
153 "input": chunk,
154 });
155
156 let response = crate::retry::post_json_with_retry(
158 &self.client,
159 &url,
160 self.config.api_key(),
161 &body,
162 &crate::retry::DEFAULT_RETRY,
163 )
164 .await
165 .map_err(|e| EmbeddingError::HttpError(e.to_string()))?;
166
167 let status = response.status();
168 if !status.is_success() {
169 let error_text = response.text().await.map_err(|e| {
171 EmbeddingError::HttpError(format!("failed to read error response body: {e}"))
172 })?;
173 return Err(EmbeddingError::ApiError(format!(
174 "HTTP {}: {}",
175 status, error_text
176 )));
177 }
178
179 let embedding_response: EmbeddingResponse = response
180 .json()
181 .await
182 .map_err(|e| EmbeddingError::ParseError(e.to_string()))?;
183
184 for item in embedding_response.data {
185 let global_index = offset + item.index as usize;
186 if global_index >= all_results.len() {
187 return Err(EmbeddingError::BatchMismatch {
189 expected: all_results.len(),
190 actual: global_index + 1,
191 });
192 }
193 all_results[global_index] = Some(item.embedding);
194 }
195 offset += chunk.len();
196 }
197
198 all_results
200 .into_iter()
201 .map(|opt| {
202 let mut v = opt.ok_or(EmbeddingError::EmptyVectorInBatch)?;
203 crate::l2_normalize(&mut v);
204 Ok(v)
205 })
206 .collect()
207 }
208
209 fn dimension(&self) -> usize {
210 self.dimension
211 }
212
213 fn model_name(&self) -> &str {
214 self.config.model()
215 }
216}
217
218#[derive(Debug, Deserialize)]
220struct EmbeddingResponse {
221 data: Vec<EmbeddingData>,
222}
223
224#[derive(Debug, Deserialize)]
225struct EmbeddingData {
226 embedding: Vec<f32>,
227 index: i32,
228}