Skip to main content

khive_runtime/
embedder_registry.rs

1//! EmbedderRegistry — pack-extensible embedding provider surface.
2//!
3//! Packs implement [`EmbedderProvider`] and register custom models via
4//! [`crate::KhiveRuntime::register_embedder`]. Built-in lattice models are pre-registered
5//! during runtime construction and require no opt-in.
6
7use std::collections::HashMap;
8use std::sync::atomic::{AtomicUsize, Ordering};
9use std::sync::mpsc::{self, Receiver, SyncSender, TrySendError};
10use std::sync::{Arc, OnceLock};
11
12use async_trait::async_trait;
13use lattice_embed::{
14    CachedEmbeddingService, EmbeddingModel, EmbeddingService, NativeEmbeddingService,
15    DEFAULT_MAX_BATCH_SIZE, MAX_TEXT_BYTES,
16};
17use tokio::sync::OnceCell;
18
19use crate::error::{RuntimeError, RuntimeResult};
20
21#[derive(Clone, Copy)]
22enum EmbeddingCall {
23    Generic,
24    Query,
25    Passage,
26}
27
28const EMBEDDING_QUEUE_CAPACITY: usize = 32;
29const EMBEDDING_MAX_JOB_BYTES: usize = DEFAULT_MAX_BATCH_SIZE * MAX_TEXT_BYTES;
30// 32 queue slots × the normal 128-text batch × 32 KiB/text = 128 MiB in flight.
31const EMBEDDING_QUEUE_BYTE_BUDGET: usize = EMBEDDING_QUEUE_CAPACITY * 128 * MAX_TEXT_BYTES;
32
33struct InFlightBytes {
34    counter: Arc<AtomicUsize>,
35    bytes: usize,
36}
37
38impl InFlightBytes {
39    fn reserve(
40        counter: Arc<AtomicUsize>,
41        byte_budget: usize,
42        bytes: usize,
43    ) -> lattice_embed::Result<Self> {
44        let mut current = counter.load(Ordering::Acquire);
45        loop {
46            let Some(next) = current.checked_add(bytes) else {
47                return Err(lattice_embed::EmbedError::Internal(format!(
48                    "embedding worker byte budget exceeded: in-flight byte count overflowed the {byte_budget}-byte budget"
49                )));
50            };
51            if next > byte_budget {
52                return Err(lattice_embed::EmbedError::Internal(format!(
53                    "embedding worker byte budget exceeded: {current} in flight + {bytes} job bytes > {byte_budget}"
54                )));
55            }
56            match counter.compare_exchange_weak(current, next, Ordering::AcqRel, Ordering::Acquire)
57            {
58                Ok(_) => return Ok(Self { counter, bytes }),
59                Err(observed) => current = observed,
60            }
61        }
62    }
63}
64
65impl Drop for InFlightBytes {
66    fn drop(&mut self) {
67        let previous = self.counter.fetch_sub(self.bytes, Ordering::AcqRel);
68        debug_assert!(previous >= self.bytes, "embedding byte counter underflow");
69    }
70}
71
72struct EmbeddingJob {
73    texts: Vec<String>,
74    model: EmbeddingModel,
75    call: EmbeddingCall,
76    reply: tokio::sync::oneshot::Sender<lattice_embed::Result<Vec<Vec<f32>>>>,
77    _in_flight: InFlightBytes,
78}
79
80/// Bounds non-cancellable native inference to one worker and a fixed queue.
81/// Callers can detach safely because closed queued jobs are skipped before inference.
82pub(crate) struct BlockingEmbeddingService<S> {
83    inner: Arc<S>,
84    worker: OnceLock<Result<SyncSender<EmbeddingJob>, String>>,
85    in_flight_bytes: Arc<AtomicUsize>,
86    byte_budget: usize,
87}
88
89impl<S> BlockingEmbeddingService<S> {
90    pub(crate) fn new(inner: Arc<S>) -> Self {
91        Self {
92            inner,
93            worker: OnceLock::new(),
94            in_flight_bytes: Arc::new(AtomicUsize::new(0)),
95            byte_budget: EMBEDDING_QUEUE_BYTE_BUDGET,
96        }
97    }
98
99    #[cfg(test)]
100    fn with_byte_budget(inner: Arc<S>, byte_budget: usize) -> Self {
101        Self {
102            inner,
103            worker: OnceLock::new(),
104            in_flight_bytes: Arc::new(AtomicUsize::new(0)),
105            byte_budget,
106        }
107    }
108}
109
110impl<S: EmbeddingService + 'static> BlockingEmbeddingService<S> {
111    fn input_bytes(texts: &[String]) -> lattice_embed::Result<usize> {
112        if texts.is_empty() {
113            return Err(lattice_embed::EmbedError::InvalidInput(
114                "no texts provided".to_owned(),
115            ));
116        }
117        let input_bytes = texts.iter().try_fold(0usize, |total, text| {
118            total.checked_add(text.len()).ok_or_else(|| {
119                lattice_embed::EmbedError::InvalidInput(format!(
120                    "embedding job input exceeds the {EMBEDDING_MAX_JOB_BYTES}-byte maximum"
121                ))
122            })
123        })?;
124        if input_bytes > EMBEDDING_MAX_JOB_BYTES {
125            return Err(lattice_embed::EmbedError::InvalidInput(format!(
126                "embedding job input is {input_bytes} bytes; maximum is {EMBEDDING_MAX_JOB_BYTES} bytes"
127            )));
128        }
129        if texts.len() > DEFAULT_MAX_BATCH_SIZE {
130            return Err(lattice_embed::EmbedError::InvalidInput(format!(
131                "batch size {} exceeds maximum {DEFAULT_MAX_BATCH_SIZE}",
132                texts.len()
133            )));
134        }
135        if let Some(text) = texts.iter().find(|text| text.len() > MAX_TEXT_BYTES) {
136            return Err(lattice_embed::EmbedError::TextTooLong {
137                length: text.len(),
138                max: MAX_TEXT_BYTES,
139            });
140        }
141        Ok(input_bytes)
142    }
143
144    fn worker(&self) -> lattice_embed::Result<&SyncSender<EmbeddingJob>> {
145        self.worker
146            .get_or_init(|| {
147                let (sender, receiver) = mpsc::sync_channel(EMBEDDING_QUEUE_CAPACITY);
148                let inner = Arc::clone(&self.inner);
149                let runtime = tokio::runtime::Handle::current();
150                std::thread::Builder::new()
151                    .name("khive-embedding".to_owned())
152                    .spawn(move || Self::run_worker(inner, runtime, receiver))
153                    .map(|_| sender)
154                    .map_err(|error| error.to_string())
155            })
156            .as_ref()
157            .map_err(|error| lattice_embed::EmbedError::Internal(error.clone()))
158    }
159
160    fn run_worker(
161        inner: Arc<S>,
162        runtime: tokio::runtime::Handle,
163        receiver: Receiver<EmbeddingJob>,
164    ) {
165        while let Ok(job) = receiver.recv() {
166            if job.reply.is_closed() {
167                continue;
168            }
169            let result = runtime.block_on(async {
170                match job.call {
171                    EmbeddingCall::Generic => inner.embed(&job.texts, job.model).await,
172                    EmbeddingCall::Query => inner.embed_query(&job.texts, job.model).await,
173                    EmbeddingCall::Passage => inner.embed_passage(&job.texts, job.model).await,
174                }
175            });
176            let _ = job.reply.send(result);
177        }
178    }
179
180    async fn run(
181        &self,
182        texts: &[String],
183        model: EmbeddingModel,
184        call: EmbeddingCall,
185    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
186        let input_bytes = Self::input_bytes(texts)?;
187        let sender = self.worker()?;
188        let in_flight = InFlightBytes::reserve(
189            Arc::clone(&self.in_flight_bytes),
190            self.byte_budget,
191            input_bytes,
192        )?;
193        let (reply, receiver) = tokio::sync::oneshot::channel();
194        let job = EmbeddingJob {
195            texts: texts.to_vec(),
196            model,
197            call,
198            reply,
199            _in_flight: in_flight,
200        };
201        sender.try_send(job).map_err(|error| match error {
202            TrySendError::Full(_) => {
203                lattice_embed::EmbedError::Internal("embedding worker queue is full".to_owned())
204            }
205            TrySendError::Disconnected(_) => lattice_embed::EmbedError::Internal(
206                "embedding worker channel is disconnected".to_owned(),
207            ),
208        })?;
209        receiver
210            .await
211            .map_err(|error| lattice_embed::EmbedError::Internal(error.to_string()))?
212    }
213}
214
215#[async_trait]
216impl<S: EmbeddingService + 'static> EmbeddingService for BlockingEmbeddingService<S> {
217    async fn embed(
218        &self,
219        texts: &[String],
220        model: EmbeddingModel,
221    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
222        self.run(texts, model, EmbeddingCall::Generic).await
223    }
224
225    async fn embed_query(
226        &self,
227        texts: &[String],
228        model: EmbeddingModel,
229    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
230        self.run(texts, model, EmbeddingCall::Query).await
231    }
232
233    async fn embed_passage(
234        &self,
235        texts: &[String],
236        model: EmbeddingModel,
237    ) -> lattice_embed::Result<Vec<Vec<f32>>> {
238        self.run(texts, model, EmbeddingCall::Passage).await
239    }
240
241    fn model_config(&self, model: EmbeddingModel) -> lattice_embed::ModelConfig {
242        self.inner.model_config(model)
243    }
244
245    fn supports_model(&self, model: EmbeddingModel) -> bool {
246        self.inner.supports_model(model)
247    }
248
249    fn name(&self) -> &'static str {
250        self.inner.name()
251    }
252}
253
254/// A source that can produce an [`EmbeddingService`] by name.
255///
256/// Packs implement this trait to register custom embedding backends.
257/// The runtime calls [`build`](EmbedderProvider::build) lazily — once per
258/// process per model — and caches the result. Subsequent calls to
259/// `KhiveRuntime::embedder(name)` are cheap.
260///
261/// Built-in lattice models are registered automatically via
262/// [`LatticeEmbedderProvider`]; packs need not re-register them.
263#[async_trait]
264pub trait EmbedderProvider: Send + Sync {
265    /// Stable, case-sensitive name for this embedder.
266    ///
267    /// Must be unique across all registered providers. The name is used as
268    /// the key in `KhiveRuntime::embedder(name)` lookups and as the storage
269    /// table suffix for vector indices. Use the model's canonical short form
270    /// (e.g. `"all-minilm-l6-v2"`, `"my-custom-encoder"`).
271    fn name(&self) -> &str;
272
273    /// Output vector dimension for this embedder.
274    ///
275    /// Must be consistent with what [`build`](Self::build) produces.
276    /// The runtime uses this to pre-register the vector store columns.
277    fn dimensions(&self) -> usize;
278
279    /// Construct the underlying [`EmbeddingService`].
280    ///
281    /// Called at most once per process. The result is cached in a
282    /// [`OnceCell`]; concurrent callers block on the first call and share
283    /// the result thereafter.
284    async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>>;
285}
286
287/// An entry in the [`EmbedderRegistry`] combining a provider with its
288/// lazy-initialized service.
289pub(crate) struct EmbedderEntry {
290    provider: Arc<dyn EmbedderProvider>,
291    cell: Arc<OnceCell<Arc<dyn EmbeddingService>>>,
292}
293
294impl Clone for EmbedderEntry {
295    fn clone(&self) -> Self {
296        Self {
297            provider: Arc::clone(&self.provider),
298            cell: Arc::clone(&self.cell),
299        }
300    }
301}
302
303/// Registry of named [`EmbedderProvider`] instances.
304///
305/// Built during `KhiveRuntime` construction and optionally extended by packs
306/// via [`crate::KhiveRuntime::register_embedder`]. The registry is internally
307/// reference-counted so `KhiveRuntime::clone()` shares the same providers
308/// and cached service instances.
309#[derive(Clone, Default)]
310pub struct EmbedderRegistry {
311    entries: HashMap<String, EmbedderEntry>,
312}
313
314impl EmbedderRegistry {
315    /// Create an empty registry.
316    pub fn new() -> Self {
317        Self {
318            entries: HashMap::new(),
319        }
320    }
321
322    /// Register a provider.
323    ///
324    /// If a provider with the same [`name`](EmbedderProvider::name) already
325    /// exists, it is replaced (last-writer wins) and any cached service is
326    /// discarded, since pack registration order is not guaranteed and packs
327    /// may legitimately override a default model under the same name.
328    /// Callers needing strict collision detection should check
329    /// [`names`](Self::names) before registering.
330    pub fn register<P: EmbedderProvider + 'static>(&mut self, provider: P) {
331        let name = provider.name().to_owned();
332        self.entries.insert(
333            name,
334            EmbedderEntry {
335                provider: Arc::new(provider),
336                cell: Arc::new(OnceCell::new()),
337            },
338        );
339    }
340
341    /// Look up a provider by name.
342    pub fn get_provider(&self, name: &str) -> Option<&dyn EmbedderProvider> {
343        self.entries.get(name).map(|e| e.provider.as_ref())
344    }
345
346    /// Returns `true` if a provider with this name is registered.
347    pub fn contains(&self, name: &str) -> bool {
348        self.entries.contains_key(name)
349    }
350
351    /// Names of all registered providers, in unspecified order.
352    pub fn names(&self) -> Vec<String> {
353        self.entries.keys().cloned().collect()
354    }
355
356    /// Return a cloned entry for `name` without holding any lock.
357    ///
358    /// The caller can then call [`EmbedderEntry::resolve`] without holding
359    /// a lock — this avoids holding a `RwLockGuard` across `await` points.
360    /// Returns `None` if `name` is not registered.
361    pub(crate) fn get_entry(&self, name: &str) -> Option<EmbedderEntry> {
362        self.entries.get(name).cloned()
363    }
364
365    /// Lazily resolve a registered provider to its live [`EmbeddingService`].
366    ///
367    /// Returns [`RuntimeError::UnknownModel`] if `name` is not registered.
368    /// The first call for a given name triggers [`EmbedderProvider::build`];
369    /// subsequent calls return the cached `Arc`.
370    ///
371    /// Prefer [`crate::KhiveRuntime::embedder`] over calling this directly from pack
372    /// handlers — the runtime method handles alias resolution and error mapping.
373    pub async fn get_service(&self, name: &str) -> RuntimeResult<Arc<dyn EmbeddingService>> {
374        let entry = self
375            .entries
376            .get(name)
377            .ok_or_else(|| RuntimeError::UnknownModel(name.to_string()))?
378            .clone();
379
380        Ok(entry.resolve().await?.0)
381    }
382}
383
384impl EmbedderEntry {
385    /// Lazily initialise and return the embedding service for this entry.
386    ///
387    /// `OnceCell::get_or_try_init` single-flights concurrent cold callers: only
388    /// one of them runs [`EmbedderProvider::build`], and every other waiter
389    /// receives the same result once it completes. Only the caller whose task
390    /// actually ran `build()` gets back `Some(duration)`; every other caller
391    /// (cache hit or wait-for-in-flight-build) gets `None`.
392    ///
393    /// Returns `RuntimeError` if `build()` fails, rather than panicking. On
394    /// failure the cell stays unset, so a later call retries the build.
395    pub(crate) async fn resolve(self) -> RuntimeResult<(Arc<dyn EmbeddingService>, Option<i64>)> {
396        let mut own_init_duration_us: Option<i64> = None;
397        let provider = Arc::clone(&self.provider);
398        let init_duration_us = &mut own_init_duration_us;
399        let svc = self
400            .cell
401            .get_or_try_init(|| async move {
402                let init_start = std::time::Instant::now();
403                let svc = provider.build().await.map_err(|e| {
404                    crate::error::RuntimeError::Internal(format!(
405                        "EmbedderProvider '{}' build() failed: {e}",
406                        provider.name()
407                    ))
408                })?;
409                *init_duration_us = Some(init_start.elapsed().as_micros() as i64);
410                Ok::<_, RuntimeError>(svc)
411            })
412            .await?;
413        Ok((Arc::clone(svc), own_init_duration_us))
414    }
415}
416
417// ── LatticeEmbedderProvider ───────────────────────────────────────────────────
418
419/// Adapter that wraps a [`lattice_embed::EmbeddingModel`] as an
420/// [`EmbedderProvider`].
421///
422/// All built-in models (MiniLM, paraphrase-multilingual, BGE variants, etc.)
423/// are registered as `LatticeEmbedderProvider` instances during
424/// `KhiveRuntime` construction. External callers do not need to use this type
425/// unless they are constructing a custom registry from scratch.
426pub struct LatticeEmbedderProvider {
427    model: EmbeddingModel,
428    /// Cached `to_string()` result so `name()` can return `&str`.
429    name: String,
430}
431
432impl LatticeEmbedderProvider {
433    /// Create a new provider wrapping the given lattice model.
434    pub fn new(model: EmbeddingModel) -> Self {
435        let name = model.to_string();
436        Self { model, name }
437    }
438}
439
440#[async_trait]
441impl EmbedderProvider for LatticeEmbedderProvider {
442    fn name(&self) -> &str {
443        &self.name
444    }
445
446    fn dimensions(&self) -> usize {
447        self.model.dimensions()
448    }
449
450    async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
451        let native = Arc::new(NativeEmbeddingService::with_model(self.model));
452        native.ensure_loaded().await?;
453        let cached = Arc::new(CachedEmbeddingService::with_default_cache(native));
454        Ok(Arc::new(BlockingEmbeddingService::new(cached)) as Arc<dyn EmbeddingService>)
455    }
456}
457
458// ── Unit tests ────────────────────────────────────────────────────────────────
459
460#[cfg(test)]
461mod tests {
462    use super::*;
463    use std::collections::HashSet;
464    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
465    use std::sync::{Condvar, Mutex};
466    use std::time::Duration;
467
468    struct ConstVecProvider {
469        name: String,
470        dims: usize,
471        build_calls: Arc<AtomicUsize>,
472    }
473
474    impl ConstVecProvider {
475        fn new(name: &str, dims: usize) -> Self {
476            Self {
477                name: name.to_owned(),
478                dims,
479                build_calls: Arc::new(AtomicUsize::new(0)),
480            }
481        }
482    }
483
484    /// A trivial embedding service that returns a constant vector of `1.0`s.
485    /// The `model` parameter is ignored — this service always returns the
486    /// same synthetic vector regardless of which model is requested.
487    struct ConstVecService {
488        dims: usize,
489    }
490
491    #[async_trait]
492    impl EmbeddingService for ConstVecService {
493        async fn embed(
494            &self,
495            texts: &[String],
496            _model: EmbeddingModel,
497        ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
498            Ok(texts.iter().map(|_| vec![1.0_f32; self.dims]).collect())
499        }
500
501        fn supports_model(&self, _model: EmbeddingModel) -> bool {
502            true
503        }
504
505        fn name(&self) -> &'static str {
506            "const-vec-service"
507        }
508    }
509
510    #[async_trait]
511    impl EmbedderProvider for ConstVecProvider {
512        fn name(&self) -> &str {
513            &self.name
514        }
515
516        fn dimensions(&self) -> usize {
517            self.dims
518        }
519
520        async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
521            self.build_calls.fetch_add(1, Ordering::SeqCst);
522            Ok(Arc::new(ConstVecService { dims: self.dims }))
523        }
524    }
525
526    struct FirstLoadBlockingService {
527        loaded: AtomicBool,
528    }
529
530    #[async_trait]
531    impl EmbeddingService for FirstLoadBlockingService {
532        async fn embed(
533            &self,
534            texts: &[String],
535            _model: EmbeddingModel,
536        ) -> lattice_embed::Result<Vec<Vec<f32>>> {
537            if !self.loaded.swap(true, Ordering::SeqCst) {
538                tokio::task::spawn_blocking(|| {})
539                    .await
540                    .map_err(|error| lattice_embed::EmbedError::Internal(error.to_string()))?;
541            }
542            Ok(texts.iter().map(|_| vec![1.0]).collect())
543        }
544
545        fn supports_model(&self, _model: EmbeddingModel) -> bool {
546            true
547        }
548
549        fn name(&self) -> &'static str {
550            "first-load-blocking-service"
551        }
552    }
553
554    struct BlockingTestService {
555        calls: Mutex<Vec<String>>,
556        entered: AtomicUsize,
557        release: (Mutex<bool>, Condvar),
558        thread_ids: Mutex<HashSet<std::thread::ThreadId>>,
559    }
560
561    impl BlockingTestService {
562        fn new() -> Self {
563            Self {
564                calls: Mutex::new(Vec::new()),
565                entered: AtomicUsize::new(0),
566                release: (Mutex::new(false), Condvar::new()),
567                thread_ids: Mutex::new(HashSet::new()),
568            }
569        }
570
571        fn release(&self) {
572            *self
573                .release
574                .0
575                .lock()
576                .expect("release lock must not be poisoned") = true;
577            self.release.1.notify_all();
578        }
579    }
580
581    #[async_trait]
582    impl EmbeddingService for BlockingTestService {
583        async fn embed(
584            &self,
585            texts: &[String],
586            _model: EmbeddingModel,
587        ) -> lattice_embed::Result<Vec<Vec<f32>>> {
588            let text = texts.first().cloned().unwrap_or_default();
589            self.thread_ids
590                .lock()
591                .expect("thread id lock must not be poisoned")
592                .insert(std::thread::current().id());
593            self.calls
594                .lock()
595                .expect("call lock must not be poisoned")
596                .push(text.clone());
597            self.entered.fetch_add(1, Ordering::Release);
598
599            if text != "later" {
600                let (released, wake) = &self.release;
601                let guard = released.lock().expect("release lock must not be poisoned");
602                let _guard = wake
603                    .wait_while(guard, |released| !*released)
604                    .expect("release lock must not be poisoned");
605            }
606
607            Ok(texts.iter().map(|_| vec![1.0]).collect())
608        }
609
610        fn supports_model(&self, _model: EmbeddingModel) -> bool {
611            true
612        }
613
614        fn name(&self) -> &'static str {
615            "blocking-test-service"
616        }
617    }
618
619    #[test]
620    fn blocking_adapter_first_use_completes_with_single_blocking_thread() {
621        let runtime = tokio::runtime::Builder::new_current_thread()
622            .enable_time()
623            .max_blocking_threads(1)
624            .build()
625            .expect("current-thread runtime must build");
626        let service = BlockingEmbeddingService::new(Arc::new(FirstLoadBlockingService {
627            loaded: AtomicBool::new(false),
628        }));
629
630        let result = runtime.block_on(async {
631            tokio::time::timeout(
632                Duration::from_secs(5),
633                service.embed(&["first use".to_owned()], EmbeddingModel::default()),
634            )
635            .await
636        });
637        runtime.shutdown_timeout(Duration::from_secs(1));
638
639        let embeddings = result
640            .expect("first-use embedding must not exhaust the blocking pool")
641            .expect("first-use embedding must succeed");
642        assert_eq!(embeddings, vec![vec![1.0]]);
643    }
644
645    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
646    async fn blocking_adapter_uses_one_worker_for_concurrent_calls() {
647        const CALL_COUNT: usize = 32;
648        let inner = Arc::new(BlockingTestService::new());
649        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
650        let mut calls = Vec::with_capacity(CALL_COUNT);
651
652        for index in 0..CALL_COUNT {
653            let service = Arc::clone(&service);
654            calls.push(tokio::spawn(async move {
655                service
656                    .embed(&[format!("request-{index}")], EmbeddingModel::default())
657                    .await
658            }));
659        }
660
661        let _ = tokio::time::timeout(Duration::from_millis(250), async {
662            while inner.entered.load(Ordering::Acquire) < 2 {
663                tokio::task::yield_now().await;
664            }
665        })
666        .await;
667        inner.release();
668
669        for call in calls {
670            call.await
671                .expect("embedding task must not panic")
672                .expect("embedding call must succeed");
673        }
674        assert_eq!(
675            inner
676                .thread_ids
677                .lock()
678                .expect("thread id lock must not be poisoned")
679                .len(),
680            1,
681            "concurrent calls must share one native worker thread"
682        );
683    }
684
685    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
686    async fn blocking_adapter_skips_timed_out_call_and_serves_later_call() {
687        let inner = Arc::new(BlockingTestService::new());
688        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
689
690        let first_service = Arc::clone(&service);
691        let first = tokio::spawn(async move {
692            first_service
693                .embed(&["first".to_owned()], EmbeddingModel::default())
694                .await
695        });
696        tokio::time::timeout(Duration::from_secs(1), async {
697            while inner.entered.load(Ordering::Acquire) == 0 {
698                tokio::task::yield_now().await;
699            }
700        })
701        .await
702        .expect("first embedding call must enter the native service");
703
704        let abandoned = tokio::time::timeout(
705            Duration::from_millis(50),
706            service.embed(&["abandoned".to_owned()], EmbeddingModel::default()),
707        )
708        .await;
709        let later_service = Arc::clone(&service);
710        let later = tokio::spawn(async move {
711            later_service
712                .embed(&["later".to_owned()], EmbeddingModel::default())
713                .await
714        });
715        inner.release();
716
717        first
718            .await
719            .expect("first embedding task must not panic")
720            .expect("first embedding call must succeed");
721        let later_result = tokio::time::timeout(Duration::from_secs(1), later)
722            .await
723            .expect("later embedding call must be served")
724            .expect("later embedding task must not panic")
725            .expect("later embedding call must succeed");
726
727        assert!(abandoned.is_err(), "queued embedding call must time out");
728        assert_eq!(later_result, vec![vec![1.0]]);
729        assert_eq!(
730            *inner.calls.lock().expect("call lock must not be poisoned"),
731            vec!["first".to_owned(), "later".to_owned()],
732            "the worker must skip a queued call whose receiver is closed"
733        );
734    }
735
736    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
737    async fn blocking_adapter_rejects_call_when_queue_is_full() {
738        let inner = Arc::new(BlockingTestService::new());
739        let service = Arc::new(BlockingEmbeddingService::new(Arc::clone(&inner)));
740        let first_service = Arc::clone(&service);
741        let first = tokio::spawn(async move {
742            first_service
743                .embed(&["first".to_owned()], EmbeddingModel::default())
744                .await
745        });
746        tokio::time::timeout(Duration::from_secs(1), async {
747            while inner.entered.load(Ordering::Acquire) == 0 {
748                tokio::task::yield_now().await;
749            }
750        })
751        .await
752        .expect("first embedding call must occupy the native worker");
753
754        let sender = service.worker().expect("worker must be running");
755        let mut queued_receivers = Vec::with_capacity(EMBEDDING_QUEUE_CAPACITY);
756        for index in 0..EMBEDDING_QUEUE_CAPACITY {
757            let (reply, receiver) = tokio::sync::oneshot::channel();
758            let queued = sender.try_send(EmbeddingJob {
759                texts: vec![format!("queued-{index}")],
760                model: EmbeddingModel::default(),
761                call: EmbeddingCall::Generic,
762                reply,
763                _in_flight: InFlightBytes::reserve(
764                    Arc::clone(&service.in_flight_bytes),
765                    service.byte_budget,
766                    0,
767                )
768                .expect("zero-byte test job must fit the byte budget"),
769            });
770            assert!(queued.is_ok(), "bounded queue must accept its capacity");
771            queued_receivers.push(receiver);
772        }
773
774        let overflow = tokio::time::timeout(
775            Duration::from_millis(100),
776            service.embed(&["overflow".to_owned()], EmbeddingModel::default()),
777        )
778        .await
779        .expect("a full embedding queue must fail without waiting")
780        .expect_err("a full embedding queue must return an embedding error");
781
782        drop(queued_receivers);
783        inner.release();
784        first
785            .await
786            .expect("first embedding task must not panic")
787            .expect("first embedding call must succeed");
788        assert!(
789            overflow.to_string().contains("queue is full"),
790            "queue saturation must use the embedding failure path: {overflow}"
791        );
792    }
793
794    #[tokio::test]
795    async fn blocking_adapter_rejects_oversized_job_before_enqueue() {
796        let service = BlockingEmbeddingService::new(Arc::new(ConstVecService { dims: 1 }));
797        let oversized =
798            "x".repeat(lattice_embed::DEFAULT_MAX_BATCH_SIZE * lattice_embed::MAX_TEXT_BYTES + 1);
799
800        let error = service
801            .embed(&[oversized], EmbeddingModel::default())
802            .await
803            .expect_err("an oversized embedding job must be rejected");
804
805        assert!(
806            error.to_string().contains("embedding job input"),
807            "oversized admission must use the embedding error path: {error}"
808        );
809        assert!(
810            service.worker.get().is_none(),
811            "oversized work must be rejected before the worker queue is initialized"
812        );
813    }
814
815    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
816    async fn blocking_adapter_byte_budget_rejects_excess_and_queued_jobs_complete() {
817        const ADMITTED_JOBS: usize = 4;
818        let byte_budget = ADMITTED_JOBS * EMBEDDING_MAX_JOB_BYTES;
819        let inner = Arc::new(BlockingTestService::new());
820        let service = Arc::new(BlockingEmbeddingService::with_byte_budget(
821            Arc::clone(&inner),
822            byte_budget,
823        ));
824        let max_texts = Arc::new(vec![
825            "x".repeat(lattice_embed::MAX_TEXT_BYTES);
826            lattice_embed::DEFAULT_MAX_BATCH_SIZE
827        ]);
828        let mut admitted = Vec::with_capacity(ADMITTED_JOBS);
829
830        for _ in 0..ADMITTED_JOBS {
831            let service = Arc::clone(&service);
832            let texts = Arc::clone(&max_texts);
833            admitted.push(tokio::spawn(async move {
834                service.embed(&texts, EmbeddingModel::default()).await
835            }));
836        }
837        tokio::time::timeout(Duration::from_secs(1), async {
838            while service.in_flight_bytes.load(Ordering::Acquire) < byte_budget {
839                tokio::task::yield_now().await;
840            }
841        })
842        .await
843        .expect("all jobs within the byte budget must be admitted");
844
845        let overflow = tokio::time::timeout(
846            Duration::from_millis(100),
847            service.embed(&max_texts, EmbeddingModel::default()),
848        )
849        .await
850        .expect("a byte-budget overflow must fail without waiting")
851        .expect_err("a byte-budget overflow must return an embedding error");
852
853        inner.release();
854        for call in admitted {
855            call.await
856                .expect("admitted embedding task must not panic")
857                .expect("admitted embedding job must complete");
858        }
859        assert!(
860            overflow.to_string().contains("byte budget"),
861            "byte saturation must use the embedding failure path: {overflow}"
862        );
863    }
864
865    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
866    async fn blocking_adapter_releases_byte_budget_after_completion_and_skipped_job() {
867        let byte_budget = "first".len() + "abandoned".len();
868        let inner = Arc::new(BlockingTestService::new());
869        let service = Arc::new(BlockingEmbeddingService::with_byte_budget(
870            Arc::clone(&inner),
871            byte_budget,
872        ));
873
874        let first_service = Arc::clone(&service);
875        let first = tokio::spawn(async move {
876            first_service
877                .embed(&["first".to_owned()], EmbeddingModel::default())
878                .await
879        });
880        tokio::time::timeout(Duration::from_secs(1), async {
881            while inner.entered.load(Ordering::Acquire) == 0 {
882                tokio::task::yield_now().await;
883            }
884        })
885        .await
886        .expect("first embedding call must occupy the native worker");
887
888        let abandoned = tokio::time::timeout(
889            Duration::from_millis(50),
890            service.embed(&["abandoned".to_owned()], EmbeddingModel::default()),
891        )
892        .await;
893        assert!(abandoned.is_err(), "queued embedding call must time out");
894        assert_eq!(
895            service.in_flight_bytes.load(Ordering::Acquire),
896            byte_budget,
897            "running and queued jobs must both consume the byte budget"
898        );
899
900        inner.release();
901        first
902            .await
903            .expect("first embedding task must not panic")
904            .expect("first embedding call must succeed");
905        tokio::time::timeout(Duration::from_secs(1), async {
906            while service.in_flight_bytes.load(Ordering::Acquire) != 0 {
907                tokio::task::yield_now().await;
908            }
909        })
910        .await
911        .expect("completed and skipped jobs must release their byte reservations");
912
913        let later = service
914            .embed(&["later".to_owned()], EmbeddingModel::default())
915            .await
916            .expect("a later call must succeed after the byte budget is released");
917        assert_eq!(later, vec![vec![1.0]]);
918    }
919
920    #[test]
921    fn register_and_get_provider_round_trip() {
922        let mut reg = EmbedderRegistry::new();
923        reg.register(ConstVecProvider::new("mock-384", 384));
924
925        assert!(reg.contains("mock-384"), "registered name must be present");
926        let provider = reg.get_provider("mock-384").expect("provider must exist");
927        assert_eq!(provider.name(), "mock-384");
928        assert_eq!(provider.dimensions(), 384);
929    }
930
931    #[test]
932    fn duplicate_name_last_wins() {
933        let mut reg = EmbedderRegistry::new();
934        reg.register(ConstVecProvider::new("shared", 128));
935        reg.register(ConstVecProvider::new("shared", 256));
936
937        let provider = reg.get_provider("shared").expect("provider must exist");
938        assert_eq!(
939            provider.dimensions(),
940            256,
941            "last registration must win; expected dims=256"
942        );
943    }
944
945    #[test]
946    fn names_returns_all_registered() {
947        let mut reg = EmbedderRegistry::new();
948        reg.register(ConstVecProvider::new("model-a", 64));
949        reg.register(ConstVecProvider::new("model-b", 128));
950        reg.register(ConstVecProvider::new("model-c", 256));
951
952        let mut names = reg.names();
953        names.sort();
954        assert_eq!(names, vec!["model-a", "model-b", "model-c"]);
955    }
956
957    #[tokio::test]
958    async fn get_service_unknown_name_returns_error() {
959        let reg = EmbedderRegistry::new();
960        let result = reg.get_service("does-not-exist").await;
961        let err = result.err().expect("expected Err for unknown name, got Ok");
962        assert!(
963            matches!(err, RuntimeError::UnknownModel(ref n) if n == "does-not-exist"),
964            "expected UnknownModel, got {err:?}"
965        );
966    }
967
968    #[tokio::test]
969    async fn get_service_calls_build_once() {
970        let counter = Arc::new(AtomicUsize::new(0));
971        let provider = ConstVecProvider {
972            name: "cached-model".to_owned(),
973            dims: 32,
974            build_calls: Arc::clone(&counter),
975        };
976        let mut reg = EmbedderRegistry::new();
977        reg.register(provider);
978
979        let _ = reg.get_service("cached-model").await.unwrap();
980        let _ = reg.get_service("cached-model").await.unwrap();
981        let _ = reg.get_service("cached-model").await.unwrap();
982
983        assert_eq!(
984            counter.load(Ordering::SeqCst),
985            1,
986            "build must be called exactly once regardless of get_service call count"
987        );
988    }
989
990    struct SlowBuildProvider {
991        name: String,
992        dims: usize,
993        build_calls: Arc<AtomicUsize>,
994    }
995
996    #[async_trait]
997    impl EmbedderProvider for SlowBuildProvider {
998        fn name(&self) -> &str {
999            &self.name
1000        }
1001
1002        fn dimensions(&self) -> usize {
1003            self.dims
1004        }
1005
1006        async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1007            self.build_calls.fetch_add(1, Ordering::SeqCst);
1008            tokio::time::sleep(Duration::from_millis(50)).await;
1009            Ok(Arc::new(ConstVecService { dims: self.dims }))
1010        }
1011    }
1012
1013    #[tokio::test(flavor = "multi_thread", worker_threads = 8)]
1014    async fn concurrent_cold_resolutions_single_flight_one_build() {
1015        const CALLERS: usize = 16;
1016        let counter = Arc::new(AtomicUsize::new(0));
1017        let mut reg = EmbedderRegistry::new();
1018        reg.register(SlowBuildProvider {
1019            name: "cold-model".to_owned(),
1020            dims: 8,
1021            build_calls: Arc::clone(&counter),
1022        });
1023        let reg = Arc::new(reg);
1024
1025        let mut callers = Vec::with_capacity(CALLERS);
1026        for _ in 0..CALLERS {
1027            let reg = Arc::clone(&reg);
1028            callers.push(tokio::spawn(
1029                async move { reg.get_service("cold-model").await },
1030            ));
1031        }
1032
1033        for caller in callers {
1034            let service = caller
1035                .await
1036                .expect("resolution task must not panic")
1037                .expect("every concurrent cold resolution must receive a working service");
1038            assert_eq!(service.name(), "const-vec-service");
1039        }
1040
1041        assert_eq!(
1042            counter.load(Ordering::SeqCst),
1043            1,
1044            "concurrent cold resolutions must share a single in-flight build()"
1045        );
1046    }
1047}