use crate::error::HostError;
pub trait Embedder: Send {
fn dim(&self) -> usize;
fn embed(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError>;
}
#[derive(Clone, Copy, Debug, Default)]
pub struct NullEmbedder;
impl Embedder for NullEmbedder {
fn dim(&self) -> usize {
0
}
fn embed(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
Ok(vec![Vec::new(); texts.len()])
}
}
#[derive(Clone)]
pub struct SharedEmbedder(std::sync::Arc<std::sync::Mutex<Box<dyn Embedder>>>);
impl SharedEmbedder {
pub fn new(inner: Box<dyn Embedder>) -> Self {
Self(std::sync::Arc::new(std::sync::Mutex::new(inner)))
}
fn inner(&self) -> std::sync::MutexGuard<'_, Box<dyn Embedder>> {
self.0.lock().unwrap_or_else(|e| e.into_inner())
}
}
impl std::fmt::Debug for SharedEmbedder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SharedEmbedder")
.field("dim", &self.dim())
.finish()
}
}
impl Embedder for SharedEmbedder {
fn dim(&self) -> usize {
self.inner().dim()
}
fn embed(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
self.inner().embed(texts)
}
}
#[derive(Debug)]
pub struct OpenAiCompatEmbedder {
url: String,
model: String,
api_key: Option<String>,
dim: usize,
agent: ureq::Agent,
}
impl OpenAiCompatEmbedder {
pub fn new(base_url: &str, model: &str, dim: usize) -> Self {
Self {
url: format!("{}/embeddings", base_url.trim_end_matches('/')),
model: model.to_string(),
api_key: None,
dim,
agent: ureq::Agent::new_with_defaults(),
}
}
pub fn with_api_key(mut self, key: impl Into<String>) -> Self {
self.api_key = Some(key.into());
self
}
}
impl Embedder for OpenAiCompatEmbedder {
fn dim(&self) -> usize {
self.dim
}
fn embed(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
if texts.is_empty() {
return Ok(Vec::new());
}
let body = serde_json::json!({ "model": self.model, "input": texts });
let mut request = self.agent.post(&self.url);
if let Some(key) = &self.api_key {
request = request.header("Authorization", &format!("Bearer {key}"));
}
let mut response = request
.send_json(&body)
.map_err(|e| HostError::Embed(format!("request to {}: {e}", self.url)))?;
let value: serde_json::Value = response
.body_mut()
.read_json()
.map_err(|e| HostError::Embed(format!("response body: {e}")))?;
let data = value
.get("data")
.and_then(|d| d.as_array())
.ok_or_else(|| HostError::Embed("response has no data array".into()))?;
if data.len() != texts.len() {
return Err(HostError::Embed(format!(
"expected {} embeddings, got {}",
texts.len(),
data.len()
)));
}
let mut out = vec![Vec::new(); texts.len()];
for item in data {
let index = item
.get("index")
.and_then(|i| i.as_u64())
.ok_or_else(|| HostError::Embed("embedding without an index".into()))?
as usize;
let raw = item
.get("embedding")
.and_then(|e| e.as_array())
.ok_or_else(|| HostError::Embed("embedding is not an array".into()))?;
if index >= out.len() || !out[index].is_empty() {
return Err(HostError::Embed(format!("bad embedding index {index}")));
}
if raw.len() != self.dim {
return Err(HostError::Embed(format!(
"dimension mismatch: server sent {}, configured {}",
raw.len(),
self.dim
)));
}
let mut v = Vec::with_capacity(raw.len());
for x in raw {
v.push(
x.as_f64().ok_or_else(|| {
HostError::Embed("embedding component is not a number".into())
})? as f32,
);
}
out[index] = v;
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Counting(usize);
impl Embedder for Counting {
fn dim(&self) -> usize {
3
}
fn embed(&mut self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
self.0 += texts.len();
Ok(vec![vec![self.0 as f32; 3]; texts.len()])
}
}
#[test]
fn clones_of_a_shared_embedder_reach_the_same_provider() {
let shared = SharedEmbedder::new(Box::new(Counting(0)));
let mut a = shared.clone();
let mut b = shared.clone();
assert_eq!(a.dim(), 3);
assert_eq!(format!("{shared:?}"), "SharedEmbedder { dim: 3 }");
assert_eq!(a.embed(&["x"]).unwrap(), vec![vec![1.0; 3]]);
assert_eq!(b.embed(&["y", "z"]).unwrap(), vec![vec![3.0; 3]; 2]);
}
#[test]
fn a_poisoned_provider_is_recovered_rather_than_propagated() {
let shared = SharedEmbedder::new(Box::new(Counting(0)));
let poisoner = shared.clone();
let _ = std::thread::spawn(move || {
let _guard = poisoner.inner();
panic!("provider blew up");
})
.join();
let mut after = shared.clone();
assert_eq!(after.embed(&["still works"]).unwrap().len(), 1);
}
#[test]
fn the_null_embedder_produces_one_empty_vector_per_text() {
let mut null = NullEmbedder;
assert_eq!(null.dim(), 0);
assert_eq!(null.embed(&["a", "b"]).unwrap(), vec![Vec::<f32>::new(); 2]);
}
}