use std::sync::{Arc, RwLock, RwLockReadGuard, RwLockWriteGuard};
use std::time::{Duration, Instant};
use crate::error::HostError;
pub const DEFAULT_EMBED_TIMEOUT: Duration = Duration::from_secs(10);
fn agent_with_timeout(timeout: Option<Duration>) -> ureq::Agent {
ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.timeout_global(timeout)
.build(),
)
}
pub trait Embedder: Send + Sync {
fn space_id(&self) -> &str;
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 space_id(&self) -> &str {
"none"
}
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 space_id(&self) -> &str {
self.0.space_id()
}
fn dim(&self) -> usize {
self.0.dim()
}
fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
self.0.embed(texts)
}
}
pub const DEFAULT_EMBED_RETRY_FIRST: Duration = Duration::from_secs(1);
pub const DEFAULT_EMBED_RETRY_MAX: Duration = Duration::from_secs(60);
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum EmbedErrorPolicy {
#[default]
Fail,
Degrade,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EmbedRetry {
Backoff {
first: Duration,
max: Duration,
},
Fixed(Duration),
Manual,
}
impl Default for EmbedRetry {
fn default() -> Self {
Self::Backoff {
first: DEFAULT_EMBED_RETRY_FIRST,
max: DEFAULT_EMBED_RETRY_MAX,
}
}
}
impl EmbedRetry {
fn wait(self, failures: u32) -> Option<Duration> {
match self {
Self::Manual => None,
Self::Fixed(after) => Some(after),
Self::Backoff { first, max } => {
let factor = 1u32
.checked_shl(failures.saturating_sub(1))
.unwrap_or(u32::MAX);
Some(first.saturating_mul(factor).min(max))
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EmbedderState {
Absent,
Active,
Suspended {
retry_at: Option<Instant>,
},
}
struct EmbedderSlot {
provider: Option<Arc<dyn Embedder>>,
suspended_until: Option<Option<Instant>>,
failures: u32,
}
impl EmbedderSlot {
fn new(provider: Option<Arc<dyn Embedder>>) -> Self {
Self {
provider,
suspended_until: None,
failures: 0,
}
}
fn usable(&mut self, now: Instant) -> Option<Arc<dyn Embedder>> {
match self.suspended_until {
None => self.provider.clone(),
Some(Some(deadline)) if deadline <= now => {
self.suspended_until = None;
self.provider.clone()
}
Some(_) => None,
}
}
fn note_failure(&mut self, retry: EmbedRetry, now: Instant) {
self.failures = self.failures.saturating_add(1);
if self.suspended_until.is_some() {
return;
}
self.suspended_until = Some(retry.wait(self.failures).map(|wait| now + wait));
}
fn state(&self) -> EmbedderState {
match (&self.provider, self.suspended_until) {
(None, _) => EmbedderState::Absent,
(Some(_), None) => EmbedderState::Active,
(Some(_), Some(retry_at)) => EmbedderState::Suspended { retry_at },
}
}
}
pub type Embedded = (Vec<f32>, String);
pub type EmbeddedBatch = (Vec<Vec<f32>>, String);
type Ready = (Arc<dyn Embedder>, String);
pub struct EmbedderGate {
slot: RwLock<EmbedderSlot>,
policy: EmbedErrorPolicy,
retry: EmbedRetry,
}
impl EmbedderGate {
pub fn new(
provider: Option<Arc<dyn Embedder>>,
policy: EmbedErrorPolicy,
retry: EmbedRetry,
) -> Self {
Self {
slot: RwLock::new(EmbedderSlot::new(provider)),
policy,
retry,
}
}
pub fn policy(&self) -> EmbedErrorPolicy {
self.policy
}
pub fn state(&self) -> EmbedderState {
let mut slot = self.write();
let _ = slot.usable(Instant::now());
slot.state()
}
pub fn suspend(&self) {
self.write().suspended_until = Some(None);
}
pub fn resume(&self) {
self.write().suspended_until = None;
}
pub fn provider(&self) -> Option<Arc<dyn Embedder>> {
self.read().provider.clone()
}
pub fn embed_one(
&self,
text: &str,
check_space: impl FnOnce(&str) -> Result<(), HostError>,
) -> Result<Option<Embedded>, HostError> {
let Some((embedder, space)) = self.ready(check_space)? else {
return Ok(None);
};
let mut vectors = match embedder.embed(&[text]) {
Ok(vectors) => vectors,
Err(error) => return self.degrade(error).map(|()| None),
};
if vectors.len() != 1 {
let got = vectors.len();
return self
.degrade(HostError::Embed(format!("expected 1 embedding, got {got}")))
.map(|()| None);
}
self.note_success();
Ok(Some((vectors.remove(0), space)))
}
pub fn embed_many(
&self,
texts: &[&str],
check_space: impl FnOnce(&str) -> Result<(), HostError>,
) -> Result<Option<EmbeddedBatch>, HostError> {
let Some((embedder, space)) = self.ready(check_space)? else {
return Ok(None);
};
if texts.is_empty() {
return Ok(Some((Vec::new(), space)));
}
let vectors = match embedder.embed(texts) {
Ok(vectors) => vectors,
Err(error) => return self.degrade(error).map(|()| None),
};
if vectors.len() != texts.len() {
let (want, got) = (texts.len(), vectors.len());
return self
.degrade(HostError::Embed(format!(
"expected {want} embeddings, got {got}"
)))
.map(|()| None);
}
self.note_success();
Ok(Some((vectors, space)))
}
pub(crate) fn install(&self, provider: Arc<dyn Embedder>) {
let mut slot = self.write();
slot.provider = Some(provider);
slot.suspended_until = None;
slot.failures = 0;
}
fn ready(
&self,
check_space: impl FnOnce(&str) -> Result<(), HostError>,
) -> Result<Option<Ready>, HostError> {
let Some(embedder) = self.usable() else {
return Ok(None);
};
if embedder.dim() == 0 {
return Ok(None);
}
let space = embedder.space_id().to_owned();
check_space(&space)?;
Ok(Some((embedder, space)))
}
fn usable(&self) -> Option<Arc<dyn Embedder>> {
{
let slot = self.read();
if slot.suspended_until.is_none() {
return slot.provider.clone();
}
}
self.write().usable(Instant::now())
}
fn degrade(&self, error: HostError) -> Result<(), HostError> {
if self.policy != EmbedErrorPolicy::Degrade || !matches!(error, HostError::Embed(_)) {
return Err(error);
}
self.write().note_failure(self.retry, Instant::now());
Ok(())
}
fn note_success(&self) {
if self.read().failures == 0 {
return;
}
self.write().failures = 0;
}
fn read(&self) -> RwLockReadGuard<'_, EmbedderSlot> {
self.slot.read().unwrap_or_else(|e| e.into_inner())
}
fn write(&self) -> RwLockWriteGuard<'_, EmbedderSlot> {
self.slot.write().unwrap_or_else(|e| e.into_inner())
}
}
#[derive(Debug)]
pub struct OpenAiCompatEmbedder {
url: String,
model: String,
space_id: String,
api_key: Option<String>,
dim: usize,
agent: ureq::Agent,
}
impl OpenAiCompatEmbedder {
pub fn new(endpoint_url: &str, model: &str, dim: usize) -> Self {
Self {
url: endpoint_url.to_string(),
model: model.to_string(),
space_id: model.to_string(),
api_key: None,
dim,
agent: agent_with_timeout(Some(DEFAULT_EMBED_TIMEOUT)),
}
}
pub fn with_timeout(mut self, timeout: Option<Duration>) -> Self {
self.agent = agent_with_timeout(timeout);
self
}
pub fn with_api_key(mut self, key: impl Into<String>) -> Self {
self.api_key = Some(key.into());
self
}
pub fn with_space_id(mut self, space_id: impl Into<String>) -> Self {
self.space_id = space_id.into();
self
}
}
impl Embedder for OpenAiCompatEmbedder {
fn space_id(&self) -> &str {
&self.space_id
}
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};
fn t0() -> Instant {
Instant::now()
}
#[test]
fn a_backoff_doubles_per_consecutive_failure_and_stops_at_its_ceiling() {
let retry = EmbedRetry::Backoff {
first: Duration::from_secs(1),
max: Duration::from_secs(8),
};
let waits: Vec<Duration> = (1..=6).map(|n| retry.wait(n).unwrap()).collect();
assert_eq!(
waits,
[
Duration::from_secs(1),
Duration::from_secs(2),
Duration::from_secs(4),
Duration::from_secs(8),
Duration::from_secs(8),
Duration::from_secs(8),
]
);
assert_eq!(retry.wait(64), Some(Duration::from_secs(8)));
assert_eq!(retry.wait(u32::MAX), Some(Duration::from_secs(8)));
}
#[test]
fn a_fixed_retry_ignores_the_failure_count_and_manual_never_retries() {
let fixed = EmbedRetry::Fixed(Duration::from_millis(250));
assert_eq!(fixed.wait(1), Some(Duration::from_millis(250)));
assert_eq!(fixed.wait(9), Some(Duration::from_millis(250)));
assert_eq!(EmbedRetry::Manual.wait(1), None);
}
#[test]
fn a_suspension_expires_on_the_next_call_after_its_deadline() {
let mut slot = EmbedderSlot::new(Some(Arc::new(NullEmbedder)));
let retry = EmbedRetry::Fixed(Duration::from_secs(30));
let now = t0();
slot.note_failure(retry, now);
assert!(slot.usable(now).is_none(), "still inside the interval");
assert!(matches!(
slot.state(),
EmbedderState::Suspended { retry_at: Some(_) }
));
assert!(slot.usable(now + Duration::from_secs(29)).is_none());
assert!(slot.usable(now + Duration::from_secs(30)).is_some());
assert_eq!(slot.state(), EmbedderState::Active);
}
#[test]
fn consecutive_failures_lengthen_the_wait_and_a_success_resets_it() {
let mut slot = EmbedderSlot::new(Some(Arc::new(NullEmbedder)));
let retry = EmbedRetry::Backoff {
first: Duration::from_secs(1),
max: Duration::from_secs(60),
};
let now = t0();
slot.note_failure(retry, now);
assert!(slot.usable(now + Duration::from_secs(1)).is_some());
slot.note_failure(retry, now + Duration::from_secs(1));
assert!(slot.usable(now + Duration::from_secs(2)).is_none());
assert!(slot.usable(now + Duration::from_secs(3)).is_some());
slot.failures = 0;
slot.note_failure(retry, now + Duration::from_secs(3));
assert!(slot.usable(now + Duration::from_secs(4)).is_some());
}
#[test]
fn an_explicit_suspension_has_no_deadline_and_survives_a_failure() {
let mut slot = EmbedderSlot::new(Some(Arc::new(NullEmbedder)));
let now = t0();
slot.suspended_until = Some(None);
slot.note_failure(EmbedRetry::Fixed(Duration::from_millis(1)), now);
assert!(slot.usable(now + Duration::from_secs(3600)).is_none());
assert_eq!(slot.state(), EmbedderState::Suspended { retry_at: None });
}
struct MiscountingEmbedder(usize);
impl Embedder for MiscountingEmbedder {
fn space_id(&self) -> &str {
"miscounting"
}
fn dim(&self) -> usize {
4
}
fn embed(&self, _texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
Ok(vec![vec![0.0; 4]; self.0])
}
}
#[test]
fn a_provider_answering_with_the_wrong_number_of_vectors_follows_the_policy() {
let strict = EmbedderGate::new(
Some(Arc::new(MiscountingEmbedder(2))),
EmbedErrorPolicy::Fail,
EmbedRetry::Manual,
);
assert!(matches!(
strict.embed_one("one text", |_| Ok(())),
Err(HostError::Embed(_))
));
assert!(matches!(
strict.embed_many(&["a", "b", "c"], |_| Ok(())),
Err(HostError::Embed(_))
));
assert_eq!(strict.state(), EmbedderState::Active);
let lenient = EmbedderGate::new(
Some(Arc::new(MiscountingEmbedder(2))),
EmbedErrorPolicy::Degrade,
EmbedRetry::Manual,
);
assert_eq!(lenient.embed_one("one text", |_| Ok(())).unwrap(), None);
assert_eq!(lenient.state(), EmbedderState::Suspended { retry_at: None });
assert_eq!(lenient.policy(), EmbedErrorPolicy::Degrade);
}
#[test]
fn an_empty_batch_answers_without_a_round_trip() {
let gate = EmbedderGate::new(
Some(Arc::new(Counting(AtomicUsize::new(0)))),
EmbedErrorPolicy::Fail,
EmbedRetry::Manual,
);
let (vectors, space) = gate.embed_many(&[], |_| Ok(())).unwrap().unwrap();
assert!(vectors.is_empty());
assert_eq!(space, "test/counting");
assert_eq!(gate.provider().unwrap().dim(), 3);
}
#[test]
fn a_refused_space_is_not_a_failure_the_policy_may_swallow() {
let gate = EmbedderGate::new(
Some(Arc::new(Counting(AtomicUsize::new(0)))),
EmbedErrorPolicy::Degrade,
EmbedRetry::Manual,
);
let refused = gate.embed_one("a text", |_| {
Err(HostError::Engine(plugmem_core::Error::UntrackedVectorSpace))
});
assert!(matches!(refused, Err(HostError::Engine(_))));
assert_eq!(gate.state(), EmbedderState::Active);
}
#[test]
fn a_slot_with_no_provider_is_absent_whatever_is_done_to_it() {
let mut slot = EmbedderSlot::new(None);
assert_eq!(slot.state(), EmbedderState::Absent);
slot.suspended_until = Some(None);
assert_eq!(slot.state(), EmbedderState::Absent);
assert!(slot.usable(t0()).is_none());
}
struct Counting(AtomicUsize);
impl Embedder for Counting {
fn space_id(&self) -> &str {
"test/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 space_id(&self) -> &str {
"test/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]);
}
}