use crate::error::HostError;
pub trait Embedder: Send + Sync {
fn dim(&self) -> usize;
fn embed(&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(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
Ok(vec![Vec::new(); texts.len()])
}
}
#[derive(Clone)]
pub struct SharedEmbedder(std::sync::Arc<dyn Embedder>);
impl SharedEmbedder {
pub fn new(inner: Box<dyn Embedder>) -> Self {
Self(std::sync::Arc::from(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.0.dim())
.finish()
}
}
impl Embedder for SharedEmbedder {
fn dim(&self) -> usize {
self.0.dim()
}
fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
self.0.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(&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::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct Counting(AtomicUsize);
impl Embedder for Counting {
fn dim(&self) -> usize {
3
}
fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
let total = self.0.fetch_add(texts.len(), Ordering::Relaxed) + texts.len();
Ok(vec![vec![total as f32; 3]; texts.len()])
}
}
struct Overlapping {
inside: AtomicUsize,
peak: AtomicUsize,
}
impl Embedder for Overlapping {
fn dim(&self) -> usize {
1
}
fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
let now = self.inside.fetch_add(1, Ordering::SeqCst) + 1;
self.peak.fetch_max(now, Ordering::SeqCst);
std::thread::sleep(std::time::Duration::from_millis(50));
self.inside.fetch_sub(1, Ordering::SeqCst);
Ok(vec![vec![0.0]; texts.len()])
}
}
#[test]
fn clones_of_a_shared_embedder_reach_the_same_provider() {
let shared = SharedEmbedder::new(Box::new(Counting(AtomicUsize::new(0))));
let a = shared.clone();
let 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 concurrent_callers_are_inside_the_provider_at_the_same_time() {
let provider = std::sync::Arc::new(Overlapping {
inside: AtomicUsize::new(0),
peak: AtomicUsize::new(0),
});
let shared = SharedEmbedder(provider.clone());
std::thread::scope(|scope| {
for _ in 0..4 {
let handle = shared.clone();
scope.spawn(move || handle.embed(&["question"]).unwrap());
}
});
assert!(
provider.peak.load(Ordering::SeqCst) > 1,
"callers serialized: peak concurrency was {}",
provider.peak.load(Ordering::SeqCst)
);
}
#[test]
fn the_null_embedder_produces_one_empty_vector_per_text() {
let null = NullEmbedder;
assert_eq!(null.dim(), 0);
assert_eq!(null.embed(&["a", "b"]).unwrap(), vec![Vec::<f32>::new(); 2]);
}
}