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(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)
}
}