1use std::sync::OnceLock;
14use std::time::Duration;
15
16use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE};
17use reqwest::{Client, StatusCode};
18use serde::{Deserialize, Serialize};
19use tokio::runtime::Runtime;
20
21use super::{Embedding, EmbeddingProvider};
22use crate::error::{CtxError, Result};
23
24const DEFAULT_HOST: &str = "http://localhost:11434";
26
27const DEFAULT_MODEL: &str = "nomic-embed-text";
29
30const REQUEST_TIMEOUT_SECS: u64 = 60;
31const CONNECT_TIMEOUT_SECS: u64 = 10;
32const MAX_RETRIES: u32 = 3;
33const RETRY_BASE_DELAY_MS: u64 = 500;
34
35static GLOBAL_RUNTIME: OnceLock<Runtime> = OnceLock::new();
37
38fn get_or_create_runtime() -> &'static Runtime {
39 GLOBAL_RUNTIME.get_or_init(|| {
40 Runtime::new().expect("Failed to create global tokio runtime for Ollama provider")
41 })
42}
43
44#[derive(Serialize)]
46struct OllamaEmbedRequest<'a> {
47 model: &'a str,
48 input: Vec<&'a str>,
49}
50
51#[derive(Deserialize)]
53struct OllamaEmbedResponse {
54 embeddings: Option<Vec<Vec<f32>>>,
55 error: Option<String>,
56}
57
58pub struct OllamaProvider {
60 client: Client,
61 host: String,
62 model: String,
63 dimension: usize,
65}
66
67fn resolve_host(config_host: Option<&str>) -> String {
70 let raw = std::env::var("OLLAMA_HOST")
71 .ok()
72 .filter(|h| !h.is_empty())
73 .or_else(|| config_host.map(str::to_string));
74 match raw {
75 Some(h) if h.starts_with("http://") || h.starts_with("https://") => h,
76 Some(h) => format!("http://{}", h),
77 None => DEFAULT_HOST.to_string(),
78 }
79}
80
81fn resolve_model(config_model: Option<&str>) -> String {
83 std::env::var("OLLAMA_EMBED_MODEL")
84 .ok()
85 .filter(|m| !m.is_empty())
86 .or_else(|| config_model.map(str::to_string))
87 .unwrap_or_else(|| DEFAULT_MODEL.to_string())
88}
89
90impl OllamaProvider {
91 pub fn from_env() -> Result<Self> {
95 Self::from_config(None, None)
96 }
97
98 pub fn from_config(config_model: Option<&str>, config_host: Option<&str>) -> Result<Self> {
101 let mut provider = Self::new_unprobed(config_model, config_host)?;
102 let probe = provider.request(&["dimension probe"])?;
103 provider.dimension = Self::dimension_from_probe(&provider.model, probe)?;
104 Ok(provider)
105 }
106
107 pub async fn from_env_async() -> Result<Self> {
110 Self::from_config_async(None, None).await
111 }
112
113 pub async fn from_config_async(
115 config_model: Option<&str>,
116 config_host: Option<&str>,
117 ) -> Result<Self> {
118 let mut provider = Self::new_unprobed(config_model, config_host)?;
119 let probe = provider.request_async(&["dimension probe"]).await?;
120 provider.dimension = Self::dimension_from_probe(&provider.model, probe)?;
121 Ok(provider)
122 }
123
124 fn new_unprobed(config_model: Option<&str>, config_host: Option<&str>) -> Result<Self> {
126 let model = resolve_model(config_model);
127 let host = resolve_host(config_host);
128
129 let mut headers = HeaderMap::new();
130 headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
131 if let Ok(token) = std::env::var("OLLAMA_API_KEY") {
133 if !token.is_empty() {
134 let value = HeaderValue::from_str(&format!("Bearer {}", token)).map_err(|e| {
135 CtxError::embedding(format!("Invalid OLLAMA_API_KEY format: {}", e))
136 })?;
137 headers.insert(AUTHORIZATION, value);
138 }
139 }
140
141 let client = Client::builder()
142 .timeout(Duration::from_secs(REQUEST_TIMEOUT_SECS))
143 .connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
144 .default_headers(headers)
145 .build()
146 .map_err(|e| CtxError::embedding(format!("Failed to build HTTP client: {}", e)))?;
147
148 Ok(Self {
149 client,
150 host,
151 model,
152 dimension: 0,
153 })
154 }
155
156 fn dimension_from_probe(model: &str, probe: Vec<Embedding>) -> Result<usize> {
157 probe
158 .first()
159 .map(|e| e.vector.len())
160 .filter(|d| *d > 0)
161 .ok_or_else(|| {
162 CtxError::embedding(format!(
163 "Ollama model '{}' returned no embedding on probe",
164 model
165 ))
166 })
167 }
168
169 pub async fn embed_async(&self, text: &str) -> Result<Embedding> {
171 self.request_async(&[text])
172 .await?
173 .pop()
174 .ok_or_else(|| CtxError::embedding("Empty response"))
175 }
176
177 fn embed_url(&self) -> String {
179 format!("{}/api/embed", self.host.trim_end_matches('/'))
180 }
181
182 fn request(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
185 if tokio::runtime::Handle::try_current().is_ok() {
186 return Err(CtxError::embedding(
187 "Cannot call sync embed() from async context. Use request_async() instead.",
188 ));
189 }
190 get_or_create_runtime().block_on(self.request_async(texts))
191 }
192
193 pub async fn request_async(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
195 let body = OllamaEmbedRequest {
196 model: &self.model,
197 input: texts.to_vec(),
198 };
199
200 let mut last_error = None;
201 for attempt in 0..MAX_RETRIES {
202 match self.send_request(&body).await {
203 Ok(embeddings) => return Ok(embeddings),
204 Err(e) => {
205 let retryable = matches!(&e, CtxError::Embedding(msg)
208 if msg.contains("server error")
209 || msg.contains("timed out")
210 || msg.contains("Connection"));
211 if retryable && attempt < MAX_RETRIES - 1 {
212 let delay = RETRY_BASE_DELAY_MS * (1 << attempt);
213 tokio::time::sleep(Duration::from_millis(delay)).await;
214 last_error = Some(e);
215 continue;
216 }
217 return Err(e);
218 }
219 }
220 }
221 Err(last_error.unwrap_or_else(|| CtxError::embedding("Max retries exceeded")))
222 }
223
224 async fn send_request(&self, body: &OllamaEmbedRequest<'_>) -> Result<Vec<Embedding>> {
226 let response = self
227 .client
228 .post(self.embed_url())
229 .json(body)
230 .send()
231 .await
232 .map_err(|e| {
233 if e.is_timeout() {
234 CtxError::embedding(format!("Request timed out: {}", e))
235 } else if e.is_connect() {
236 CtxError::embedding(format!(
237 "Connection to Ollama at {} failed: {}. Is `ollama serve` running?",
238 self.host, e
239 ))
240 } else {
241 CtxError::embedding(e.to_string())
242 }
243 })?;
244
245 let status = response.status();
246 match status {
247 StatusCode::OK => {
248 let parsed: OllamaEmbedResponse = response
249 .json()
250 .await
251 .map_err(|e| CtxError::embedding(format!("Failed to parse response: {}", e)))?;
252 self.parse_response(parsed)
253 }
254 StatusCode::NOT_FOUND => {
255 Err(CtxError::ModelNotFound(format!(
257 "Ollama model '{}' not found. Pull it with: ollama pull {}",
258 self.model, self.model
259 )))
260 }
261 s if s.is_server_error() => {
262 let body = response.text().await.unwrap_or_default();
263 Err(CtxError::embedding(format!(
264 "server error ({}): {}",
265 status, body
266 )))
267 }
268 _ => {
269 let body = response.text().await.unwrap_or_default();
270 if let Ok(parsed) = serde_json::from_str::<OllamaEmbedResponse>(&body) {
272 if let Some(err) = parsed.error {
273 return Err(Self::classify_error(&self.model, err));
274 }
275 }
276 Err(CtxError::embedding(format!("HTTP {}: {}", status, body)))
277 }
278 }
279 }
280
281 fn classify_error(model: &str, message: String) -> CtxError {
283 if message.contains("not found") || message.contains("try pulling") {
284 CtxError::ModelNotFound(format!(
285 "Ollama model '{}' not found. Pull it with: ollama pull {}",
286 model, model
287 ))
288 } else {
289 CtxError::embedding(message)
290 }
291 }
292
293 fn parse_response(&self, response: OllamaEmbedResponse) -> Result<Vec<Embedding>> {
295 if let Some(err) = response.error {
296 return Err(Self::classify_error(&self.model, err));
297 }
298 let embeddings = response
299 .embeddings
300 .ok_or_else(|| CtxError::embedding("No embeddings in Ollama response"))?;
301 if embeddings.is_empty() {
302 return Err(CtxError::embedding(
303 "Ollama returned an empty embeddings list",
304 ));
305 }
306 if self.dimension != 0 {
308 for vector in &embeddings {
309 if vector.len() != self.dimension {
310 return Err(CtxError::DimensionMismatch {
311 expected: self.dimension,
312 actual: vector.len(),
313 });
314 }
315 }
316 }
317 Ok(embeddings.into_iter().map(Embedding::new).collect())
318 }
319}
320
321impl EmbeddingProvider for OllamaProvider {
322 fn name(&self) -> &str {
323 "ollama"
324 }
325
326 fn dimension(&self) -> usize {
327 self.dimension
328 }
329
330 fn embed(&self, text: &str) -> Result<Embedding> {
331 self.request(&[text])?
332 .pop()
333 .ok_or_else(|| CtxError::embedding("Empty response"))
334 }
335
336 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
337 const BATCH_SIZE: usize = 64;
340 let mut all = Vec::with_capacity(texts.len());
341 for chunk in texts.chunks(BATCH_SIZE) {
342 all.extend(self.request(chunk)?);
343 }
344 Ok(all)
345 }
346}
347
348#[cfg(test)]
349mod tests {
350 use super::*;
351
352 #[test]
353 fn host_precedence_env_over_config_over_default() {
354 std::env::set_var("OLLAMA_HOST", "localhost:11434");
355 assert_eq!(resolve_host(None), "http://localhost:11434"); assert_eq!(resolve_host(Some("http://cfg:1")), "http://localhost:11434"); std::env::remove_var("OLLAMA_HOST");
358 assert_eq!(resolve_host(Some("gpu-box:11434")), "http://gpu-box:11434"); assert_eq!(resolve_host(None), DEFAULT_HOST); }
361
362 #[test]
363 fn model_precedence_env_over_config_over_default() {
364 std::env::remove_var("OLLAMA_EMBED_MODEL");
365 assert_eq!(resolve_model(None), DEFAULT_MODEL);
366 assert_eq!(
367 resolve_model(Some("qwen3-embedding:8b")),
368 "qwen3-embedding:8b"
369 ); std::env::set_var("OLLAMA_EMBED_MODEL", "mxbai-embed-large");
371 assert_eq!(
372 resolve_model(Some("qwen3-embedding:8b")),
373 "mxbai-embed-large"
374 ); std::env::remove_var("OLLAMA_EMBED_MODEL");
376 }
377
378 fn offline_provider(dimension: usize) -> OllamaProvider {
381 OllamaProvider {
382 client: Client::new(),
383 host: DEFAULT_HOST.to_string(),
384 model: "test-model".to_string(),
385 dimension,
386 }
387 }
388
389 #[test]
390 fn parse_response_success() {
391 let provider = offline_provider(3);
392 let parsed = OllamaEmbedResponse {
393 embeddings: Some(vec![vec![0.1, 0.2, 0.3], vec![0.4, 0.5, 0.6]]),
394 error: None,
395 };
396 let out = provider.parse_response(parsed).unwrap();
397 assert_eq!(out.len(), 2);
398 assert_eq!(out[0].vector, vec![0.1, 0.2, 0.3]);
399 }
400
401 #[test]
402 fn parse_response_dimension_mismatch() {
403 let provider = offline_provider(3);
404 let parsed = OllamaEmbedResponse {
405 embeddings: Some(vec![vec![0.1, 0.2]]), error: None,
407 };
408 assert!(matches!(
409 provider.parse_response(parsed).unwrap_err(),
410 CtxError::DimensionMismatch {
411 expected: 3,
412 actual: 2
413 }
414 ));
415 }
416
417 #[test]
418 fn parse_response_model_not_found() {
419 let provider = offline_provider(0);
420 let parsed = OllamaEmbedResponse {
421 embeddings: None,
422 error: Some("model \"foo\" not found, try pulling it first".to_string()),
423 };
424 assert!(matches!(
425 provider.parse_response(parsed).unwrap_err(),
426 CtxError::ModelNotFound(_)
427 ));
428 }
429
430 #[test]
431 fn parse_response_empty() {
432 let provider = offline_provider(0);
433 let parsed = OllamaEmbedResponse {
434 embeddings: Some(vec![]),
435 error: None,
436 };
437 assert!(provider.parse_response(parsed).is_err());
438 }
439}