#[cfg(feature = "ollama")]
use serde::Deserialize;
#[derive(Debug, thiserror::Error)]
pub enum EmbedError {
#[error("embedding backend error: {0}")]
Backend(String),
#[error("embedding backend returned an empty vector")]
Empty,
}
pub trait Embedder {
fn dimension(&self) -> usize;
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError>;
}
#[derive(Debug, Clone)]
pub struct HashEmbedder {
dimension: usize,
}
impl HashEmbedder {
#[must_use]
pub fn new(dimension: usize) -> Self {
Self { dimension }
}
}
impl Embedder for HashEmbedder {
fn dimension(&self) -> usize {
self.dimension
}
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
let mut vector = vec![0.0_f32; self.dimension];
if self.dimension == 0 {
return Ok(vector);
}
let modulus = self.dimension as u64;
for token in text.split_whitespace() {
let bucket = usize::try_from(crate::id::stable_id(token) % modulus).unwrap_or(0);
vector[bucket] += 1.0;
}
velesdb_core::simd_native::normalize_inplace_native(&mut vector);
Ok(vector)
}
}
pub type DynEmbedder = Box<dyn Embedder + Send + Sync>;
impl<T: Embedder + ?Sized> Embedder for Box<T> {
fn dimension(&self) -> usize {
(**self).dimension()
}
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
(**self).embed(text)
}
}
#[cfg(feature = "ollama")]
pub const DEFAULT_OLLAMA_URL: &str = "http://localhost:11434";
#[cfg(feature = "ollama")]
pub const DEFAULT_OLLAMA_MODEL: &str = "all-minilm";
#[cfg(feature = "ollama")]
#[derive(Debug, Clone)]
pub struct OllamaEmbedder {
base_url: String,
model: String,
dimension: usize,
}
#[cfg(feature = "ollama")]
impl OllamaEmbedder {
pub fn new(base_url: impl Into<String>, model: impl Into<String>) -> Result<Self, EmbedError> {
let base_url = base_url.into();
let model = model.into();
let dimension = request_embedding(&base_url, &model, "dimension probe")?.len();
if dimension == 0 {
return Err(EmbedError::Empty);
}
Ok(Self {
base_url,
model,
dimension,
})
}
}
#[cfg(feature = "ollama")]
impl Embedder for OllamaEmbedder {
fn dimension(&self) -> usize {
self.dimension
}
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
request_embedding(&self.base_url, &self.model, text)
}
}
#[cfg(feature = "ollama")]
fn build_request_body(model: &str, text: &str) -> String {
serde_json::json!({ "model": model, "prompt": text }).to_string()
}
#[cfg(feature = "ollama")]
#[derive(Deserialize)]
struct EmbeddingResponse {
embedding: Vec<f32>,
}
#[cfg(feature = "ollama")]
fn parse_embedding_response(body: &str) -> Result<Vec<f32>, EmbedError> {
let parsed: EmbeddingResponse = serde_json::from_str(body)
.map_err(|err| EmbedError::Backend(format!("invalid embeddings response: {err}")))?;
if parsed.embedding.is_empty() {
return Err(EmbedError::Empty);
}
Ok(parsed.embedding)
}
#[cfg(feature = "ollama")]
fn request_embedding(base_url: &str, model: &str, text: &str) -> Result<Vec<f32>, EmbedError> {
let url = format!("{base_url}/api/embeddings");
let body = build_request_body(model, text);
let response = ureq::post(&url)
.set("Content-Type", "application/json")
.send_string(&body)
.map_err(|err| EmbedError::Backend(format!("ollama request failed: {err}")))?;
let payload = response
.into_string()
.map_err(|err| EmbedError::Backend(format!("reading ollama response failed: {err}")))?;
parse_embedding_response(&payload)
}
#[cfg(all(test, feature = "ollama"))]
mod ollama_tests {
use super::*;
#[test]
fn request_body_carries_model_and_prompt() {
let body = build_request_body("all-minilm", "hello world");
let json: serde_json::Value = serde_json::from_str(&body).expect("valid json");
assert_eq!(json["model"], "all-minilm");
assert_eq!(json["prompt"], "hello world");
}
#[test]
fn parses_a_well_formed_embedding() {
let vector = parse_embedding_response(r#"{"embedding":[0.1,0.2,0.3]}"#).expect("parse");
assert_eq!(vector.len(), 3);
assert!((vector[0] - 0.1_f32).abs() < f32::EPSILON);
}
#[test]
fn rejects_an_empty_embedding() {
let parsed = parse_embedding_response(r#"{"embedding":[]}"#);
assert!(matches!(parsed, Err(EmbedError::Empty)));
}
#[test]
fn rejects_a_malformed_response() {
let parsed = parse_embedding_response(r#"{"oops":true}"#);
assert!(matches!(parsed, Err(EmbedError::Backend(_))));
}
#[test]
#[ignore = "requires a local Ollama with an embedding model (ollama pull all-minilm)"]
fn embeds_through_a_running_ollama() {
let embedder = OllamaEmbedder::new(DEFAULT_OLLAMA_URL, DEFAULT_OLLAMA_MODEL)
.expect("connect to ollama");
let vector = embedder
.embed("parking_lot avoids lock poisoning")
.expect("embed");
assert_eq!(vector.len(), embedder.dimension());
assert!(vector
.iter()
.any(|&component| component.abs() > f32::EPSILON));
}
}